josie / simplegit

// 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")
}