From 893d610891cf727f34885cb7e2a2da444e4cc280 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Bla=C5=BE=20Dular?= <22869613+xBlaz3kx@users.noreply.github.com> Date: Sun, 14 Dec 2025 11:35:54 +0100 Subject: [PATCH] chore: improve DSN handling for SQLite (#26011) --- server/db/db.go | 35 ++++++++++++++++++++---- server/db/db_test.go | 65 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 94 insertions(+), 6 deletions(-) create mode 100644 server/db/db_test.go diff --git a/server/db/db.go b/server/db/db.go index 8df101419..19b5011f2 100644 --- a/server/db/db.go +++ b/server/db/db.go @@ -22,21 +22,44 @@ func New(driver, dsn string) (*gorm.DB, error) { switch driver { case "sqlite": - file, err := homedir.Expand(dsn) + + // Example DSNs: + //"path/to/database.db" + // "~/database.db", + // "database.db?cache=shared&journal_mode=WAL" + // ":memory:" + + // Split database path and connection parameters + dbPath, connectionParams, _ := strings.Cut(dsn, "?") + + file, err := homedir.Expand(dbPath) if err != nil { return nil, err } + if err := os.MkdirAll(filepath.Dir(file), 0700); err != nil { + return nil, err + } + // Store the expanded file path for later use FilePath = file - if err := os.MkdirAll(filepath.Dir(file), 0700); err != nil { - return nil, err + + // Add busy_timeout pragma if not already present + if !strings.Contains(connectionParams, "_pragma=busy_timeout") { + // Append '&' if there are existing connection parameters + if len(connectionParams) > 0 { + connectionParams += "&" + } + + // Add busy_timeout pragma to connection parameters + connectionParams += "_pragma=busy_timeout(5000)" } - util.NewLogger("main").INFO.Println("using sqlite database:", file) + connectionStr := file + "?" + connectionParams - // avoid busy errors - dialect = sqlite.Open(file + "?_pragma=busy_timeout(5000)") + util.NewLogger("main").INFO.Println("using sqlite database:", connectionStr) + + dialect = sqlite.Open(connectionStr) // case "postgres": // dialect = postgres.Open(dsn) // case "mysql": diff --git a/server/db/db_test.go b/server/db/db_test.go new file mode 100644 index 000000000..4695cf2b4 --- /dev/null +++ b/server/db/db_test.go @@ -0,0 +1,65 @@ +package db + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestUnitNewDriver(t *testing.T) { + tmpDir := t.TempDir() + + tests := []struct { + name string + driver string + dsn string + expectedFilePath string + wantErr bool + }{ + { + name: "SQLite In-Memory", + driver: "sqlite", + dsn: ":memory:", + expectedFilePath: ":memory:", + wantErr: false, + }, + { + name: "SQLite File", + driver: "sqlite", + dsn: tmpDir + "/evcc.db", + expectedFilePath: tmpDir + "/evcc.db", + wantErr: false, + }, + { + name: "SQLite with connection parameters", + driver: "sqlite", + dsn: tmpDir + "evcc.db?_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)", + wantErr: false, + expectedFilePath: tmpDir + "evcc.db", + }, + { + name: "Unsupported Driver", + driver: "postgresql", + dsn: "/var/lib/evcc/evcc.db", + wantErr: true, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + // Reset file path + FilePath = "" + + driver, err := New(test.driver, test.dsn) + if test.wantErr { + assert.Error(t, err) + assert.Nil(t, driver) + } else { + assert.NoError(t, err) + assert.NotNil(t, driver) + } + + assert.Equal(t, test.expectedFilePath, FilePath) + }) + } +}