feat(core): add WithMigrationBaseline hook before goose Up

This commit is contained in:
ryan
2026-08-30 11:41:00 +08:00
parent 4ce7110e23
commit a617457a3c
5 changed files with 313 additions and 52 deletions
+44 -4
View File
@@ -41,6 +41,10 @@ import (
const (
defaultShutdownTimeout = 15 * time.Second
defaultHTTPAddr = "127.0.0.1:8000"
// migrationAdvisoryLockKey serializes baseline + plugin Up across Postgres
// sessions (ASCII "wave"). SQLite is single-writer and needs no extra lock.
migrationAdvisoryLockKey int64 = 0x77617665
)
// runProfileApp prepares and runs the application for a given profile.
@@ -140,17 +144,21 @@ type sharedStore struct {
func (s *sharedStore) Tablename() string { return "w_schema_versions" }
func (s *sharedStore) CreateVersionTable(ctx context.Context, db goosedb.DBTxConn) error {
_, err := db.ExecContext(ctx, schemaVersionsDDL(s.dialect))
return err
}
func schemaVersionsDDL(dialect string) string {
timeType := "TIMESTAMPTZ"
if s.dialect == "sqlite3" || s.dialect == "sqlite" {
if dialect == "sqlite3" || dialect == "sqlite" {
timeType = "DATETIME"
}
_, err := db.ExecContext(ctx, fmt.Sprintf(`CREATE TABLE IF NOT EXISTS w_schema_versions (
return fmt.Sprintf(`CREATE TABLE IF NOT EXISTS w_schema_versions (
plugin_id VARCHAR(64) NOT NULL,
version_id BIGINT NOT NULL,
applied_at %s NOT NULL DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY (plugin_id, version_id)
)`, timeType))
return err
)`, timeType)
}
//nolint:mnd
@@ -265,6 +273,38 @@ func (e *gooseEngine) Migrate(ctx *core.Context, entries []core.MigrationEntry)
dialect := gooseDialect(ctx)
dialectStr := string(dialect)
goCtx := context.Background()
if ctx != nil {
goCtx = ctx.GoContext()
}
if goCtx == nil {
goCtx = context.Background()
}
bootstrap := &sharedStore{dialect: dialectStr}
if err := bootstrap.CreateVersionTable(goCtx, sqlDB); err != nil {
return fmt.Errorf("migration: create version table: %w", err)
}
if dialect == goose.DialectPostgres {
conn, lockErr := sqlDB.Conn(goCtx)
if lockErr != nil {
return fmt.Errorf("migration: pin connection for advisory lock: %w", lockErr)
}
defer func() { _ = conn.Close() }()
if _, lockErr = conn.ExecContext(goCtx, "SELECT pg_advisory_lock($1)", migrationAdvisoryLockKey); lockErr != nil {
return fmt.Errorf("migration: advisory lock: %w", lockErr)
}
defer func() {
_, _ = conn.ExecContext(context.Background(), "SELECT pg_advisory_unlock($1)", migrationAdvisoryLockKey)
}()
}
if fn := ctx.MigrationBaseline(); fn != nil {
if err := fn(ctx); err != nil {
return fmt.Errorf("migration baseline: %w", err)
}
}
for _, entry := range entries {
store := &sharedStore{
+127
View File
@@ -0,0 +1,127 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cmd
import (
"Wavelet/core"
"Wavelet/core/contracts"
"context"
"path/filepath"
"testing"
"testing/fstest"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
type migrateTestDB struct {
db *gorm.DB
}
func (s migrateTestDB) GORM() *gorm.DB { return s.db }
func (s migrateTestDB) DB(ctx context.Context) *gorm.DB { return s.db.WithContext(ctx) }
func (s migrateTestDB) Named(string) *gorm.DB { return s.db }
type migrateTestPlugin struct {
db *gorm.DB
fs fstest.MapFS
}
func (p *migrateTestPlugin) Name() string { return "t" }
func (p *migrateTestPlugin) Apply(ctx *core.Context) error {
core.Provide[contracts.DBService](ctx, migrateTestDB{db: p.db})
ctx.Migrations().Register("t", p.fs)
return nil
}
func sqliteTableExists(t *testing.T, db *gorm.DB, name string) bool {
t.Helper()
var n int
err := db.Raw("SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?", name).Scan(&n).Error
require.NoError(t, err)
return n > 0
}
func testMigrationFS() fstest.MapFS {
return fstest.MapFS{
"migrations/sqlite/00001_init.sql": &fstest.MapFile{Data: []byte(`-- +goose Up
CREATE TABLE t_up (id INTEGER PRIMARY KEY);
-- +goose Down
DROP TABLE t_up;
`)},
}
}
func openMigrateTestDB(t *testing.T) *gorm.DB {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "migrate.db")
gdb, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
require.NoError(t, err)
return gdb
}
func TestGooseEngineMigrateOrderCreateTableBaselineUp(t *testing.T) {
gdb := openMigrateTestDB(t)
var order []string
app := core.NewApp(
core.WithMigrationEngine(&gooseEngine{}),
core.WithMigrationBaseline(func(*core.Context) error {
require.True(t, sqliteTableExists(t, gdb, "w_schema_versions"), "version table must exist before baseline")
require.False(t, sqliteTableExists(t, gdb, "t_up"), "plugin Up must not run before baseline")
order = append(order, "create-table", "baseline")
return nil
}),
core.WithPlugins(&migrateTestPlugin{db: gdb, fs: testMigrationFS()}),
)
require.NoError(t, app.Prepare())
require.NoError(t, app.ApplyPlugins())
require.NoError(t, app.RunMigrations())
require.True(t, sqliteTableExists(t, gdb, "t_up"), "plugin Up must run after baseline")
order = append(order, "up")
assert.Equal(t, []string{"create-table", "baseline", "up"}, order)
}
func TestGooseEngineBaselineErrorSkipsUp(t *testing.T) {
gdb := openMigrateTestDB(t)
app := core.NewApp(
core.WithMigrationEngine(&gooseEngine{}),
core.WithMigrationBaseline(func(*core.Context) error {
require.True(t, sqliteTableExists(t, gdb, "w_schema_versions"), "version table must exist before baseline")
return assert.AnError
}),
core.WithPlugins(&migrateTestPlugin{db: gdb, fs: testMigrationFS()}),
)
require.NoError(t, app.Prepare())
require.NoError(t, app.ApplyPlugins())
err := app.RunMigrations()
require.Error(t, err)
assert.ErrorContains(t, err, "migration baseline")
assert.False(t, sqliteTableExists(t, gdb, "t_up"), "plugin Up must not run when baseline fails")
}
func TestGooseEngineNilBaselineStillMigrates(t *testing.T) {
gdb := openMigrateTestDB(t)
app := core.NewApp(
core.WithMigrationEngine(&gooseEngine{}),
core.WithPlugins(&migrateTestPlugin{db: gdb, fs: testMigrationFS()}),
)
require.NoError(t, app.Prepare())
require.NoError(t, app.ApplyPlugins())
require.NoError(t, app.RunMigrations())
assert.True(t, sqliteTableExists(t, gdb, "w_schema_versions"))
assert.True(t, sqliteTableExists(t, gdb, "t_up"))
}