diff --git a/cmd/root.go b/cmd/root.go index ee63f278b..9e3a28921 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -400,7 +400,7 @@ func runRoot(cmd *cobra.Command, args []string) { // publish system infos valueChan <- util.Param{Key: keys.Version, Val: util.FormattedVersion()} valueChan <- util.Param{Key: keys.Config, Val: viper.ConfigFileUsed()} - valueChan <- util.Param{Key: keys.Database, Val: db.FilePath} + valueChan <- util.Param{Key: keys.Database, Val: db.FilePath()} valueChan <- util.Param{Key: keys.System, Val: util.System()} valueChan <- util.Param{Key: keys.Timezone, Val: time.Now().Format("MST -07:00")} valueChan <- util.Param{Key: keys.Experimental, Val: isExperimental()} diff --git a/go.mod b/go.mod index e03d87245..17f1781a0 100644 --- a/go.mod +++ b/go.mod @@ -70,7 +70,7 @@ require ( github.com/koron/go-ssdp v0.1.0 github.com/korylprince/ipnetgen v1.0.1 github.com/libp2p/zeroconf/v2 v2.2.0 - github.com/libtnb/sqlite v1.0.4 + github.com/libtnb/sqlite v1.1.1 github.com/lorenzodonini/ocpp-go v0.19.0 github.com/lunixbochs/struc v0.0.0-20241101090106-8d528fa2c543 github.com/mabunixda/wattpilot v1.8.5 diff --git a/go.sum b/go.sum index 68ed65b9e..dd266eb1f 100644 --- a/go.sum +++ b/go.sum @@ -476,8 +476,8 @@ github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ= github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI= github.com/libp2p/zeroconf/v2 v2.2.0 h1:Cup06Jv6u81HLhIj1KasuNM/RHHrJ8T7wOTS4+Tv53Q= github.com/libp2p/zeroconf/v2 v2.2.0/go.mod h1:fuJqLnUwZTshS3U/bMRJ3+ow/v9oid1n0DmyYyNO1Xs= -github.com/libtnb/sqlite v1.0.4 h1:SPFG8aSNNaUjO1ises/AFeWV+q1etBAVrtXPci0Uhvs= -github.com/libtnb/sqlite v1.0.4/go.mod h1:PNlcqMVWU+pMIUR499fHvRLUdxrNVVVXwEZkW1c9stE= +github.com/libtnb/sqlite v1.1.1 h1:p1Ud8vHTsPzgmChPUuJFkesrzl9ajD/pYrmoWf1cgig= +github.com/libtnb/sqlite v1.1.1/go.mod h1:PNlcqMVWU+pMIUR499fHvRLUdxrNVVVXwEZkW1c9stE= github.com/lightstep/lightstep-tracer-common/golang/gogo v0.0.0-20190605223551-bc2310a04743/go.mod h1:qklhhLq1aX+mtWk9cPHPzaBjWImj5ULL6C7HFJtXQMM= github.com/lightstep/lightstep-tracer-go v0.18.1/go.mod h1:jlF1pusYV4pidLvZ+XD0UBX0ZE6WURAspgAczcDHrL4= github.com/lunixbochs/struc v0.0.0-20241101090106-8d528fa2c543 h1:GxMuVb9tJajC1QpbQwYNY1ZAo1EIE8I+UclBjOfjz/M= diff --git a/server/db/db.go b/server/db/db.go index 41ebb4032..cf92b321a 100644 --- a/server/db/db.go +++ b/server/db/db.go @@ -17,9 +17,13 @@ import ( var ( Instance *gorm.DB - FilePath string // Store the actual SQLite file path + filePath string // Store the actual SQLite file path ) +func FilePath() string { + return filePath +} + func New(driver, dsn string) (*gorm.DB, error) { var dialect gorm.Dialector @@ -45,12 +49,13 @@ func New(driver, dsn string) (*gorm.DB, error) { } // Store the expanded file path for later use - FilePath = file + filePath = file - addParam := func(typ, param string) { + // TODO WAL mode "journal_mode(WAL)", "synchronous(NORMAL)" + for _, pragma := range []string{"foreign_keys(1)", "auto_vacuum(INCREMENTAL)"} { // Add busy_timeout pragma if not already present - if short, _, _ := strings.Cut(param, "("); strings.Contains(params, typ+"="+short) { - return + if short, _, _ := strings.Cut(pragma, "("); strings.Contains(params, "_pragma="+short) { + continue } // Append '&' if there are existing connection parameters @@ -59,17 +64,9 @@ func New(driver, dsn string) (*gorm.DB, error) { } // Add busy_timeout pragma to connection parameters - params += typ + "=" + param + params += "_pragma=" + pragma } - // TODO "foreign_keys(1)" is only set in metrics migrator to ensure home entity exists - for _, pragma := range []string{"busy_timeout(5000)", "synchronous(NORMAL)"} { - addParam("_pragma", pragma) - } - - // https://github.com/libtnb/sqlite/issues/15 - addParam("_time_format", "sqlite") - connectionStr := file + "?" + params util.NewLogger("main").INFO.Println("using sqlite database:", connectionStr) @@ -116,7 +113,12 @@ func Close() error { return db.Close() } -func Backup(ctx context.Context, target string) error { +type backuper interface { + NewBackup(string) (*sqlite3.Backup, error) + NewRestore(string) (*sqlite3.Backup, error) +} + +func runWithBackuper(ctx context.Context, fun func(backuper) (*sqlite3.Backup, error)) error { live, err := Instance.DB() if err != nil { return err @@ -129,17 +131,12 @@ func Backup(ctx context.Context, target string) error { defer conn.Close() return conn.Raw(func(driverConn any) error { - type backuper interface { - NewBackup(string) (*sqlite3.Backup, error) - NewRestore(string) (*sqlite3.Backup, error) - } - conn, ok := driverConn.(backuper) if !ok { return errors.New("invalid db type") } - bck, err := conn.NewBackup(target) + bck, err := fun(conn) if err != nil { return err } @@ -151,3 +148,15 @@ func Backup(ctx context.Context, target string) error { return bck.Finish() }) } + +func Backup(ctx context.Context, target string) error { + return runWithBackuper(ctx, func(conn backuper) (*sqlite3.Backup, error) { + return conn.NewBackup(target) + }) +} + +func Restore(ctx context.Context, target string) error { + return runWithBackuper(ctx, func(conn backuper) (*sqlite3.Backup, error) { + return conn.NewRestore(target) + }) +} diff --git a/server/db/db_test.go b/server/db/db_test.go index 4695cf2b4..3366c961f 100644 --- a/server/db/db_test.go +++ b/server/db/db_test.go @@ -48,7 +48,7 @@ func TestUnitNewDriver(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { // Reset file path - FilePath = "" + filePath = "" driver, err := New(test.driver, test.dsn) if test.wantErr { @@ -59,7 +59,7 @@ func TestUnitNewDriver(t *testing.T) { assert.NotNil(t, driver) } - assert.Equal(t, test.expectedFilePath, FilePath) + assert.Equal(t, test.expectedFilePath, FilePath()) }) } } diff --git a/server/http_site_handler.go b/server/http_site_handler.go index 8a89323ad..2cd0dd9de 100644 --- a/server/http_site_handler.go +++ b/server/http_site_handler.go @@ -1,9 +1,9 @@ package server import ( + "context" "encoding/json" "errors" - "fmt" "io" "io/fs" "net/http" @@ -340,95 +340,130 @@ func getBackup(authObject auth.Auth) http.HandlerFunc { return } - f, err := os.Open(db.FilePath) + filename := "evcc-backup-" + time.Now().Format("2006-01-02--15-04") + ".db" + + tmpFile, err := os.CreateTemp("", "evcc-backup-*.db") if err != nil { - http.Error(w, "Could not open DB file: "+err.Error(), http.StatusInternalServerError) + http.Error(w, "Creating backup failed", http.StatusInternalServerError) + return + } + tmpName := tmpFile.Name() + tmpFile.Close() + defer os.Remove(tmpName) + + if err := db.Backup(r.Context(), tmpName); err != nil { + http.Error(w, "Backup failed", http.StatusInternalServerError) + return + } + + f, err := os.Open(tmpName) + if err != nil { + http.Error(w, "Opening backup failed", http.StatusInternalServerError) return } defer f.Close() - filename := "evcc-backup-" + time.Now().Format("2006-01-02--15-04") + ".db" - w.Header().Set("Content-Type", "application/octet-stream") w.Header().Set("Content-Disposition", `attachment; filename="`+filename+`"`) if _, err := io.Copy(w, f); err != nil { - http.Error(w, "Error streaming DB file: "+err.Error(), http.StatusInternalServerError) + http.Error(w, "Streaming backup failed", http.StatusInternalServerError) return } } } // createLocalDatabaseBackup creates a local backup in case of catastrophic error in reset or restore -func createLocalDatabaseBackup() error { - backupPath := db.FilePath + ".bak" - - src, err := os.Open(db.FilePath) - if err != nil { - return fmt.Errorf("failed to open database file: %w", err) - } - defer src.Close() - - dst, err := os.Create(backupPath) - if err != nil { - return fmt.Errorf("failed to create backup file: %w", err) - } - defer dst.Close() - - if _, err := io.Copy(dst, src); err != nil { - // clean up partial backup on error - os.Remove(backupPath) - return fmt.Errorf("failed to copy database: %w", err) - } - - return nil +func createLocalDatabaseBackup(ctx context.Context) error { + return db.Backup(ctx, db.FilePath()+".bak") } func restoreDatabase(authObject auth.Auth, shutdown func()) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - // Parse multipart form - err := r.ParseMultipartForm(32 << 20) // 32MB max memory + // cap upload size to bound disk usage + r.Body = http.MaxBytesReader(w, r.Body, 256<<20) + + mr, err := r.MultipartReader() if err != nil { http.Error(w, "Failed to parse form: "+err.Error(), http.StatusBadRequest) return } - if !adminPasswordValid(authObject, r.FormValue("password")) { + var ( + password string + tmpName string + ) + + for { + part, err := mr.NextPart() + if err == io.EOF { + break + } + if err != nil { + http.Error(w, "Upload failed", http.StatusBadRequest) + return + } + + switch part.FormName() { + case "password": + b, err := io.ReadAll(io.LimitReader(part, 1<<10)) + part.Close() + if err != nil { + http.Error(w, "Upload failed", http.StatusBadRequest) + return + } + password = string(b) + + case "file": + tmpFile, err := os.CreateTemp("", "evcc-restore-*.db") + if err != nil { + part.Close() + http.Error(w, "Failed to create temp file", http.StatusInternalServerError) + return + } + tmpName = tmpFile.Name() + defer os.Remove(tmpName) + + _, copyErr := io.Copy(tmpFile, part) + closeErr := tmpFile.Close() + part.Close() + if copyErr != nil || closeErr != nil { + http.Error(w, "Upload failed", http.StatusBadRequest) + return + } + + default: + part.Close() + } + } + + if !adminPasswordValid(authObject, password) { http.Error(w, "Invalid password", http.StatusUnauthorized) return } - file, _, err := r.FormFile("file") - if err != nil { - http.Error(w, "Failed to get uploaded file: "+err.Error(), http.StatusBadRequest) + if tmpName == "" { + http.Error(w, "Missing file", http.StatusBadRequest) return } - defer file.Close() settings.Persist() - // close db connection to avoid corruption - if err := db.Close(); err != nil { - jsonError(w, http.StatusInternalServerError, err) - return - } - // create local backup before overwriting - if err := createLocalDatabaseBackup(); err != nil { - http.Error(w, "Failed to create local backup: "+err.Error(), http.StatusInternalServerError) + if err := createLocalDatabaseBackup(r.Context()); err != nil { + http.Error(w, "Backup failed", http.StatusInternalServerError) return } - // overwrite DB file - f, err := os.Create(db.FilePath) - if err != nil { - http.Error(w, "Could not open DB file for writing: "+err.Error(), http.StatusInternalServerError) + if err := db.Restore(r.Context(), tmpName); err != nil { + http.Error(w, "Restore failed", http.StatusInternalServerError) return } - defer f.Close() - if _, err := io.Copy(f, file); err != nil { - http.Error(w, "Failed to write DB file: "+err.Error(), http.StatusInternalServerError) + // close DB so the WAL is checkpointed into the main file before shutdown + // hooks (e.g. settings.Persist) can overwrite the restored content + if err := db.Close(); err != nil { + http.Error(w, "DB close failed", http.StatusInternalServerError) return } @@ -457,7 +492,7 @@ func resetDatabase(authObject auth.Auth, shutdown func()) http.HandlerFunc { settings.Persist() - if err := createLocalDatabaseBackup(); err != nil { + if err := createLocalDatabaseBackup(r.Context()); err != nil { jsonError(w, http.StatusInternalServerError, err) return }