From 10625c3757953592ebe93401b0e2cc448b200159 Mon Sep 17 00:00:00 2001 From: Aly Raffauf Date: Sun, 9 Aug 2026 14:55:57 -0400 Subject: [PATCH] feat: migrate state schema --- internal/state/migrations.go | 81 ++++++++++++++ internal/state/migrations/001_initial.sql | 49 +++++++++ internal/state/migrations_test.go | 127 ++++++++++++++++++++++ 3 files changed, 257 insertions(+) create mode 100644 internal/state/migrations.go create mode 100644 internal/state/migrations/001_initial.sql create mode 100644 internal/state/migrations_test.go diff --git a/internal/state/migrations.go b/internal/state/migrations.go new file mode 100644 index 0000000..f2a1f26 --- /dev/null +++ b/internal/state/migrations.go @@ -0,0 +1,81 @@ +package state + +import ( + "database/sql" + "fmt" + + _ "embed" +) + +// initialMigrationSQL is the embedded Section 8.4 schema. It is one of the two +// package-variable exceptions permitted by Section 12.1. +// +//go:embed migrations/001_initial.sql +var initialMigrationSQL string + +// currentSchemaVersion is the schema version this build's migrations produce. +// PRAGMA user_version is managed against it. +const currentSchemaVersion = 1 + +// Migrate applies the embedded schema migration to the database when needed. +// Re-running on a current database is a no-op. An unknown newer schema +// (user_version greater than current) is rejected so a newer Cattery never +// silently corrupts an older binary's state. +func Migrate(database *Database) error { + version, err := readUserVersion(database) + if err != nil { + return err + } + if version > currentSchemaVersion { + return errUnknownSchema(version) + } + if version == currentSchemaVersion { + return nil + } + return applyMigration(database) +} + +func applyMigration(database *Database) error { + transaction, err := beginExclusive(database) + if err != nil { + return err + } + if _, err := transaction.Exec(initialMigrationSQL); err != nil { + _ = transaction.Rollback() + return fmt.Errorf("state: apply migration: %w", err) + } + if err := setUserVersion(transaction, currentSchemaVersion); err != nil { + _ = transaction.Rollback() + return err + } + return transaction.Commit() +} + +func beginExclusive(database *Database) (*sql.Tx, error) { + if _, err := database.conn.Exec("PRAGMA locking_mode = EXCLUSIVE"); err != nil { + return nil, fmt.Errorf("state: exclusive locking mode: %w", err) + } + return database.conn.Begin() +} + +func readUserVersion(database *Database) (int, error) { + var version int + if err := database.conn.QueryRow("PRAGMA user_version").Scan(&version); err != nil { + return 0, fmt.Errorf("state: read user_version: %w", err) + } + return version, nil +} + +func setUserVersion(transaction *sql.Tx, version int) error { + statement := fmt.Sprintf("PRAGMA user_version = %d", version) + if _, err := transaction.Exec(statement); err != nil { + return fmt.Errorf("state: set user_version: %w", err) + } + return nil +} + +func errUnknownSchema(version int) error { + return fmt.Errorf( + "state: database schema version %d is newer than supported version %d", + version, currentSchemaVersion) +} diff --git a/internal/state/migrations/001_initial.sql b/internal/state/migrations/001_initial.sql new file mode 100644 index 0000000..466747e --- /dev/null +++ b/internal/state/migrations/001_initial.sql @@ -0,0 +1,49 @@ +CREATE TABLE metadata ( + key TEXT PRIMARY KEY, + value BLOB NOT NULL +); + +CREATE TABLE repositories ( + id INTEGER PRIMARY KEY, + root_path TEXT NOT NULL, + home_path TEXT NOT NULL, + is_default INTEGER NOT NULL DEFAULT 0 CHECK (is_default IN (0, 1)), + created_at TEXT NOT NULL, + last_seen_at TEXT NOT NULL, + UNIQUE (root_path, home_path) +); + +CREATE UNIQUE INDEX repositories_one_default +ON repositories(home_path) +WHERE is_default = 1; + +CREATE TABLE files ( + repository_id INTEGER NOT NULL REFERENCES repositories(id) ON DELETE CASCADE, + target_path TEXT NOT NULL, + group_name TEXT NOT NULL DEFAULT '', + source_path TEXT NOT NULL, + source_kind TEXT NOT NULL CHECK (source_kind IN ('ordinary', 'secret')), + layer TEXT NOT NULL CHECK (layer IN ('base', 'darwin', 'linux')), + baseline_content_hash BLOB NOT NULL CHECK (length(baseline_content_hash) = 32), + baseline_source_hash BLOB NOT NULL CHECK (length(baseline_source_hash) = 32), + executable_bits INTEGER NOT NULL, + status TEXT NOT NULL CHECK (status IN ('active', 'retired')), + applied_at TEXT NOT NULL, + retired_at TEXT, + PRIMARY KEY (repository_id, target_path) +); + +CREATE INDEX files_by_scope +ON files(repository_id, group_name, status); + +CREATE TABLE aliases ( + repository_id INTEGER NOT NULL REFERENCES repositories(id) ON DELETE CASCADE, + alias_path TEXT NOT NULL, + canonical_target_path TEXT NOT NULL, + group_name TEXT NOT NULL DEFAULT '', + layer TEXT NOT NULL CHECK (layer IN ('all', 'darwin', 'linux')), + status TEXT NOT NULL CHECK (status IN ('active', 'retired')), + applied_at TEXT NOT NULL, + retired_at TEXT, + PRIMARY KEY (repository_id, alias_path) +); diff --git a/internal/state/migrations_test.go b/internal/state/migrations_test.go new file mode 100644 index 0000000..3b2183d --- /dev/null +++ b/internal/state/migrations_test.go @@ -0,0 +1,127 @@ +package state + +import ( + "database/sql" + "fmt" + "testing" +) + +func TestStateMigration(t *testing.T) { + scenarios := []struct { + name string + run func(*testing.T) + }{ + {"fresh database applies", testMigrationFresh}, + {"current database is no-op", testMigrationIdempotent}, + {"forced failure rolls back", testMigrationRollback}, + {"interrupted migration recovers", testMigrationRecoversAfterInterruption}, + {"newer schema rejected", testMigrationRejectsNewer}, + } + for _, scenario := range scenarios { + t.Run(scenario.name, scenario.run) + } +} + +func testMigrationFresh(t *testing.T) { + database := openTestDatabase(t) + defer database.Close() + if err := Migrate(database); err != nil { + t.Fatalf("Migrate: %v", err) + } + assertSchemaCurrent(t, database.conn) +} + +func testMigrationIdempotent(t *testing.T) { + database := openTestDatabase(t) + defer database.Close() + if err := Migrate(database); err != nil { + t.Fatalf("first Migrate: %v", err) + } + if err := Migrate(database); err != nil { + t.Fatalf("second Migrate: %v", err) + } + assertSchemaCurrent(t, database.conn) +} + +func testMigrationRollback(t *testing.T) { + database := openTestDatabase(t) + defer database.Close() + execOn(t, database.conn, "CREATE TABLE metadata (key TEXT)") + if err := Migrate(database); err == nil { + t.Fatal("Migrate succeeded against a conflicting schema") + } + if version := userVersion(t, database.conn); version != 0 { + t.Fatalf("user_version = %d after rollback, want 0", version) + } + if tableExists(t, database.conn, "repositories") { + t.Fatal("rollback left a partial repositories table") + } +} + +func testMigrationRecoversAfterInterruption(t *testing.T) { + database := openTestDatabase(t) + defer database.Close() + execOn(t, database.conn, "CREATE TABLE metadata (key TEXT)") + if err := Migrate(database); err == nil { + t.Fatal("Migrate succeeded against an interrupted schema") + } + execOn(t, database.conn, "DROP TABLE metadata") + if err := Migrate(database); err != nil { + t.Fatalf("Migrate after interruption cleanup: %v", err) + } + assertSchemaCurrent(t, database.conn) +} + +func testMigrationRejectsNewer(t *testing.T) { + database := openTestDatabase(t) + defer database.Close() + statement := fmt.Sprintf("PRAGMA user_version = %d", currentSchemaVersion+1) + execOn(t, database.conn, statement) + if err := Migrate(database); err == nil { + t.Fatal("Migrate accepted a newer schema version") + } + if tableExists(t, database.conn, "metadata") { + t.Fatal("newer-schema rejection must not create any tables") + } +} + +func assertSchemaCurrent(t *testing.T, conn *sql.DB) { + t.Helper() + if version := userVersion(t, conn); version != currentSchemaVersion { + t.Fatalf("user_version = %d, want %d", version, currentSchemaVersion) + } + for _, table := range []string{"metadata", "repositories", "files", "aliases"} { + if !tableExists(t, conn, table) { + t.Fatalf("table %q missing after migration", table) + } + } +} + +func userVersion(t *testing.T, conn *sql.DB) int { + t.Helper() + var version int + if err := conn.QueryRow("PRAGMA user_version").Scan(&version); err != nil { + t.Fatalf("read user_version: %v", err) + } + return version +} + +func tableExists(t *testing.T, conn *sql.DB, name string) bool { + t.Helper() + var found string + query := "SELECT name FROM sqlite_master WHERE type = ? AND name = ?" + if err := conn.QueryRow(query, "table", name).Scan(&found); err != nil { + if err == sql.ErrNoRows { + return false + } + t.Fatalf("query table %s: %v", name, err) + } + return found == name +} + +func execOn(t *testing.T, conn *sql.DB, statement string) { + t.Helper() + if _, err := conn.Exec(statement); err != nil { + t.Fatalf("%q: %v", statement, err) + } +} -- 2.51.2