diff --git a/api/database/database.go b/api/database/database.go index 417b0878..78d29d45 100644 --- a/api/database/database.go +++ b/api/database/database.go @@ -5,9 +5,9 @@ import ( "log" "net/url" "os" - "strings" "time" + "github.com/photoview/photoview/api/database/drivers" "github.com/photoview/photoview/api/graphql/models" "github.com/pkg/errors" @@ -66,9 +66,8 @@ func SetupDatabase() (*gorm.DB, error) { config.Logger = logger.Default.LogMode(logger.Info) var databaseDialect gorm.Dialector - databaseDriver := strings.ToLower(os.Getenv("PHOTOVIEW_DATABASE_DRIVER")) - switch databaseDriver { - case "mysql": + switch drivers.DatabaseDriver() { + case drivers.DatabaseDriverMysql: mysqlAddress, err := getMysqlAddress() if err != nil { return nil, err @@ -76,7 +75,7 @@ func SetupDatabase() (*gorm.DB, error) { log.Printf("Connecting to database: %s", mysqlAddress) databaseDialect = mysql.Open(mysqlAddress.String()) - case "sqlite": + case drivers.DatabaseDriverSqlite: sqliteAddress, err := getSqliteAddress() if err != nil { return nil, err diff --git a/api/database/drivers/database_drivers.go b/api/database/drivers/database_drivers.go new file mode 100644 index 00000000..140cdd26 --- /dev/null +++ b/api/database/drivers/database_drivers.go @@ -0,0 +1,28 @@ +package drivers + +import ( + "os" + "strings" +) + +type DatabaseDriverType string + +const ( + DatabaseDriverMysql DatabaseDriverType = "mysql" + DatabaseDriverSqlite DatabaseDriverType = "sqlite" +) + +func DatabaseDriver() DatabaseDriverType { + + var driver DatabaseDriverType + driverString := strings.ToLower(os.Getenv("PHOTOVIEW_DATABASE_DRIVER")) + + switch driverString { + case "mysql": + driver = DatabaseDriverMysql + case "sqlite": + driver = DatabaseDriverSqlite + } + + return driver +} diff --git a/api/graphql/models/site_info.go b/api/graphql/models/site_info.go index 43ddb7f8..76d2dd7e 100644 --- a/api/graphql/models/site_info.go +++ b/api/graphql/models/site_info.go @@ -1,6 +1,7 @@ package models import ( + db_drivers "github.com/photoview/photoview/api/database/drivers" "github.com/pkg/errors" "gorm.io/gorm" ) @@ -22,10 +23,16 @@ func GetSiteInfo(db *gorm.DB) (*SiteInfo, error) { if err := db.First(&siteInfo).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { + + defaultConcurrentWorkers := 3 + if db_drivers.DatabaseDriver() == db_drivers.DatabaseDriverSqlite { + defaultConcurrentWorkers = 1 + } + siteInfo = SiteInfo{ InitialSetup: true, PeriodicScanInterval: 0, - ConcurrentWorkers: 3, + ConcurrentWorkers: defaultConcurrentWorkers, } if err := db.Create(&siteInfo).Error; err != nil { diff --git a/api/graphql/resolvers/scanner.go b/api/graphql/resolvers/scanner.go index 179f1be3..5eca1cb5 100644 --- a/api/graphql/resolvers/scanner.go +++ b/api/graphql/resolvers/scanner.go @@ -4,6 +4,7 @@ import ( "context" "time" + "github.com/photoview/photoview/api/database/drivers" "github.com/photoview/photoview/api/graphql/models" "github.com/photoview/photoview/api/scanner" "github.com/pkg/errors" @@ -66,6 +67,10 @@ func (r *mutationResolver) SetScannerConcurrentWorkers(ctx context.Context, work return 0, errors.New("concurrent workers must at least be 1") } + if workers > 1 && drivers.DatabaseDriver() == drivers.DatabaseDriverSqlite { + return 0, errors.New("multiple workers not supported for SQLite databases") + } + if err := r.Database.Session(&gorm.Session{AllowGlobalUpdate: true}).Model(&models.SiteInfo{}).Update("concurrent_workers", workers).Error; err != nil { return 0, err }