mirror of
https://github.com/multipleof4/ntfy.git
synced 2026-10-08 21:05:21 +00:00
Merge pull request #1868 from binwiederhier/schema-migration
Schema migration library
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -14,6 +14,8 @@ const (
|
||||
subscriptionIDPrefix = "wps_"
|
||||
subscriptionIDLength = 10
|
||||
subscriptionEndpointLimitPerSubscriberIP = 10
|
||||
|
||||
schemaStore = "webpush"
|
||||
)
|
||||
|
||||
// Errors returned by the store
|
||||
|
||||
+28
-58
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
+28
-54
@@ -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
|
||||
|
||||
@@ -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"), "")
|
||||
|
||||
Reference in New Issue
Block a user