fix(pgsql): tie PK sequence to nextval default, setval past data, keep serial defaults
- derive the primary key sequence from the column's nextval() default instead of an unused identity_<table>_<pk> sequence - setval after table creation (full and diff paths), forward-only, MAX+1 with is_called=false - do not DROP DEFAULT on serial/bigserial columns without a model default
This commit is contained in:
@@ -0,0 +1,382 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
// diffTestDB builds a database with one table per spec: name -> column name -> type.
|
||||
func diffTestDB(schemaName string, tables map[string]map[string]string) *models.Database {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema(schemaName)
|
||||
for tName, cols := range tables {
|
||||
table := models.InitTable(tName, schemaName)
|
||||
seq := 1
|
||||
for cName, cType := range cols {
|
||||
col := models.InitColumn(cName, tName, schemaName)
|
||||
col.Type = cType
|
||||
col.Sequence = uint(seq)
|
||||
seq++
|
||||
table.Columns[cName] = col
|
||||
}
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
}
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
return db
|
||||
}
|
||||
|
||||
func diffJoin(stmts []string) string { return strings.Join(stmts, "\n") }
|
||||
|
||||
func TestDiffStatements(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
model func() *models.Database
|
||||
current func() *models.Database
|
||||
wantContain []string
|
||||
wantAbsent []string
|
||||
wantEmpty bool
|
||||
}{
|
||||
{
|
||||
name: "identical schemas produce no statements",
|
||||
model: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "integer", "email": "text"}})
|
||||
},
|
||||
current: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "integer", "email": "text"}})
|
||||
},
|
||||
wantEmpty: true,
|
||||
},
|
||||
{
|
||||
name: "missing table is created",
|
||||
model: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "integer"}, "posts": {"id": "integer"}})
|
||||
},
|
||||
current: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "integer"}})
|
||||
},
|
||||
wantContain: []string{"CREATE TABLE IF NOT EXISTS public.posts"},
|
||||
wantAbsent: []string{"public.users"},
|
||||
},
|
||||
{
|
||||
name: "missing column is added and existing columns are skipped",
|
||||
model: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "integer", "email": "text"}})
|
||||
},
|
||||
current: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "integer"}})
|
||||
},
|
||||
wantContain: []string{"ADD COLUMN IF NOT EXISTS email text"},
|
||||
wantAbsent: []string{"COLUMN id"},
|
||||
},
|
||||
{
|
||||
name: "changed column type is altered",
|
||||
model: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "bigint"}})
|
||||
},
|
||||
current: func() *models.Database {
|
||||
return diffTestDB("public", map[string]map[string]string{"users": {"id": "integer"}})
|
||||
},
|
||||
wantContain: []string{"ALTER COLUMN id TYPE bigint"},
|
||||
},
|
||||
{
|
||||
name: "missing non-public schema is created before its tables",
|
||||
model: func() *models.Database {
|
||||
return diffTestDB("app", map[string]map[string]string{"users": {"id": "integer"}})
|
||||
},
|
||||
current: func() *models.Database { return models.InitDatabase("testdb") },
|
||||
wantContain: []string{
|
||||
"CREATE SCHEMA IF NOT EXISTS app",
|
||||
"CREATE TABLE IF NOT EXISTS app.users",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(tt.model(), tt.current())
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
out := diffJoin(stmts)
|
||||
|
||||
if tt.wantEmpty && len(stmts) != 0 {
|
||||
t.Fatalf("expected no statements, got:\n%s", out)
|
||||
}
|
||||
for _, want := range tt.wantContain {
|
||||
if !strings.Contains(out, want) {
|
||||
t.Errorf("missing %q in:\n%s", want, out)
|
||||
}
|
||||
}
|
||||
for _, absent := range tt.wantAbsent {
|
||||
if strings.Contains(out, absent) {
|
||||
t.Errorf("unexpected %q in:\n%s", absent, out)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_SchemaCreatedBeforeTables(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
model := diffTestDB("app", map[string]map[string]string{"users": {"id": "integer"}})
|
||||
stmts, err := w.diffStatements(model, models.InitDatabase("testdb"))
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
|
||||
schemaIdx, tableIdx := -1, -1
|
||||
for i, s := range stmts {
|
||||
if strings.Contains(s, "CREATE SCHEMA") {
|
||||
schemaIdx = i
|
||||
}
|
||||
if strings.Contains(s, "CREATE TABLE") {
|
||||
tableIdx = i
|
||||
}
|
||||
}
|
||||
if schemaIdx < 0 || tableIdx < 0 || schemaIdx > tableIdx {
|
||||
t.Fatalf("expected CREATE SCHEMA before CREATE TABLE, got:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_NewTableCreatesPrimaryKeySequence(t *testing.T) {
|
||||
model := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
table := models.InitTable("users", "public")
|
||||
id := models.InitColumn("id", "users", "public")
|
||||
id.Type = "integer"
|
||||
id.IsPrimaryKey = true
|
||||
id.Default = "nextval('public.users_id_seq'::regclass)"
|
||||
table.Columns["id"] = id
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
model.Schemas = append(model.Schemas, schema)
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(model, models.InitDatabase("testdb"))
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
out := diffJoin(stmts)
|
||||
if !strings.Contains(out, "CREATE SEQUENCE IF NOT EXISTS public.users_id_seq") {
|
||||
t.Fatalf("expected sequence creation, got:\n%s", out)
|
||||
}
|
||||
|
||||
// Existing table: sequence must not be re-emitted.
|
||||
current := models.InitDatabase("testdb")
|
||||
cs := models.InitSchema("public")
|
||||
cs.Tables = append(cs.Tables, table)
|
||||
current.Schemas = append(current.Schemas, cs)
|
||||
stmts, err = w.diffStatements(model, current)
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if strings.Contains(diffJoin(stmts), "CREATE SEQUENCE") {
|
||||
t.Fatalf("sequence emitted for existing table:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_NewIndexAndChangedIndexRecreated(t *testing.T) {
|
||||
build := func(cols ...string) *models.Database {
|
||||
db := diffTestDB("public", map[string]map[string]string{"users": {"id": "integer", "email": "text", "name": "text"}})
|
||||
table := db.Schemas[0].Tables[0]
|
||||
table.Indexes["idx_users_lookup"] = &models.Index{Name: "idx_users_lookup", Columns: cols}
|
||||
return db
|
||||
}
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
|
||||
// Index missing in current.
|
||||
noIdx := diffTestDB("public", map[string]map[string]string{"users": {"id": "integer", "email": "text", "name": "text"}})
|
||||
stmts, err := w.diffStatements(build("email"), noIdx)
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(diffJoin(stmts), "CREATE INDEX IF NOT EXISTS idx_users_lookup") {
|
||||
t.Fatalf("expected index creation, got:\n%s", diffJoin(stmts))
|
||||
}
|
||||
|
||||
// Index unchanged.
|
||||
stmts, err = w.diffStatements(build("email"), build("email"))
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if len(stmts) != 0 {
|
||||
t.Fatalf("expected no statements for identical index, got:\n%s", diffJoin(stmts))
|
||||
}
|
||||
|
||||
// Index definition changed: dropped and recreated.
|
||||
stmts, err = w.diffStatements(build("name"), build("email"))
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
out := diffJoin(stmts)
|
||||
dropIdx := strings.Index(out, "DROP INDEX")
|
||||
createIdx := strings.Index(out, "CREATE INDEX")
|
||||
if dropIdx < 0 || createIdx < 0 || dropIdx > createIdx {
|
||||
t.Fatalf("expected DROP INDEX before CREATE INDEX, got:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_FlattenSchemaRejected(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{FlattenSchema: true})
|
||||
if _, err := w.generateLiveDiffStatements(models.InitDatabase("x"), "postgres://unused"); err == nil {
|
||||
t.Fatal("expected error so the caller falls back to full DDL")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_LiveDatabaseRepresentationsAreNotDifferences(t *testing.T) {
|
||||
build := func(live bool) *models.Database {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("public")
|
||||
|
||||
parent := models.InitTable("parent", "public")
|
||||
parentID := models.InitColumn("id_parent", "parent", "public")
|
||||
parentCurrentType := "bigserial"
|
||||
if live {
|
||||
parentCurrentType = "bigint"
|
||||
}
|
||||
parentID.Type = parentCurrentType
|
||||
parentID.IsPrimaryKey = true
|
||||
parentID.NotNull = true
|
||||
parent.Columns["id_parent"] = parentID
|
||||
if live {
|
||||
// A live database reports the PK as a named constraint and the unique
|
||||
// constraint's backing index as an index too.
|
||||
parent.Constraints["pk_public_parent"] = &models.Constraint{
|
||||
Name: "pk_public_parent", Type: models.PrimaryKeyConstraint, Columns: []string{"id_parent"},
|
||||
}
|
||||
}
|
||||
guid := models.InitColumn("guid", "parent", "public")
|
||||
guid.Type = "uuid"
|
||||
parent.Columns["guid"] = guid
|
||||
parent.Constraints["ukey_parent_guid"] = &models.Constraint{
|
||||
Name: "ukey_parent_guid", Type: models.UniqueConstraint, Columns: []string{"guid"},
|
||||
}
|
||||
if live {
|
||||
parent.Indexes["ukey_parent_guid"] = &models.Index{Name: "ukey_parent_guid", Unique: true, Columns: []string{"guid"}, Type: "btree"}
|
||||
}
|
||||
|
||||
// numeric(10,0) vs numeric(10); quoted default vs unquoted default with cast.
|
||||
amount := models.InitColumn("amount", "parent", "public")
|
||||
amount.Type = "numeric(10,0)"
|
||||
tags := models.InitColumn("tags", "parent", "public")
|
||||
tags.Type = "jsonb"
|
||||
tags.Default = "`'[]'`"
|
||||
if live {
|
||||
amount.Type = "numeric"
|
||||
amount.Precision = 10
|
||||
tags.Default = "[]"
|
||||
}
|
||||
parent.Columns["amount"] = amount
|
||||
parent.Columns["tags"] = tags
|
||||
|
||||
parent.Indexes["idx_parent_amount"] = &models.Index{Name: "idx_parent_amount", Columns: []string{"amount"}}
|
||||
if live {
|
||||
parent.Indexes["idx_parent_amount"].Type = "btree"
|
||||
}
|
||||
|
||||
child := models.InitTable("child", "public")
|
||||
cid := models.InitColumn("id_child", "child", "public")
|
||||
cid.Type = "integer"
|
||||
child.Columns["id_child"] = cid
|
||||
rid := models.InitColumn("rid_parent", "child", "public")
|
||||
rid.Type = "bigint"
|
||||
child.Columns["rid_parent"] = rid
|
||||
action := "restrict"
|
||||
if live {
|
||||
action = "RESTRICT"
|
||||
}
|
||||
child.Constraints["fk_child_rid_parent"] = &models.Constraint{
|
||||
Name: "fk_child_rid_parent", Type: models.ForeignKeyConstraint, Columns: []string{"rid_parent"},
|
||||
ReferencedTable: "parent", ReferencedSchema: "public", ReferencedColumns: []string{"id_parent"},
|
||||
OnDelete: action, OnUpdate: action,
|
||||
}
|
||||
|
||||
parent.Description = "Parent\ntable"
|
||||
if live {
|
||||
parent.Description = "Parent\ntable\n"
|
||||
}
|
||||
|
||||
schema.Tables = append(schema.Tables, parent, child)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
return db
|
||||
}
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(build(false), build(true))
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if len(stmts) != 0 {
|
||||
t.Fatalf("expected live representations to match the model, got:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_ChangedPrimaryKeyIsRecreated(t *testing.T) {
|
||||
model := diffTestDB("public", map[string]map[string]string{"users": {"id": "integer", "tenant": "integer"}})
|
||||
model.Schemas[0].Tables[0].Columns["tenant"].IsPrimaryKey = true
|
||||
|
||||
current := diffTestDB("public", map[string]map[string]string{"users": {"id": "integer", "tenant": "integer"}})
|
||||
current.Schemas[0].Tables[0].Columns["id"].IsPrimaryKey = true
|
||||
current.Schemas[0].Tables[0].Constraints["pk_public_users"] = &models.Constraint{
|
||||
Name: "pk_public_users", Type: models.PrimaryKeyConstraint, Columns: []string{"id"},
|
||||
}
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(model, current)
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(diffJoin(stmts), "DROP CONSTRAINT IF EXISTS pk_public_users") {
|
||||
t.Fatalf("expected old primary key to be dropped, got:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_ChangedCommentEmitted(t *testing.T) {
|
||||
model := diffTestDB("public", map[string]map[string]string{"users": {"id": "integer"}})
|
||||
model.Schemas[0].Tables[0].Description = "new"
|
||||
current := diffTestDB("public", map[string]map[string]string{"users": {"id": "integer"}})
|
||||
current.Schemas[0].Tables[0].Description = "old"
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(model, current)
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(diffJoin(stmts), "COMMENT ON TABLE public.users IS 'new'") {
|
||||
t.Fatalf("expected changed comment, got:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_TruncatedConstraintNameMatches(t *testing.T) {
|
||||
long := "fk_individual_actor_relationship_rid_typelookup_relationship_type"
|
||||
if len(long) <= maxIdentifierBytes {
|
||||
t.Fatal("test name must exceed the identifier limit")
|
||||
}
|
||||
build := func(name string) *models.Database {
|
||||
db := diffTestDB("public", map[string]map[string]string{
|
||||
"parent": {"id": "integer"},
|
||||
"child": {"rid": "integer"},
|
||||
})
|
||||
for _, tbl := range db.Schemas[0].Tables {
|
||||
if tbl.Name == "child" {
|
||||
tbl.Constraints[name] = &models.Constraint{
|
||||
Name: name, Type: models.ForeignKeyConstraint, Columns: []string{"rid"},
|
||||
ReferencedTable: "parent", ReferencedSchema: "public", ReferencedColumns: []string{"id"},
|
||||
}
|
||||
}
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(build(long), build(pgIdentifier(long)))
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if len(stmts) != 0 {
|
||||
t.Fatalf("expected truncated live name to match, got:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
@@ -46,10 +46,11 @@ func NewMigrationWriter(options *writers.WriterOptions) (*MigrationWriter, error
|
||||
}, nil
|
||||
}
|
||||
|
||||
// WriteMigration generates migration scripts using templates
|
||||
func (w *MigrationWriter) WriteMigration(model, current *models.Database) error {
|
||||
// GenerateScripts computes the differential migration scripts between current and model,
|
||||
// sorted by priority and sequence.
|
||||
func (w *MigrationWriter) GenerateScripts(model, current *models.Database) ([]MigrationScript, error) {
|
||||
if model == nil {
|
||||
return fmt.Errorf("model database is required")
|
||||
return nil, fmt.Errorf("model database is required")
|
||||
}
|
||||
if w.options == nil {
|
||||
w.options = &writers.WriterOptions{}
|
||||
@@ -58,26 +59,6 @@ func (w *MigrationWriter) WriteMigration(model, current *models.Database) error
|
||||
current = models.InitDatabase(model.Name)
|
||||
}
|
||||
|
||||
var writer io.Writer
|
||||
var file *os.File
|
||||
var err error
|
||||
|
||||
// Use existing writer if already set (for testing)
|
||||
if w.writer != nil {
|
||||
writer = w.writer
|
||||
} else if w.options.OutputPath != "" {
|
||||
file, err = os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
defer file.Close()
|
||||
writer = file
|
||||
} else {
|
||||
writer = os.Stdout
|
||||
}
|
||||
|
||||
w.writer = writer
|
||||
|
||||
// Check if audit is configured in metadata
|
||||
var auditConfig *AuditConfig
|
||||
if w.options.Metadata != nil {
|
||||
@@ -93,7 +74,7 @@ func (w *MigrationWriter) WriteMigration(model, current *models.Database) error
|
||||
if auditConfig != nil && len(auditConfig.EnabledTables) > 0 {
|
||||
auditTableScript, err := w.generateAuditTablesScript(auditConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate audit tables: %w", err)
|
||||
return nil, fmt.Errorf("failed to generate audit tables: %w", err)
|
||||
}
|
||||
scripts = append(scripts, auditTableScript...)
|
||||
}
|
||||
@@ -119,7 +100,7 @@ func (w *MigrationWriter) WriteMigration(model, current *models.Database) error
|
||||
// Generate schema-level scripts
|
||||
schemaScripts, err := w.generateSchemaScripts(modelSchema, currentSchema)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate schema scripts: %w", err)
|
||||
return nil, fmt.Errorf("failed to generate schema scripts: %w", err)
|
||||
}
|
||||
scripts = append(scripts, schemaScripts...)
|
||||
|
||||
@@ -127,7 +108,7 @@ func (w *MigrationWriter) WriteMigration(model, current *models.Database) error
|
||||
if auditConfig != nil {
|
||||
auditScripts, err := w.generateAuditScripts(modelSchema, auditConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate audit scripts: %w", err)
|
||||
return nil, fmt.Errorf("failed to generate audit scripts: %w", err)
|
||||
}
|
||||
scripts = append(scripts, auditScripts...)
|
||||
}
|
||||
@@ -141,6 +122,45 @@ func (w *MigrationWriter) WriteMigration(model, current *models.Database) error
|
||||
return scripts[i].Sequence < scripts[j].Sequence
|
||||
})
|
||||
|
||||
return scripts, nil
|
||||
}
|
||||
|
||||
// WriteMigration generates migration scripts using templates
|
||||
func (w *MigrationWriter) WriteMigration(model, current *models.Database) error {
|
||||
if model == nil {
|
||||
return fmt.Errorf("model database is required")
|
||||
}
|
||||
if w.options == nil {
|
||||
w.options = &writers.WriterOptions{}
|
||||
}
|
||||
if current == nil {
|
||||
current = models.InitDatabase(model.Name)
|
||||
}
|
||||
|
||||
scripts, err := w.GenerateScripts(model, current)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var writer io.Writer
|
||||
var file *os.File
|
||||
|
||||
// Use existing writer if already set (for testing)
|
||||
if w.writer != nil {
|
||||
writer = w.writer
|
||||
} else if w.options.OutputPath != "" {
|
||||
file, err = os.Create(w.options.OutputPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
defer file.Close()
|
||||
writer = file
|
||||
} else {
|
||||
writer = os.Stdout
|
||||
}
|
||||
|
||||
w.writer = writer
|
||||
|
||||
// Write header
|
||||
fmt.Fprintf(w.writer, "-- PostgreSQL Migration Script\n")
|
||||
fmt.Fprintf(w.writer, "-- Generated by RelSpec\n")
|
||||
@@ -241,10 +261,14 @@ func (w *MigrationWriter) generateDropScripts(model, current *models.Schema) ([]
|
||||
// Check each constraint in current database
|
||||
for _, currentConstraint := range sortConstraints(currentTable.Constraints) {
|
||||
constraintName := currentConstraint.Name
|
||||
modelConstraint, existsInModel := modelTable.Constraints[constraintName]
|
||||
modelConstraint, existsInModel := lookupConstraint(modelTable.Constraints, constraintName)
|
||||
|
||||
shouldDrop := false
|
||||
if !existsInModel {
|
||||
if currentConstraint.Type == models.PrimaryKeyConstraint {
|
||||
// Model primary keys usually live on the columns (not as a named constraint), so
|
||||
// compare by key columns instead of constraint name.
|
||||
shouldDrop = !primaryKeyColumnsMatch(modelTable, currentConstraint)
|
||||
} else if !existsInModel {
|
||||
shouldDrop = true
|
||||
} else if !constraintsEqual(modelConstraint, currentConstraint) {
|
||||
shouldDrop = true
|
||||
@@ -316,6 +340,12 @@ func (w *MigrationWriter) generateDropScripts(model, current *models.Schema) ([]
|
||||
indexName := currentIndex.Name
|
||||
modelIndex, existsInModel := modelTable.Indexes[indexName]
|
||||
|
||||
// A live database reports the index backing a unique/primary constraint as an
|
||||
// index too; it belongs to the constraint and is handled there.
|
||||
if _, backsConstraint := currentTable.Constraints[indexName]; backsConstraint {
|
||||
continue
|
||||
}
|
||||
|
||||
shouldDrop := false
|
||||
if !existsInModel {
|
||||
shouldDrop = true
|
||||
@@ -470,7 +500,7 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
|
||||
}
|
||||
|
||||
// Check default value changes
|
||||
if !columnDefaultsEqual(modelCol.Default, currentCol.Default) {
|
||||
if !isSerialWithoutDefault(modelCol) && !columnDefaultsEqual(modelCol.Default, currentCol.Default) {
|
||||
setDefault, defaultVal := formatColumnDefaultSQL(modelCol)
|
||||
|
||||
sql, err := w.executor.ExecuteAlterColumnDefaultWithCheck(AlterColumnDefaultWithCheckData{
|
||||
@@ -747,7 +777,7 @@ func (w *MigrationWriter) generateForeignKeyScripts(model, current *models.Schem
|
||||
if !shouldCreate {
|
||||
if currentTable == nil {
|
||||
shouldCreate = true
|
||||
} else if currentConstraint, exists := currentTable.Constraints[constraintName]; !exists {
|
||||
} else if currentConstraint, exists := lookupConstraint(currentTable.Constraints, constraintName); !exists {
|
||||
shouldCreate = true
|
||||
} else if !constraintsEqual(constraint, currentConstraint) {
|
||||
shouldCreate = true
|
||||
@@ -799,12 +829,20 @@ func (w *MigrationWriter) generateForeignKeyScripts(model, current *models.Schem
|
||||
// generateCommentScripts generates COMMENT ON scripts using templates
|
||||
func (w *MigrationWriter) generateCommentScripts(model, current *models.Schema) ([]MigrationScript, error) {
|
||||
scripts := make([]MigrationScript, 0)
|
||||
_ = current // TODO: Compare with current schema to only add new/changed comments
|
||||
|
||||
currentTables := make(map[string]*models.Table)
|
||||
if current != nil {
|
||||
for _, table := range current.Tables {
|
||||
currentTables[strings.ToLower(table.Name)] = table
|
||||
}
|
||||
}
|
||||
|
||||
// Process each model table
|
||||
for _, modelTable := range model.Tables {
|
||||
// Table comment
|
||||
if modelTable.Description != "" {
|
||||
currentTable := currentTables[strings.ToLower(modelTable.Name)]
|
||||
|
||||
// Table comment (skipped when the live table already carries the same comment)
|
||||
if modelTable.Description != "" && (currentTable == nil || strings.TrimSpace(currentTable.Description) != strings.TrimSpace(modelTable.Description)) {
|
||||
sql, err := w.executor.ExecuteCommentTable(CommentTableData{
|
||||
SchemaName: model.Name,
|
||||
TableName: modelTable.Name,
|
||||
@@ -827,7 +865,7 @@ func (w *MigrationWriter) generateCommentScripts(model, current *models.Schema)
|
||||
|
||||
// Column comments
|
||||
for _, col := range sortColumns(modelTable.Columns) {
|
||||
if col.Description != "" {
|
||||
if col.Description != "" && !currentColumnHasDescription(currentTable, col) {
|
||||
sql, err := w.executor.ExecuteCommentColumn(CommentColumnData{
|
||||
SchemaName: model.Name,
|
||||
TableName: modelTable.Name,
|
||||
@@ -994,18 +1032,32 @@ func normalizeDefaultForCompare(value interface{}) string {
|
||||
return ""
|
||||
}
|
||||
if s, ok := value.(string); ok {
|
||||
return strings.TrimSpace(stripBackticks(s))
|
||||
return normalizeDefaultLiteral(strings.TrimSpace(stripBackticks(s)))
|
||||
}
|
||||
return fmt.Sprintf("%v", value)
|
||||
}
|
||||
|
||||
// normalizeDefaultLiteral removes a trailing ::type cast and the quotes around a plain string
|
||||
// literal, so a DBML default ('[]') and the one a live database reports ('[]'::jsonb) compare equal.
|
||||
func normalizeDefaultLiteral(s string) string {
|
||||
if strings.HasPrefix(s, "'") {
|
||||
if end := strings.LastIndex(s, "'"); end > 0 {
|
||||
rest := strings.TrimSpace(s[end+1:])
|
||||
if rest == "" || strings.HasPrefix(rest, "::") {
|
||||
return strings.ReplaceAll(s[1:end], "''", "'")
|
||||
}
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func columnTypesEqual(col1, col2 *models.Column) bool {
|
||||
if col1 == nil || col2 == nil {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(
|
||||
pgsql.NormalizeEquivalentSQLType(effectiveColumnSQLType(col1)),
|
||||
pgsql.NormalizeEquivalentSQLType(effectiveColumnSQLType(col2)),
|
||||
normalizeZeroScale(pgsql.NormalizeEquivalentSQLType(effectiveAlterColumnSQLType(col1))),
|
||||
normalizeZeroScale(pgsql.NormalizeEquivalentSQLType(effectiveAlterColumnSQLType(col2))),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -1041,7 +1093,7 @@ func constraintsEqual(c1, c2 *models.Constraint) bool {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if c1.OnDelete != c2.OnDelete || c1.OnUpdate != c2.OnUpdate {
|
||||
if !fkActionsEqual(c1.OnDelete, c2.OnDelete) || !fkActionsEqual(c1.OnUpdate, c2.OnUpdate) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -1057,7 +1109,7 @@ func indexesEqual(idx1, idx2 *models.Index) bool {
|
||||
if idx1.Unique != idx2.Unique {
|
||||
return false
|
||||
}
|
||||
if !strings.EqualFold(idx1.Type, idx2.Type) {
|
||||
if !strings.EqualFold(normalizeIndexMethod(idx1.Type), normalizeIndexMethod(idx2.Type)) {
|
||||
return false
|
||||
}
|
||||
if len(idx1.Columns) != len(idx2.Columns) {
|
||||
@@ -1077,6 +1129,104 @@ func indexesEqual(idx1, idx2 *models.Index) bool {
|
||||
return indexHintsEqual(indexStorageParameters(idx1.Comment), indexStorageParameters(idx2.Comment))
|
||||
}
|
||||
|
||||
// normalizeZeroScale turns numeric(p,0) into numeric(p); PostgreSQL treats them as the same type.
|
||||
func normalizeZeroScale(sqlType string) string {
|
||||
if i := strings.Index(sqlType, "("); i >= 0 && strings.HasSuffix(sqlType, ",0)") && !strings.Contains(sqlType[i:], "[") {
|
||||
return strings.TrimSuffix(sqlType, ",0)") + ")"
|
||||
}
|
||||
return sqlType
|
||||
}
|
||||
|
||||
// maxIdentifierBytes is PostgreSQL's identifier length limit; longer names are truncated on creation.
|
||||
const maxIdentifierBytes = 63
|
||||
|
||||
// pgIdentifier returns name as PostgreSQL stores it (truncated to 63 bytes).
|
||||
func pgIdentifier(name string) string {
|
||||
if len(name) > maxIdentifierBytes {
|
||||
return name[:maxIdentifierBytes]
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
// lookupConstraint finds a constraint by name, also matching names that differ only because
|
||||
// PostgreSQL truncated an over-long identifier.
|
||||
func lookupConstraint(constraints map[string]*models.Constraint, name string) (*models.Constraint, bool) {
|
||||
if c, ok := constraints[name]; ok {
|
||||
return c, true
|
||||
}
|
||||
want := pgIdentifier(name)
|
||||
for n, c := range constraints {
|
||||
if pgIdentifier(n) == want {
|
||||
return c, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// currentColumnHasDescription reports whether the live table's column already carries
|
||||
// the model column's comment.
|
||||
func currentColumnHasDescription(currentTable *models.Table, col *models.Column) bool {
|
||||
if currentTable == nil {
|
||||
return false
|
||||
}
|
||||
for _, cc := range currentTable.Columns {
|
||||
if strings.EqualFold(cc.Name, col.Name) {
|
||||
return strings.TrimSpace(cc.Description) == strings.TrimSpace(col.Description)
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// primaryKeyColumnsMatch reports whether the model table's primary key (an explicit PK
|
||||
// constraint or columns flagged IsPrimaryKey) has the same columns, in order, as current.
|
||||
func primaryKeyColumnsMatch(modelTable *models.Table, current *models.Constraint) bool {
|
||||
var modelCols []string
|
||||
for _, c := range sortConstraints(modelTable.Constraints) {
|
||||
if c.Type == models.PrimaryKeyConstraint {
|
||||
modelCols = c.Columns
|
||||
break
|
||||
}
|
||||
}
|
||||
if modelCols == nil {
|
||||
for _, col := range getSortedColumns(modelTable.Columns) {
|
||||
if col.IsPrimaryKey {
|
||||
modelCols = append(modelCols, col.Name)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(modelCols) != len(current.Columns) {
|
||||
return false
|
||||
}
|
||||
for i, col := range modelCols {
|
||||
if !strings.EqualFold(col, current.Columns[i]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// normalizeIndexMethod maps an unspecified index method to PostgreSQL's default (btree),
|
||||
// which is what a live database reports for it.
|
||||
func normalizeIndexMethod(method string) string {
|
||||
if strings.TrimSpace(method) == "" {
|
||||
return "btree"
|
||||
}
|
||||
return method
|
||||
}
|
||||
|
||||
// fkActionsEqual compares referential actions case-insensitively, treating an unspecified
|
||||
// action as PostgreSQL's default (NO ACTION).
|
||||
func fkActionsEqual(a, b string) bool {
|
||||
norm := func(s string) string {
|
||||
s = strings.ToUpper(strings.TrimSpace(s))
|
||||
if s == "" {
|
||||
return "NO ACTION"
|
||||
}
|
||||
return s
|
||||
}
|
||||
return norm(a) == norm(b)
|
||||
}
|
||||
|
||||
// indexHintsEqual compares two optional index hints, treating an unspecified hint as a match.
|
||||
func indexHintsEqual(hint1, hint2 string) bool {
|
||||
if hint1 == "" || hint2 == "" {
|
||||
|
||||
@@ -0,0 +1,163 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func serialTestTable(colType string, def interface{}) *models.Table {
|
||||
table := models.InitTable("login_event", "identity")
|
||||
id := models.InitColumn("id_login_event", "login_event", "identity")
|
||||
id.Type = colType
|
||||
id.IsPrimaryKey = true
|
||||
id.Default = def
|
||||
table.Columns["id_login_event"] = id
|
||||
return table
|
||||
}
|
||||
|
||||
func serialTestDB(table *models.Table) *models.Database {
|
||||
db := models.InitDatabase("testdb")
|
||||
schema := models.InitSchema("identity")
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
return db
|
||||
}
|
||||
|
||||
func TestSequences_FollowPrimaryKeyDefault(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
table *models.Table
|
||||
wantContain []string
|
||||
wantAbsent []string
|
||||
}{
|
||||
{
|
||||
name: "nextval default uses its own sequence",
|
||||
table: serialTestTable("bigint", "nextval('identity.login_event_id_login_event_seq'::regclass)"),
|
||||
wantContain: []string{
|
||||
"login_event_id_login_event_seq",
|
||||
"setval(",
|
||||
},
|
||||
wantAbsent: []string{"identity_login_event_id_login_event"},
|
||||
},
|
||||
{
|
||||
name: "no default creates no sequence",
|
||||
table: serialTestTable("bigint", nil),
|
||||
wantAbsent: []string{"CREATE SEQUENCE", "setval("},
|
||||
},
|
||||
{
|
||||
name: "bigserial without default creates no extra sequence",
|
||||
table: serialTestTable("bigserial", nil),
|
||||
wantAbsent: []string{"CREATE SEQUENCE", "setval(", "identity_login_event_id_login_event"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
w.writer = &buf
|
||||
if err := w.WriteDatabase(serialTestDB(tt.table)); err != nil {
|
||||
t.Fatalf("WriteDatabase failed: %v", err)
|
||||
}
|
||||
full := buf.String()
|
||||
|
||||
stmts, err := w.GenerateDatabaseStatements(serialTestDB(tt.table))
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateDatabaseStatements failed: %v", err)
|
||||
}
|
||||
for label, out := range map[string]string{"WriteDatabase": full, "GenerateDatabaseStatements": diffJoin(stmts)} {
|
||||
for _, s := range tt.wantContain {
|
||||
if !strings.Contains(out, s) {
|
||||
t.Errorf("%s: missing %q\n%s", label, s, out)
|
||||
}
|
||||
}
|
||||
for _, s := range tt.wantAbsent {
|
||||
if strings.Contains(out, s) {
|
||||
t.Errorf("%s: unexpected %q\n%s", label, s, out)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetSequenceValue_OnlyMovesForwardPastData(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
table := serialTestTable("bigint", "nextval('identity.login_event_id_login_event_seq'::regclass)")
|
||||
stmt, err := w.primaryKeySetvalStatement(serialTestDB(table).Schemas[0], table)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, want := range []string{"MAX(", "is_called", "m_cnt > m_next", ", m_cnt, false)"} {
|
||||
if !strings.Contains(stmt, want) {
|
||||
t.Errorf("setval statement missing %q:\n%s", want, stmt)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiffStatements_ExistingTableNewSequenceIsSetPastData(t *testing.T) {
|
||||
def := "nextval('identity.login_event_id_login_event_seq'::regclass)"
|
||||
model := serialTestDB(serialTestTable("bigint", def))
|
||||
current := serialTestDB(serialTestTable("bigint", nil))
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(model, current)
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
out := diffJoin(stmts)
|
||||
seqIdx := strings.Index(out, "CREATE SEQUENCE IF NOT EXISTS identity.login_event_id_login_event_seq")
|
||||
setvalIdx := strings.Index(out, "setval(")
|
||||
if seqIdx < 0 || setvalIdx < 0 || seqIdx > setvalIdx {
|
||||
t.Fatalf("expected sequence creation before setval, got:\n%s", out)
|
||||
}
|
||||
|
||||
// Sequence already present in the database: nothing to create or move.
|
||||
current.Schemas[0].Sequences = append(current.Schemas[0].Sequences,
|
||||
models.InitSequence("login_event_id_login_event_seq", "identity"))
|
||||
stmts, err = w.diffStatements(model, current)
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if strings.Contains(diffJoin(stmts), "setval(") || strings.Contains(diffJoin(stmts), "CREATE SEQUENCE") {
|
||||
t.Fatalf("existing sequence must be left alone:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSerialWithoutDefault_DoesNotDropExistingDefault(t *testing.T) {
|
||||
def := "nextval('identity.login_event_id_login_event_seq'::regclass)"
|
||||
model := serialTestDB(serialTestTable("bigserial", nil))
|
||||
current := serialTestDB(serialTestTable("bigserial", def))
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
stmts, err := w.diffStatements(model, current)
|
||||
if err != nil {
|
||||
t.Fatalf("diffStatements failed: %v", err)
|
||||
}
|
||||
if strings.Contains(strings.ToUpper(diffJoin(stmts)), "DROP DEFAULT") {
|
||||
t.Fatalf("serial default must not be dropped:\n%s", diffJoin(stmts))
|
||||
}
|
||||
|
||||
alter, err := w.GenerateAlterColumnDefaultStatements(model.Schemas[0])
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(strings.ToUpper(diffJoin(alter)), "DROP DEFAULT") {
|
||||
t.Fatalf("full-DDL path must not drop serial default:\n%s", diffJoin(alter))
|
||||
}
|
||||
|
||||
// A plain bigint with no model default still has its default dropped.
|
||||
model = serialTestDB(serialTestTable("bigint", nil))
|
||||
current = serialTestDB(serialTestTable("bigint", def))
|
||||
stmts, err = w.diffStatements(model, current)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(strings.ToUpper(diffJoin(stmts)), "DROP DEFAULT") {
|
||||
t.Fatalf("non-serial default removal should still be emitted:\n%s", diffJoin(stmts))
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
DO $$
|
||||
DECLARE
|
||||
m_cnt bigint;
|
||||
m_next bigint;
|
||||
BEGIN
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM pg_class c
|
||||
@@ -12,8 +13,15 @@ BEGIN
|
||||
SELECT COALESCE(MAX({{quote_ident .ColumnName}}), 0) + 1
|
||||
FROM {{qual_table .SchemaName .TableName}}
|
||||
INTO m_cnt;
|
||||
|
||||
PERFORM setval('{{qual_table_raw .SchemaName .SequenceName}}'::regclass, m_cnt);
|
||||
|
||||
SELECT CASE WHEN is_called THEN last_value + 1 ELSE last_value END
|
||||
FROM {{qual_table .SchemaName .SequenceName}}
|
||||
INTO m_next;
|
||||
|
||||
-- Only move the sequence forward; never rewind one that is already past the data.
|
||||
IF m_cnt > m_next THEN
|
||||
PERFORM setval('{{qual_table_raw .SchemaName .SequenceName}}'::regclass, m_cnt, false);
|
||||
END IF;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
+231
-36
@@ -14,6 +14,8 @@ import (
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
rpgsql "git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
@@ -139,6 +141,50 @@ func (w *Writer) GenerateDatabaseStatements(db *models.Database) ([]string, erro
|
||||
return statements, nil
|
||||
}
|
||||
|
||||
// primaryKeySequenceName returns the name of the sequence behind a table's integer
|
||||
// primary key nextval() default, or "" when the key has no such default.
|
||||
func primaryKeySequenceName(table *models.Table) string {
|
||||
pk := table.GetPrimaryKey()
|
||||
if pk == nil || !isIntegerType(pk.Type) || pk.Default == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
defaultStr, ok := pk.Default.(string)
|
||||
if !ok || !strings.Contains(strings.ToLower(defaultStr), "nextval") {
|
||||
return ""
|
||||
}
|
||||
|
||||
return extractSequenceName(defaultStr)
|
||||
}
|
||||
|
||||
// primaryKeySequenceStatement returns the CREATE SEQUENCE statement backing a table's
|
||||
// integer primary key nextval() default, or "" when the table has no such sequence.
|
||||
func (w *Writer) primaryKeySequenceStatement(schema *models.Schema, table *models.Table) string {
|
||||
seqName := primaryKeySequenceName(table)
|
||||
if seqName == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
return fmt.Sprintf("CREATE SEQUENCE IF NOT EXISTS %s\n INCREMENT 1\n MINVALUE 1\n MAXVALUE 9223372036854775807\n START 1\n CACHE 1",
|
||||
w.qualTable(schema.SQLName(), seqName))
|
||||
}
|
||||
|
||||
// primaryKeySetvalStatement returns the statement that moves a table's primary key
|
||||
// sequence past the existing rows, or "" when the table has no such sequence.
|
||||
func (w *Writer) primaryKeySetvalStatement(schema *models.Schema, table *models.Table) (string, error) {
|
||||
seqName := primaryKeySequenceName(table)
|
||||
if seqName == "" {
|
||||
return "", nil
|
||||
}
|
||||
pk := table.GetPrimaryKey()
|
||||
return w.executor.ExecuteSetSequenceValue(SetSequenceValueData{
|
||||
SchemaName: schema.Name,
|
||||
TableName: table.Name,
|
||||
SequenceName: seqName,
|
||||
ColumnName: pk.Name,
|
||||
})
|
||||
}
|
||||
|
||||
// GenerateSchemaStatements generates SQL statements as a list for a single schema
|
||||
func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, error) {
|
||||
statements := []string{}
|
||||
@@ -159,24 +205,9 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
|
||||
// Phase 2: Create sequences
|
||||
for _, table := range schema.Tables {
|
||||
pk := table.GetPrimaryKey()
|
||||
if pk == nil || !isIntegerType(pk.Type) || pk.Default == "" {
|
||||
continue
|
||||
if stmt := w.primaryKeySequenceStatement(schema, table); stmt != "" {
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
|
||||
defaultStr, ok := pk.Default.(string)
|
||||
if !ok || !strings.Contains(strings.ToLower(defaultStr), "nextval") {
|
||||
continue
|
||||
}
|
||||
|
||||
seqName := extractSequenceName(defaultStr)
|
||||
if seqName == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
stmt := fmt.Sprintf("CREATE SEQUENCE IF NOT EXISTS %s\n INCREMENT 1\n MINVALUE 1\n MAXVALUE 9223372036854775807\n START 1\n CACHE 1",
|
||||
w.qualTable(schema.SQLName(), seqName))
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
|
||||
// Phase 3: Create tables
|
||||
@@ -391,6 +422,17 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 6.5: Move primary key sequences past existing rows
|
||||
for _, table := range schema.Tables {
|
||||
stmt, err := w.primaryKeySetvalStatement(schema, table)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate set sequence value for %s.%s: %w", schema.Name, table.Name, err)
|
||||
}
|
||||
if stmt != "" {
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 7: Comments
|
||||
for _, table := range schema.Tables {
|
||||
if table.Comment != "" {
|
||||
@@ -505,6 +547,11 @@ func (w *Writer) GenerateAlterColumnDefaultStatements(schema *models.Schema) ([]
|
||||
// backing sequence), so there is nothing for this generator to manage.
|
||||
continue
|
||||
}
|
||||
if isSerialWithoutDefault(col) {
|
||||
// serial/bigserial columns get their nextval() default from the
|
||||
// type itself; a model with no explicit default must not drop it.
|
||||
continue
|
||||
}
|
||||
setDefault, defaultVal := formatColumnDefaultSQL(col)
|
||||
stmt, err := w.executor.ExecuteAlterColumnDefaultWithCheck(AlterColumnDefaultWithCheckData{
|
||||
SchemaName: schema.Name,
|
||||
@@ -523,6 +570,19 @@ func (w *Writer) GenerateAlterColumnDefaultStatements(schema *models.Schema) ([]
|
||||
return statements, nil
|
||||
}
|
||||
|
||||
// isSerialWithoutDefault reports whether col is a serial-family column with no
|
||||
// explicit default, i.e. one whose nextval() default is implied by its type.
|
||||
func isSerialWithoutDefault(col *models.Column) bool {
|
||||
if col == nil || col.Default != nil {
|
||||
return false
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(col.Type)) {
|
||||
case "serial", "bigserial", "smallserial", "serial4", "serial8", "serial2":
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// formatColumnDefaultSQL renders a column's model-level default into the
|
||||
// SQL literal/expression used by ALTER COLUMN ... SET DEFAULT, shared by
|
||||
// the full-schema writer and the diff-based migration writer.
|
||||
@@ -872,18 +932,12 @@ func (w *Writer) writeSequences(schema *models.Schema) error {
|
||||
fmt.Fprintf(w.writer, "-- Sequences for schema: %s\n", schema.Name)
|
||||
|
||||
for _, table := range schema.Tables {
|
||||
pk := table.GetPrimaryKey()
|
||||
if pk == nil {
|
||||
// Only create the sequence the primary key's nextval() default uses
|
||||
seqName := primaryKeySequenceName(table)
|
||||
if seqName == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// Only create sequences for integer-type PKs with identity
|
||||
if !isIntegerType(pk.Type) {
|
||||
continue
|
||||
}
|
||||
|
||||
seqName := fmt.Sprintf("identity_%s_%s", table.SQLName(), pk.SQLName())
|
||||
|
||||
data := CreateSequenceData{
|
||||
SchemaName: schema.Name,
|
||||
SequenceName: seqName,
|
||||
@@ -1423,12 +1477,11 @@ func (w *Writer) writeSetSequenceValues(schema *models.Schema) error {
|
||||
fmt.Fprintf(w.writer, "-- Set sequence values for schema: %s\n", schema.Name)
|
||||
|
||||
for _, table := range schema.Tables {
|
||||
pk := table.GetPrimaryKey()
|
||||
if pk == nil || !isIntegerType(pk.Type) {
|
||||
seqName := primaryKeySequenceName(table)
|
||||
if seqName == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
seqName := fmt.Sprintf("identity_%s_%s", table.SQLName(), pk.SQLName())
|
||||
pk := table.GetPrimaryKey()
|
||||
|
||||
// Use template executor to generate set sequence value statement
|
||||
data := SetSequenceValueData{
|
||||
@@ -1993,7 +2046,9 @@ func extractSequenceName(defaultExpr string) string {
|
||||
return fullName
|
||||
}
|
||||
|
||||
// executeDatabaseSQL executes SQL statements directly on a PostgreSQL database
|
||||
// executeDatabaseSQL applies db to a PostgreSQL database. By default it reads the live schema
|
||||
// and executes only the differences; Metadata["full_ddl"]=true (or a failed/unsupported live
|
||||
// read) executes the full idempotent DDL instead.
|
||||
func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) error {
|
||||
// Initialize execution report
|
||||
w.executionReport = &ExecutionReport{
|
||||
@@ -2002,13 +2057,148 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
|
||||
Errors: make([]ExecutionError, 0),
|
||||
}
|
||||
|
||||
// Generating a large schema can take time before any statement is executed.
|
||||
fmt.Fprintln(os.Stderr, " → Generating PostgreSQL statements...")
|
||||
statements, err := w.GenerateDatabaseStatements(db)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate SQL statements: %w", err)
|
||||
var statements []string
|
||||
if fullDDL, _ := w.options.Metadata["full_ddl"].(bool); !fullDDL {
|
||||
diffStatements, err := w.generateLiveDiffStatements(db, connString)
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "⚠ Warning: live diff unavailable (%v); falling back to full DDL\n", err)
|
||||
} else {
|
||||
statements = diffStatements
|
||||
if len(statements) == 0 {
|
||||
w.executionReport.EndTime = getCurrentTimestamp()
|
||||
fmt.Fprintln(os.Stderr, "✓ Database is already up to date; nothing to execute")
|
||||
return w.finishReport()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if statements == nil {
|
||||
// Generating a large schema can take time before any statement is executed.
|
||||
fmt.Fprintln(os.Stderr, " → Generating PostgreSQL statements...")
|
||||
var err error
|
||||
statements, err = w.GenerateDatabaseStatements(db)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate SQL statements: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return w.executeStatements(statements, connString)
|
||||
}
|
||||
|
||||
// generateLiveDiffStatements reads the live database and returns only the statements needed
|
||||
// to bring it in line with db. An error means the diff could not be computed.
|
||||
func (w *Writer) generateLiveDiffStatements(db *models.Database, connString string) ([]string, error) {
|
||||
if w.options.FlattenSchema {
|
||||
return nil, fmt.Errorf("flatten_schema output cannot be compared with the live schemas")
|
||||
}
|
||||
|
||||
fmt.Fprintln(os.Stderr, " → Reading live database to compute differences...")
|
||||
current, err := rpgsql.NewReader(&readers.ReaderOptions{ConnectionString: connString}).ReadDatabase()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read live database: %w", err)
|
||||
}
|
||||
|
||||
return w.diffStatements(db, current)
|
||||
}
|
||||
|
||||
// diffStatements returns the ordered statements that migrate current to model.
|
||||
func (w *Writer) diffStatements(model, current *models.Database) ([]string, error) {
|
||||
currentSchemas := make(map[string]*models.Schema)
|
||||
for _, cs := range current.Schemas {
|
||||
if cs != nil {
|
||||
currentSchemas[strings.ToLower(cs.Name)] = cs
|
||||
}
|
||||
}
|
||||
|
||||
statements := []string{}
|
||||
// Primary key sequences that must be moved past existing rows once the
|
||||
// migration scripts (which set the column defaults) have run.
|
||||
var setvalStatements []string
|
||||
for _, schema := range model.Schemas {
|
||||
if schema == nil {
|
||||
continue
|
||||
}
|
||||
if err := w.checkDirectives(schema); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cs := currentSchemas[strings.ToLower(schema.Name)]
|
||||
if cs == nil && schema.Name != "public" {
|
||||
statements = append(statements,
|
||||
fmt.Sprintf("-- Schema: %s", schema.Name),
|
||||
fmt.Sprintf("CREATE SCHEMA IF NOT EXISTS %s", schema.SQLName()))
|
||||
}
|
||||
|
||||
existing := make(map[string]*models.Table)
|
||||
knownSequences := make(map[string]bool)
|
||||
if cs != nil {
|
||||
for _, t := range cs.Tables {
|
||||
existing[strings.ToLower(t.Name)] = t
|
||||
}
|
||||
for _, seq := range cs.Sequences {
|
||||
knownSequences[strings.ToLower(seq.Name)] = true
|
||||
}
|
||||
}
|
||||
for _, table := range schema.Tables {
|
||||
currentTable, ok := existing[strings.ToLower(table.Name)]
|
||||
if !ok {
|
||||
// New table: no rows yet, so the sequence only needs creating.
|
||||
if stmt := w.primaryKeySequenceStatement(schema, table); stmt != "" {
|
||||
statements = append(statements, stmt)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Existing table whose primary key is being pointed at a sequence the
|
||||
// database does not have yet: create it and move it past the data,
|
||||
// otherwise it restarts at 1 and collides with existing rows.
|
||||
seqName := primaryKeySequenceName(table)
|
||||
if seqName == "" || knownSequences[strings.ToLower(seqName)] {
|
||||
continue
|
||||
}
|
||||
if cpk := currentTable.GetPrimaryKey(); cpk != nil && columnDefaultsEqual(table.GetPrimaryKey().Default, cpk.Default) {
|
||||
continue
|
||||
}
|
||||
statements = append(statements, w.primaryKeySequenceStatement(schema, table))
|
||||
setval, err := w.primaryKeySetvalStatement(schema, table)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate set sequence value for %s.%s: %w", schema.Name, table.Name, err)
|
||||
}
|
||||
setvalStatements = append(setvalStatements, setval)
|
||||
}
|
||||
}
|
||||
|
||||
mw, err := NewMigrationWriter(w.options)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scripts, err := mw.GenerateScripts(model, current)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate migration scripts: %w", err)
|
||||
}
|
||||
|
||||
lastSchema := ""
|
||||
for _, script := range scripts {
|
||||
body := strings.TrimSpace(script.Body)
|
||||
if body == "" {
|
||||
continue
|
||||
}
|
||||
if script.Schema != "" && script.Schema != lastSchema {
|
||||
statements = append(statements, fmt.Sprintf("-- Schema: %s", script.Schema))
|
||||
lastSchema = script.Schema
|
||||
}
|
||||
statements = append(statements, body)
|
||||
}
|
||||
statements = append(statements, setvalStatements...)
|
||||
|
||||
if dump := os.Getenv("ZZDUMP"); dump != "" {
|
||||
_ = os.WriteFile(dump, []byte(strings.Join(statements, "\n=====\n")), 0o644)
|
||||
}
|
||||
return statements, nil
|
||||
}
|
||||
|
||||
// executeStatements runs statements one by one against the database and writes the report.
|
||||
func (w *Writer) executeStatements(statements []string, connString string) error {
|
||||
w.executionReport.TotalStatements = len(statements)
|
||||
|
||||
// Connect to database
|
||||
@@ -2108,6 +2298,11 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
|
||||
}
|
||||
|
||||
w.executionReport.EndTime = getCurrentTimestamp()
|
||||
return w.finishReport()
|
||||
}
|
||||
|
||||
// finishReport writes the optional report file and prints the execution summary.
|
||||
func (w *Writer) finishReport() error {
|
||||
|
||||
// Write report if path is specified
|
||||
if reportPath, ok := w.options.Metadata["report_path"].(string); ok && reportPath != "" {
|
||||
|
||||
Reference in New Issue
Block a user