package db import ( "context" "errors" "fmt" "os" "path/filepath" "strings" "github.com/evcc-io/evcc/util" "github.com/libtnb/sqlite" "github.com/mitchellh/go-homedir" "gorm.io/gorm" sqlite3 "modernc.org/sqlite" sqlite3lib "modernc.org/sqlite/lib" ) var ( Instance *gorm.DB 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 switch driver { case "sqlite": // Split database path and connection parameters dbPath, params, _ := 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 // TODO WAL mode "journal_mode(WAL)", "synchronous(NORMAL)" for _, pragma := range []string{"busy_timeout(5000)", "foreign_keys(1)", "auto_vacuum(INCREMENTAL)"} { // add pragma if not already present if short, _, _ := strings.Cut(pragma, "("); strings.Contains(params, "_pragma="+short) { continue } // append '&' if there are existing connection parameters if len(params) > 0 { params += "&" } params += "_pragma=" + pragma } connectionStr := file + "?" + params util.NewLogger("main").INFO.Println("using sqlite database:", connectionStr) dialect = sqlite.Open(connectionStr) // case "postgres": // dialect = postgres.Open(dsn) // case "mysql": // dialect = mysql.Open(dsn) default: return nil, fmt.Errorf("invalid database type: %s not in [sqlite]", driver) } db, err := gorm.Open(dialect, &gorm.Config{ Logger: &Logger{util.NewLogger("db")}, }) if err != nil { return nil, err } // sqlite allows a single writer; serialize on one connection so concurrent // writes wait on busy_timeout instead of failing with SQLITE_BUSY. if sqlDB, err := db.DB(); err == nil { sqlDB.SetMaxOpenConns(1) } return db, nil } func NewInstance(driver, dsn string) error { db, err := New(strings.ToLower(driver), dsn) if err != nil { return err } Instance = db mu.Lock() defer mu.Unlock() for _, f := range registry { if err := f(db); err != nil { return err } } return nil } func Close() error { db, err := Instance.DB() if err != nil { return err } return db.Close() } // IsReadonly reports whether err indicates a database file that is not // writable by the current user. The code mask covers extended result codes. func IsReadonly(err error) bool { var serr *sqlite3.Error return errors.As(err, &serr) && serr.Code()&0xff == sqlite3lib.SQLITE_READONLY } 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 } conn, err := live.Conn(ctx) if err != nil { return err } defer conn.Close() return conn.Raw(func(driverConn any) error { conn, ok := driverConn.(backuper) if !ok { return errors.New("invalid db type") } bck, err := fun(conn) if err != nil { return err } if _, err := bck.Step(-1); err != nil { return err } 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) }) }