162 lines
3.3 KiB
Go
162 lines
3.3 KiB
Go
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)
|
|
}
|
|
|
|
return gorm.Open(dialect, &gorm.Config{
|
|
Logger: &Logger{util.NewLogger("db")},
|
|
})
|
|
}
|
|
|
|
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)
|
|
})
|
|
}
|