diff --git a/db/pg/pg.go b/db/pg/pg.go index 3b034736..015910d6 100644 --- a/db/pg/pg.go +++ b/db/pg/pg.go @@ -13,6 +13,12 @@ import ( "heckel.io/ntfy/v2/db" ) +// Advisory lock keys. PostgreSQL advisory locks share one database-wide key space, so every +// ntfy key is defined here, following the "ntfy"+2586+letter scheme +const ( + SchemaLockKey = int64(0x6e7466792586a) // Schema setup serialization (transaction-scoped, see db/schema) +) + // Open opens a PostgreSQL connection pool for a primary database. It pings the database // to verify connectivity before returning. func Open(dsn string) (*db.Host, error) { diff --git a/db/schema/schema.go b/db/schema/schema.go new file mode 100644 index 00000000..f4f97b7c --- /dev/null +++ b/db/schema/schema.go @@ -0,0 +1,127 @@ +// Package schema tracks and migrates database schemas, and Migrate creates or upgrades a +// store's schema inside a single transaction. On PostgreSQL, all stores share one database, so +// versions live in a shared schema_version table keyed by store name. On SQLite, every store is +// its own database file, so the version lives in the schemaVersion table keyed by id = 1 -- the +// exact layout pre-existing ntfy SQLite databases already use, which lets them migrate onto +// this framework without an adoption step. +package schema + +import ( + "database/sql" + "errors" + "fmt" + + "heckel.io/ntfy/v2/db/pg" +) + +// Queries by dialect; the version table is shared by all stores on Postgres, and per database +// file on SQLite +const ( + postgresCreateVersionTableQuery = `CREATE TABLE IF NOT EXISTS schema_version (store TEXT PRIMARY KEY, version INT NOT NULL)` + postgresSelectVersionQuery = `SELECT version FROM schema_version WHERE store = $1` + postgresUpsertVersionQuery = `INSERT INTO schema_version (store, version) VALUES ($1, $2) ON CONFLICT (store) DO UPDATE SET version = EXCLUDED.version` + postgresAdvisoryLockQuery = `SELECT pg_advisory_xact_lock($1)` + + sqliteCreateVersionTableQuery = `CREATE TABLE IF NOT EXISTS schemaVersion (id INT PRIMARY KEY, version INT NOT NULL)` + sqliteSelectVersionQuery = `SELECT version FROM schemaVersion WHERE id = 1` + sqliteUpsertVersionQuery = `INSERT INTO schemaVersion (id, version) VALUES (1, ?) ON CONFLICT (id) DO UPDATE SET version = excluded.version` +) + +// Dialect selects the SQL flavor Migrate speaks to the version table. +type Dialect int + +// Supported dialects; SQLite is the zero value +const ( + SQLite Dialect = iota + Postgres +) + +// MigrateFunc applies one schema change inside the setup transaction: the initial creation of +// a store's tables, or one step upgrading a store from version N to N+1. Migrations needing +// config capture it via closure, e.g. func migrations(cacheDuration time.Duration) map[int]MigrateFunc. +type MigrateFunc func(tx *sql.Tx) error + +// Migrate creates or upgrades the named store's schema to targetVersion in one transaction: a +// fresh database gets the create func; an existing one is upgraded step by step through the +// migrations map (keyed by the FROM version; always append, never insert in the middle); a +// schema migrated by newer code is refused. A failed migration rolls back atomically. Caveats: +// statements that refuse transaction blocks (Postgres CREATE INDEX CONCURRENTLY, VACUUM) cannot +// go through Migrate, and DDL holds exclusive locks until commit, so keep migrations fast. +func Migrate(db *sql.DB, dialect Dialect, store string, targetVersion int, create MigrateFunc, migrations map[int]MigrateFunc) error { + if dialect != Postgres && dialect != SQLite { + return fmt.Errorf("unsupported schema dialect %d", dialect) + } + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("cannot begin %s schema transaction: %w", store, err) + } + defer tx.Rollback() + if dialect == Postgres { + // Serialize setup across nodes: CREATE TABLE IF NOT EXISTS is not atomic, and + // concurrently cold-booting nodes would otherwise race on DDL and crash + if _, err := tx.Exec(postgresAdvisoryLockQuery, pg.SchemaLockKey); err != nil { + return fmt.Errorf("cannot acquire %s schema advisory lock: %w", store, err) + } + } + if _, err := tx.Exec(createVersionTableQuery(dialect)); err != nil { + return fmt.Errorf("cannot create schema version table: %w", err) + } + version, err := readVersion(tx, dialect, store) + if errors.Is(err, sql.ErrNoRows) { + // Fresh database: create the store's tables at the target version + if err := create(tx); err != nil { + return fmt.Errorf("cannot create %s schema: %w", store, err) + } + if err := writeVersion(tx, dialect, store, targetVersion); err != nil { + return fmt.Errorf("cannot write %s schema version: %w", store, err) + } + return tx.Commit() + } else if err != nil { + return fmt.Errorf("cannot read %s schema version: %w", store, err) + } + if version == targetVersion { + return tx.Commit() + } + if version > targetVersion { + return fmt.Errorf("unexpected %s schema version %d, this version of ntfy supports up to %d", store, version, targetVersion) + } + for v := version; v < targetVersion; v++ { + migrate, ok := migrations[v] + if !ok { + return fmt.Errorf("cannot find %s migration step from version %d to %d", store, v, v+1) + } + if err := migrate(tx); err != nil { + return fmt.Errorf("%s migration step from version %d to %d failed: %w", store, v, v+1, err) + } + } + if err := writeVersion(tx, dialect, store, targetVersion); err != nil { + return fmt.Errorf("cannot write %s schema version: %w", store, err) + } + return tx.Commit() +} + +func createVersionTableQuery(dialect Dialect) string { + if dialect == Postgres { + return postgresCreateVersionTableQuery + } + return sqliteCreateVersionTableQuery +} + +func readVersion(tx *sql.Tx, dialect Dialect, store string) (version int, err error) { + if dialect == Postgres { + err = tx.QueryRow(postgresSelectVersionQuery, store).Scan(&version) + } else { + err = tx.QueryRow(sqliteSelectVersionQuery).Scan(&version) + } + return +} + +func writeVersion(tx *sql.Tx, dialect Dialect, store string, version int) error { + var err error + if dialect == Postgres { + _, err = tx.Exec(postgresUpsertVersionQuery, store, version) + } else { + _, err = tx.Exec(sqliteUpsertVersionQuery, version) + } + return err +} diff --git a/db/schema/schema_test.go b/db/schema/schema_test.go new file mode 100644 index 00000000..9d848e94 --- /dev/null +++ b/db/schema/schema_test.go @@ -0,0 +1,193 @@ +package schema_test + +import ( + "database/sql" + "fmt" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/db/pg" + "heckel.io/ntfy/v2/db/schema" + dbtest "heckel.io/ntfy/v2/db/test" + + _ "github.com/mattn/go-sqlite3" +) + +const testCreateQuery = `CREATE TABLE IF NOT EXISTS things (id TEXT PRIMARY KEY, name TEXT NOT NULL)` + +func testCreate(tx *sql.Tx) error { + _, err := tx.Exec(testCreateQuery) + return err +} + +func openTestPostgres(t *testing.T) *sql.DB { + t.Helper() + host, err := pg.Open(dbtest.CreateTestPostgresSchema(t)) + require.Nil(t, err) + t.Cleanup(func() { host.DB.Close() }) + return host.DB +} + +func openTestSQLite(t *testing.T) *sql.DB { + t.Helper() + d, err := sql.Open("sqlite3", filepath.Join(t.TempDir(), "test.db")) + require.Nil(t, err) + t.Cleanup(func() { d.Close() }) + return d +} + +func forEachDialect(t *testing.T, f func(t *testing.T, d *sql.DB, dialect schema.Dialect)) { + t.Run("postgres", func(t *testing.T) { + f(t, openTestPostgres(t), schema.Postgres) + }) + t.Run("sqlite", func(t *testing.T) { + f(t, openTestSQLite(t), schema.SQLite) + }) +} + +func TestMigrate_FreshCreate(t *testing.T) { + forEachDialect(t, func(t *testing.T, d *sql.DB, dialect schema.Dialect) { + // A fresh database jumps straight to the target version; migration steps are not consulted + require.Nil(t, schema.Migrate(d, dialect, "things", 3, testCreate, nil)) + _, err := d.Exec(`INSERT INTO things (id, name) VALUES ('a', 'thing a')`) + require.Nil(t, err) + require.Equal(t, 3, storeVersion(t, d, dialect, "things")) + // Idempotent: a second node boots against the migrated schema + require.Nil(t, schema.Migrate(d, dialect, "things", 3, testCreate, nil)) + }) +} + +func TestMigrate_AppliesMigrationSteps(t *testing.T) { + forEachDialect(t, func(t *testing.T, d *sql.DB, dialect schema.Dialect) { + require.Nil(t, schema.Migrate(d, dialect, "things", 1, testCreate, nil)) + // A newer version of the code migrates 1 -> 3 step by step, in order + migrations := map[int]schema.MigrateFunc{ + 1: func(tx *sql.Tx) error { + _, err := tx.Exec(`ALTER TABLE things ADD COLUMN color TEXT NOT NULL DEFAULT ''`) + return err + }, + 2: func(tx *sql.Tx) error { + _, err := tx.Exec(`ALTER TABLE things ADD COLUMN size INT NOT NULL DEFAULT 0`) + return err + }, + } + require.Nil(t, schema.Migrate(d, dialect, "things", 3, testCreate, migrations)) + _, err := d.Exec(`INSERT INTO things (id, name, color, size) VALUES ('b', 'thing b', 'red', 2)`) + require.Nil(t, err) + require.Equal(t, 3, storeVersion(t, d, dialect, "things")) + }) +} + +func TestMigrate_ClosureCarriesConfig(t *testing.T) { + // Migrations needing config take it via closure at map-construction time; there is no + // params plumbing in the framework itself + migrationsFor := func(defaultName string) map[int]schema.MigrateFunc { + return map[int]schema.MigrateFunc{ + 1: func(tx *sql.Tx) error { + _, err := tx.Exec(fmt.Sprintf(`ALTER TABLE things ADD COLUMN nick TEXT NOT NULL DEFAULT '%s'`, defaultName)) + return err + }, + } + } + forEachDialect(t, func(t *testing.T, d *sql.DB, dialect schema.Dialect) { + require.Nil(t, schema.Migrate(d, dialect, "things", 1, testCreate, nil)) + _, err := d.Exec(`INSERT INTO things (id, name) VALUES ('a', 'thing a')`) + require.Nil(t, err) + require.Nil(t, schema.Migrate(d, dialect, "things", 2, testCreate, migrationsFor("configured-default"))) + var nick string + require.Nil(t, d.QueryRow(`SELECT nick FROM things WHERE id = 'a'`).Scan(&nick)) + require.Equal(t, "configured-default", nick) + }) +} + +func TestMigrate_InvalidDialect(t *testing.T) { + d := openTestSQLite(t) + err := schema.Migrate(d, schema.Dialect(99), "things", 1, testCreate, nil) + require.Error(t, err) +} + +func TestMigrate_RefusesFutureVersion(t *testing.T) { + forEachDialect(t, func(t *testing.T, d *sql.DB, dialect schema.Dialect) { + require.Nil(t, schema.Migrate(d, dialect, "things", 2, testCreate, map[int]schema.MigrateFunc{})) + err := schema.Migrate(d, dialect, "things", 1, testCreate, nil) + require.Error(t, err) + }) +} + +func TestMigrate_MissingStepFails(t *testing.T) { + forEachDialect(t, func(t *testing.T, d *sql.DB, dialect schema.Dialect) { + require.Nil(t, schema.Migrate(d, dialect, "things", 1, testCreate, nil)) + err := schema.Migrate(d, dialect, "things", 3, testCreate, nil) // No step 1 -> 2 registered + require.Error(t, err) + }) +} + +func TestMigrate_StoresAreIndependent(t *testing.T) { + // Postgres only: stores share one database, tracked as rows in schema_version. On SQLite + // every store has its own database file, so independence is by file. + d := openTestPostgres(t) + require.Nil(t, schema.Migrate(d, schema.Postgres, "things", 1, testCreate, nil)) + require.Nil(t, schema.Migrate(d, schema.Postgres, "gadgets", 4, func(tx *sql.Tx) error { + _, err := tx.Exec(`CREATE TABLE IF NOT EXISTS gadgets (id TEXT PRIMARY KEY)`) + return err + }, nil)) + require.Equal(t, 1, storeVersion(t, d, schema.Postgres, "things")) + require.Equal(t, 4, storeVersion(t, d, schema.Postgres, "gadgets")) +} + +func TestMigrate_SQLiteReadsExistingSchemaVersionTable(t *testing.T) { + // Existing ntfy SQLite databases (message, user, webpush) track their version in a + // schemaVersion (id, version) table keyed by id = 1; the framework uses that table as-is + // on SQLite, so existing databases migrate without any adoption step + d := openTestSQLite(t) + _, err := d.Exec(testCreateQuery) + require.Nil(t, err) + _, err = d.Exec(`CREATE TABLE schemaVersion (id INT PRIMARY KEY, version INT NOT NULL)`) + require.Nil(t, err) + _, err = d.Exec(`INSERT INTO schemaVersion VALUES (1, 1)`) + require.Nil(t, err) + migrations := map[int]schema.MigrateFunc{ + 1: func(tx *sql.Tx) error { + _, err := tx.Exec(`ALTER TABLE things ADD COLUMN color TEXT NOT NULL DEFAULT ''`) + return err + }, + } + require.Nil(t, schema.Migrate(d, schema.SQLite, "things", 2, testCreate, migrations)) + _, err = d.Exec(`INSERT INTO things (id, name, color) VALUES ('a', 'thing a', 'red')`) + require.Nil(t, err) + require.Equal(t, 2, storeVersion(t, d, schema.SQLite, "things")) +} + +func TestMigrate_ConcurrentFreshCreate(t *testing.T) { + // Postgres only: concurrent cold-boots must not race on DDL (CREATE TABLE IF NOT EXISTS is + // not atomic); Migrate serializes via an advisory lock. SQLite has a single writer. + schemaDSN := dbtest.CreateTestPostgresSchema(t) + const n = 8 + errs := make(chan error, n) + for i := 0; i < n; i++ { + go func() { + host, err := pg.Open(schemaDSN) + if err != nil { + errs <- err + return + } + defer host.DB.Close() + errs <- schema.Migrate(host.DB, schema.Postgres, "things", 1, testCreate, nil) + }() + } + for i := 0; i < n; i++ { + require.Nil(t, <-errs) + } +} + +func storeVersion(t *testing.T, d *sql.DB, dialect schema.Dialect, store string) int { + t.Helper() + var version int + if dialect == schema.Postgres { + require.Nil(t, d.QueryRow(`SELECT version FROM schema_version WHERE store = $1`, store).Scan(&version), fmt.Sprintf("store %s", store)) + } else { + require.Nil(t, d.QueryRow(`SELECT version FROM schemaVersion WHERE id = 1`).Scan(&version), fmt.Sprintf("store %s", store)) + } + return version +} diff --git a/webpush/store_sqlite.go b/webpush/store_sqlite.go index 7677f1ce..e03a43b1 100644 --- a/webpush/store_sqlite.go +++ b/webpush/store_sqlite.go @@ -5,7 +5,6 @@ import ( "fmt" _ "github.com/mattn/go-sqlite3" // SQLite driver - "heckel.io/ntfy/v2/db" )