// Package db owns the SQLite schema and all queries. Migrations are
// embedded in the binary and applied at startup.
package db
import (
"database/sql"
"embed"
"fmt"
"io/fs"
"os"
"path/filepath"
"sort"
"strings"
_ "modernc.org/sqlite"
)
//go:embed migrations/*.sql
var migrationsFS embed.FS
// Open opens the SQLite database at path, applies any pending migrations,
// and returns a ready connection pool.
func Open(path string) (*sql.DB, error) {
dsn := "file:" + path + "?_pragma=foreign_keys(1)&_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)"
database, err := sql.Open("sqlite", dsn)
if err != nil {
return nil, fmt.Errorf("open sqlite %s: %w", path, err)
}
if err := database.Ping(); err != nil {
database.Close()
return nil, fmt.Errorf("open sqlite %s: %w", path, err)
}
// SQLite creates the file world-readable by default; it holds session
// tokens and password hashes, so tighten it regardless of umask.
if err := os.Chmod(path, 0o600); err != nil && !os.IsNotExist(err) {
database.Close()
return nil, fmt.Errorf("chmod sqlite %s: %w", path, err)
}
if err := migrate(database); err != nil {
database.Close()
return nil, err
}
return database, nil
}
// migrate applies every embedded migration not yet recorded in
// schema_migrations, in filename order, each in its own transaction.
func migrate(database *sql.DB) error {
if _, err := database.Exec(`CREATE TABLE IF NOT EXISTS schema_migrations (
version TEXT PRIMARY KEY,
applied_at INTEGER NOT NULL
)`); err != nil {
return fmt.Errorf("create schema_migrations: %w", err)
}
names, err := fs.Glob(migrationsFS, "migrations/*.sql")
if err != nil {
return fmt.Errorf("list migrations: %w", err)
}
sort.Strings(names)
for _, name := range names {
version := filepath.Base(name)
var applied int
if err := database.QueryRow(
`SELECT COUNT(*) FROM schema_migrations WHERE version = ?`, version,
).Scan(&applied); err != nil {
return fmt.Errorf("check migration %s: %w", version, err)
}
if applied > 0 {
continue
}
script, err := migrationsFS.ReadFile(name)
if err != nil {
return fmt.Errorf("read migration %s: %w", version, err)
}
if err := apply(database, version, string(script)); err != nil {
return err
}
}
return nil
}
func apply(database *sql.DB, version, script string) error {
tx, err := database.Begin()
if err != nil {
return fmt.Errorf("begin migration %s: %w", version, err)
}
defer tx.Rollback()
if _, err := tx.Exec(script); err != nil {
return fmt.Errorf("apply migration %s: %w", version, err)
}
if _, err := tx.Exec(
`INSERT INTO schema_migrations (version, applied_at) VALUES (?, unixepoch())`, version,
); err != nil {
return fmt.Errorf("record migration %s: %w", version, err)
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("commit migration %s: %w", version, err)
}
return nil
}
// rowScanner is satisfied by *sql.Row and *sql.Rows, letting entity
// scanners serve both single-row lookups and list queries.
type rowScanner interface {
Scan(dest ...any) error
}
// listQuery runs query and appends each row scanned by scan.
func listQuery[T any](database *sql.DB, query string, args []any, scan func(rowScanner, *T) error) ([]T, error) {
rows, err := database.Query(query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var items []T
for rows.Next() {
var item T
if err := scan(rows, &item); err != nil {
return nil, err
}
items = append(items, item)
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
// execScoped runs a scoped UPDATE or DELETE and maps zero affected rows to
// ErrNotFound, so callers never mistake "row absent" for success.
func execScoped(database *sql.DB, op, query string, args ...any) error {
result, err := database.Exec(query, args...)
if err != nil {
return fmt.Errorf("%s: %w", op, err)
}
affected, err := result.RowsAffected()
if err != nil {
return fmt.Errorf("%s: %w", op, err)
}
if affected == 0 {
return fmt.Errorf("%s: %w", op, ErrNotFound)
}
return nil
}
// execLastID runs an INSERT and returns the new row id.
func execLastID(database *sql.DB, query string, args ...any) (int64, error) {
result, err := database.Exec(query, args...)
if err != nil {
return 0, err
}
return result.LastInsertId()
}
// isUniqueViolation reports whether err is SQLite's UNIQUE constraint
// failure — a per-repo number collision when it surfaces from an insert.
func isUniqueViolation(err error) bool {
return err != nil && strings.Contains(err.Error(), "UNIQUE constraint failed")
}