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..6fa0b226 --- /dev/null +++ b/db/schema/schema.go @@ -0,0 +1,105 @@ +// 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. +package schema + +import ( + "database/sql" + "errors" + "fmt" + + "heckel.io/ntfy/v2/db/pg" +) + +const ( + 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` + + 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)` // Transaction-scoped lock to avoid migration races +) + +// Migrate creates or upgrades the named store's schema to targetVersion in one transaction, or +// creates a new database using the "create" function. +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..28e6c376 --- /dev/null +++ b/db/schema/schema_test.go @@ -0,0 +1,190 @@ +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: schema.AsMigrateFunc(fmt.Sprintf(`ALTER TABLE things ADD COLUMN nick TEXT NOT NULL DEFAULT '%s'`, defaultName)), + } + } + 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/db/schema/types.go b/db/schema/types.go new file mode 100644 index 00000000..430a02cd --- /dev/null +++ b/db/schema/types.go @@ -0,0 +1,25 @@ +package schema + +import "database/sql" + +// 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 + +// AsMigrateFunc converts a simple query to a migration function +func AsMigrateFunc(query string) MigrateFunc { + return func(tx *sql.Tx) error { + _, err := tx.Exec(query) + return err + } +} diff --git a/webpush/store.go b/webpush/store.go index 1a9825f5..a05fee6c 100644 --- a/webpush/store.go +++ b/webpush/store.go @@ -14,6 +14,8 @@ const ( subscriptionIDPrefix = "wps_" subscriptionIDLength = 10 subscriptionEndpointLimitPerSubscriberIP = 10 + + schemaStore = "webpush" ) // Errors returned by the store diff --git a/webpush/store_postgres.go b/webpush/store_postgres.go index 84168d89..dbb67faa 100644 --- a/webpush/store_postgres.go +++ b/webpush/store_postgres.go @@ -1,39 +1,11 @@ package webpush import ( - "database/sql" - "fmt" - "heckel.io/ntfy/v2/db" + "heckel.io/ntfy/v2/db/schema" ) const ( - postgresCreateTablesQuery = ` - CREATE TABLE IF NOT EXISTS webpush_subscription ( - id TEXT PRIMARY KEY, - endpoint TEXT NOT NULL UNIQUE, - key_auth TEXT NOT NULL, - key_p256dh TEXT NOT NULL, - user_id TEXT NOT NULL, - subscriber_ip TEXT NOT NULL, - updated_at BIGINT NOT NULL, - warned_at BIGINT NOT NULL DEFAULT 0 - ); - CREATE INDEX IF NOT EXISTS idx_webpush_subscriber_ip ON webpush_subscription (subscriber_ip); - CREATE INDEX IF NOT EXISTS idx_webpush_updated_at ON webpush_subscription (updated_at); - CREATE INDEX IF NOT EXISTS idx_webpush_user_id ON webpush_subscription (user_id); - CREATE TABLE IF NOT EXISTS webpush_subscription_topic ( - subscription_id TEXT NOT NULL REFERENCES webpush_subscription (id) ON DELETE CASCADE, - topic TEXT NOT NULL, - PRIMARY KEY (subscription_id, topic) - ); - CREATE INDEX IF NOT EXISTS idx_webpush_topic ON webpush_subscription_topic (topic); - CREATE TABLE IF NOT EXISTS schema_version ( - store TEXT PRIMARY KEY, - version INT NOT NULL - ); - ` - postgresSelectSubscriptionIDByEndpointQuery = `SELECT id FROM webpush_subscription WHERE endpoint = $1` postgresSelectSubscriptionCountBySubscriberIPQuery = `SELECT COUNT(*) FROM webpush_subscription WHERE subscriber_ip = $1` postgresSelectSubscriptionsForTopicQuery = ` @@ -66,16 +38,38 @@ const ( postgresDeleteSubscriptionTopicWithoutSubscriptionQuery = `DELETE FROM webpush_subscription_topic WHERE subscription_id NOT IN (SELECT id FROM webpush_subscription)` ) -// PostgreSQL schema management queries +// Schema version and queries const ( - pgCurrentSchemaVersion = 1 - postgresInsertSchemaVersionQuery = `INSERT INTO schema_version (store, version) VALUES ('webpush', $1)` - postgresSelectSchemaVersionQuery = `SELECT version FROM schema_version WHERE store = 'webpush'` + postgresCurrentSchemaVersion = 1 +) + +var ( + postgresCreateTables = schema.AsMigrateFunc(` + CREATE TABLE IF NOT EXISTS webpush_subscription ( + id TEXT PRIMARY KEY, + endpoint TEXT NOT NULL UNIQUE, + key_auth TEXT NOT NULL, + key_p256dh TEXT NOT NULL, + user_id TEXT NOT NULL, + subscriber_ip TEXT NOT NULL, + updated_at BIGINT NOT NULL, + warned_at BIGINT NOT NULL DEFAULT 0 + ); + CREATE INDEX IF NOT EXISTS idx_webpush_subscriber_ip ON webpush_subscription (subscriber_ip); + CREATE INDEX IF NOT EXISTS idx_webpush_updated_at ON webpush_subscription (updated_at); + CREATE INDEX IF NOT EXISTS idx_webpush_user_id ON webpush_subscription (user_id); + CREATE TABLE IF NOT EXISTS webpush_subscription_topic ( + subscription_id TEXT NOT NULL REFERENCES webpush_subscription (id) ON DELETE CASCADE, + topic TEXT NOT NULL, + PRIMARY KEY (subscription_id, topic) + ); + CREATE INDEX IF NOT EXISTS idx_webpush_topic ON webpush_subscription_topic (topic); + `) ) // NewPostgresStore creates a new PostgreSQL-backed web push store using an existing database connection pool. func NewPostgresStore(d *db.DB) (*Store, error) { - if err := setupPostgres(d.Primary()); err != nil { + if err := schema.Migrate(d.Primary(), schema.Postgres, schemaStore, postgresCurrentSchemaVersion, postgresCreateTables, nil); err != nil { return nil, err } return &Store{ @@ -97,27 +91,3 @@ func NewPostgresStore(d *db.DB) (*Store, error) { }, }, nil } - -func setupPostgres(d *sql.DB) error { - var schemaVersion int - err := d.QueryRow(postgresSelectSchemaVersionQuery).Scan(&schemaVersion) - if err != nil { - return setupNewPostgres(d) - } - if schemaVersion > pgCurrentSchemaVersion { - return fmt.Errorf("unexpected schema version: version %d is higher than current version %d", schemaVersion, pgCurrentSchemaVersion) - } - return nil -} - -func setupNewPostgres(d *sql.DB) error { - return db.ExecTx(d, func(tx *sql.Tx) error { - if _, err := tx.Exec(postgresCreateTablesQuery); err != nil { - return err - } - if _, err := tx.Exec(postgresInsertSchemaVersionQuery, pgCurrentSchemaVersion); err != nil { - return err - } - return nil - }) -} diff --git a/webpush/store_sqlite.go b/webpush/store_sqlite.go index 7677f1ce..6a5f9683 100644 --- a/webpush/store_sqlite.go +++ b/webpush/store_sqlite.go @@ -2,39 +2,13 @@ package webpush import ( "database/sql" - "fmt" _ "github.com/mattn/go-sqlite3" // SQLite driver - "heckel.io/ntfy/v2/db" + "heckel.io/ntfy/v2/db/schema" ) const ( - sqliteCreateTablesQuery = ` - CREATE TABLE IF NOT EXISTS subscription ( - id TEXT PRIMARY KEY, - endpoint TEXT NOT NULL, - key_auth TEXT NOT NULL, - key_p256dh TEXT NOT NULL, - user_id TEXT NOT NULL, - subscriber_ip TEXT NOT NULL, - updated_at INT NOT NULL, - warned_at INT NOT NULL DEFAULT 0 - ); - CREATE UNIQUE INDEX IF NOT EXISTS idx_endpoint ON subscription (endpoint); - CREATE INDEX IF NOT EXISTS idx_subscriber_ip ON subscription (subscriber_ip); - CREATE TABLE IF NOT EXISTS subscription_topic ( - subscription_id TEXT NOT NULL, - topic TEXT NOT NULL, - PRIMARY KEY (subscription_id, topic), - FOREIGN KEY (subscription_id) REFERENCES subscription (id) ON DELETE CASCADE - ); - CREATE INDEX IF NOT EXISTS idx_topic ON subscription_topic (topic); - CREATE TABLE IF NOT EXISTS schemaVersion ( - id INT PRIMARY KEY, - version INT NOT NULL - ); - ` sqliteBuiltinStartupQueries = ` PRAGMA foreign_keys = ON; ` @@ -71,11 +45,33 @@ const ( sqliteDeleteSubscriptionTopicWithoutSubscriptionQuery = `DELETE FROM subscription_topic WHERE subscription_id NOT IN (SELECT id FROM subscription)` ) -// SQLite schema management queries +// Schema version and queries const ( - sqliteCurrentSchemaVersion = 1 - sqliteInsertSchemaVersionQuery = `INSERT INTO schemaVersion VALUES (1, ?)` - sqliteSelectSchemaVersionQuery = `SELECT version FROM schemaVersion WHERE id = 1` + sqliteCurrentSchemaVersion = 1 +) + +var ( + sqliteCreateTables = schema.AsMigrateFunc(` + CREATE TABLE IF NOT EXISTS subscription ( + id TEXT PRIMARY KEY, + endpoint TEXT NOT NULL, + key_auth TEXT NOT NULL, + key_p256dh TEXT NOT NULL, + user_id TEXT NOT NULL, + subscriber_ip TEXT NOT NULL, + updated_at INT NOT NULL, + warned_at INT NOT NULL DEFAULT 0 + ); + CREATE UNIQUE INDEX IF NOT EXISTS idx_endpoint ON subscription (endpoint); + CREATE INDEX IF NOT EXISTS idx_subscriber_ip ON subscription (subscriber_ip); + CREATE TABLE IF NOT EXISTS subscription_topic ( + subscription_id TEXT NOT NULL, + topic TEXT NOT NULL, + PRIMARY KEY (subscription_id, topic), + FOREIGN KEY (subscription_id) REFERENCES subscription (id) ON DELETE CASCADE + ); + CREATE INDEX IF NOT EXISTS idx_topic ON subscription_topic (topic); + `) ) // NewSQLiteStore creates a new SQLite-backed web push store. @@ -84,7 +80,7 @@ func NewSQLiteStore(filename, startupQueries string) (*Store, error) { if err != nil { return nil, err } - if err := setupSQLite(d); err != nil { + if err := schema.Migrate(d, schema.SQLite, schemaStore, sqliteCurrentSchemaVersion, sqliteCreateTables, nil); err != nil { return nil, err } if err := runSQLiteStartupQueries(d, startupQueries); err != nil { @@ -110,28 +106,6 @@ func NewSQLiteStore(filename, startupQueries string) (*Store, error) { }, nil } -func setupSQLite(db *sql.DB) error { - var schemaVersion int - if err := db.QueryRow(sqliteSelectSchemaVersionQuery).Scan(&schemaVersion); err != nil { - return setupNewSQLite(db) - } else if schemaVersion > sqliteCurrentSchemaVersion { - return fmt.Errorf("unexpected schema version: version %d is higher than current version %d", schemaVersion, sqliteCurrentSchemaVersion) - } - return nil -} - -func setupNewSQLite(sqlDB *sql.DB) error { - return db.ExecTx(sqlDB, func(tx *sql.Tx) error { - if _, err := tx.Exec(sqliteCreateTablesQuery); err != nil { - return err - } - if _, err := tx.Exec(sqliteInsertSchemaVersionQuery, sqliteCurrentSchemaVersion); err != nil { - return err - } - return nil - }) -} - func runSQLiteStartupQueries(db *sql.DB, startupQueries string) error { if _, err := db.Exec(startupQueries); err != nil { return err diff --git a/webpush/store_test.go b/webpush/store_test.go index 348f9998..2323de56 100644 --- a/webpush/store_test.go +++ b/webpush/store_test.go @@ -1,6 +1,7 @@ package webpush_test import ( + "database/sql" "fmt" "net/netip" "path/filepath" @@ -14,6 +15,82 @@ import ( const testWebPushEndpoint = "https://updates.push.services.mozilla.com/wpush/v1/AAABBCCCDDEEEFFF" +// Schema layout as written by ntfy releases before the db/schema framework; used to verify +// that existing databases open cleanly without an adoption step +const ( + testPreFrameworkSQLiteSchema = ` + CREATE TABLE subscription ( + id TEXT PRIMARY KEY, + endpoint TEXT NOT NULL, + key_auth TEXT NOT NULL, + key_p256dh TEXT NOT NULL, + user_id TEXT NOT NULL, + subscriber_ip TEXT NOT NULL, + updated_at INT NOT NULL, + warned_at INT NOT NULL DEFAULT 0 + ); + CREATE UNIQUE INDEX idx_endpoint ON subscription (endpoint); + CREATE TABLE subscription_topic ( + subscription_id TEXT NOT NULL, + topic TEXT NOT NULL, + PRIMARY KEY (subscription_id, topic) + ); + CREATE TABLE schemaVersion (id INT PRIMARY KEY, version INT NOT NULL); + INSERT INTO schemaVersion VALUES (1, 1); + ` + testPreFrameworkPostgresSchema = ` + CREATE TABLE webpush_subscription ( + id TEXT PRIMARY KEY, + endpoint TEXT NOT NULL UNIQUE, + key_auth TEXT NOT NULL, + key_p256dh TEXT NOT NULL, + user_id TEXT NOT NULL, + subscriber_ip TEXT NOT NULL, + updated_at BIGINT NOT NULL, + warned_at BIGINT NOT NULL DEFAULT 0 + ); + CREATE TABLE webpush_subscription_topic ( + subscription_id TEXT NOT NULL REFERENCES webpush_subscription (id) ON DELETE CASCADE, + topic TEXT NOT NULL, + PRIMARY KEY (subscription_id, topic) + ); + CREATE TABLE schema_version (store TEXT PRIMARY KEY, version INT NOT NULL); + INSERT INTO schema_version (store, version) VALUES ('webpush', 1); + ` +) + +func TestStoreSQLiteOpensExistingDatabase(t *testing.T) { + filename := filepath.Join(t.TempDir(), "webpush.db") + d, err := sql.Open("sqlite3", filename) + require.Nil(t, err) + _, err = d.Exec(testPreFrameworkSQLiteSchema) + require.Nil(t, err) + require.Nil(t, d.Close()) + store, err := webpush.NewSQLiteStore(filename, "") + require.Nil(t, err) + defer store.Close() + requireStoreUsable(t, store) +} + +func TestStorePostgresOpensExistingDatabase(t *testing.T) { + testDB := dbtest.CreateTestPostgres(t) + _, err := testDB.Exec(testPreFrameworkPostgresSchema) + require.Nil(t, err) + store, err := webpush.NewPostgresStore(testDB) + require.Nil(t, err) + requireStoreUsable(t, store) +} + +func requireStoreUsable(t *testing.T, store *webpush.Store) { + t.Helper() + err := store.UpsertSubscription(testWebPushEndpoint, "auth-key", "p256dh-key", "u_1234", netip.MustParseAddr("1.2.3.4"), []string{"mytopic"}) + require.Nil(t, err) + subs, err := store.SubscriptionsForTopic("mytopic") + require.Nil(t, err) + require.Len(t, subs, 1) + require.Equal(t, testWebPushEndpoint, subs[0].Endpoint) +} + func forEachBackend(t *testing.T, f func(t *testing.T, store *webpush.Store)) { t.Run("sqlite", func(t *testing.T) { store, err := webpush.NewSQLiteStore(filepath.Join(t.TempDir(), "webpush.db"), "")