package db import ( "crypto/sha256" "database/sql" "encoding/hex" "fmt" "os" "path/filepath" "sort" "strings" _ "modernc.org/sqlite" ) type DB struct { *sql.DB } func (d *DB) SQLDB() *sql.DB { return d.DB } func Open(dbPath string) (*DB, error) { if err := os.MkdirAll(filepath.Dir(dbPath), 0700); err != nil { return nil, err } db, err := sql.Open("sqlite", dbPath+"?_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)") if err != nil { return nil, fmt.Errorf("opening db: %w", err) } if _, err := db.Exec("PRAGMA foreign_keys = ON"); err != nil { db.Close() return nil, fmt.Errorf("enabling foreign_keys: %w", err) } return &DB{db}, nil } func (db *DB) Close() error { return db.DB.Close() } func (db *DB) RunMigrations() error { return db.runMigrationsInternal(migrationsFS, "migrations") } func (db *DB) runMigrationsInternal(mfs embedFS, migrationsRoot string) error { entries, err := mfs.ReadDir(migrationsRoot) if err != nil { return fmt.Errorf("reading migrations dir: %w", err) } var names []string for _, e := range entries { if !e.IsDir() && strings.HasSuffix(e.Name(), ".sql") { names = append(names, e.Name()) } } sort.Strings(names) if _, err := db.Exec(` CREATE TABLE IF NOT EXISTS schema_migrations ( version TEXT PRIMARY KEY, checksum TEXT, applied_at DATETIME DEFAULT CURRENT_TIMESTAMP ) `); err != nil { return fmt.Errorf("creating schema_migrations table: %w", err) } hasChecksumCol := false if rows, err := db.Query("PRAGMA table_info(schema_migrations)"); err == nil { for rows.Next() { var cid int var cname string rows.Scan(&cid, &cname, new(string), new(int), new(interface{}), new(int)) if cname == "checksum" { hasChecksumCol = true } } rows.Close() } for _, name := range names { if hasChecksumCol { var storedChecksum string row := db.QueryRow("SELECT checksum FROM schema_migrations WHERE version = ?", name) if err := row.Scan(&storedChecksum); err == nil && storedChecksum != "" { continue } } else { var count int row := db.QueryRow("SELECT COUNT(*) FROM schema_migrations WHERE version = ?", name) if err := row.Scan(&count); err == nil && count > 0 { continue } } data, err := mfs.ReadFile(filepath.Join(migrationsRoot, name)) if err != nil { return fmt.Errorf("reading migration %s: %w", name, err) } checksum := sha256.Sum256(data) checksumHex := hex.EncodeToString(checksum[:]) tx, err := db.Begin() if err != nil { return fmt.Errorf("starting transaction for migration %s: %w", name, err) } if _, err := tx.Exec(string(data)); err != nil { tx.Rollback() return fmt.Errorf("applying migration %s: %w", name, err) } if hasChecksumCol { if _, err := tx.Exec( "INSERT INTO schema_migrations (version, checksum) VALUES (?, ?)", name, checksumHex, ); err != nil { tx.Rollback() return fmt.Errorf("recording migration %s: %w", name, err) } } else { if _, err := tx.Exec( "INSERT INTO schema_migrations (version) VALUES (?)", name, ); err != nil { tx.Rollback() return fmt.Errorf("recording migration %s: %w", name, err) } } if err := tx.Commit(); err != nil { return fmt.Errorf("committing migration %s: %w", name, err) } } return nil } type embedFS interface { ReadDir(name string) ([]os.DirEntry, error) ReadFile(name string) ([]byte, error) }