From d7d1d99ebc7d8dea7620510acec6718a81b819f2 Mon Sep 17 00:00:00 2001 From: Hein Date: Fri, 2 Oct 2026 22:37:37 +0200 Subject: [PATCH] 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__ 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 --- pkg/writers/pgsql/diff_statements_test.go | 382 ++++++++++++++++++ pkg/writers/pgsql/migration_writer.go | 228 +++++++++-- pkg/writers/pgsql/serial_sequence_test.go | 163 ++++++++ .../pgsql/templates/set_sequence_value.tmpl | 12 +- pkg/writers/pgsql/writer.go | 267 ++++++++++-- 5 files changed, 975 insertions(+), 77 deletions(-) create mode 100644 pkg/writers/pgsql/diff_statements_test.go create mode 100644 pkg/writers/pgsql/serial_sequence_test.go diff --git a/pkg/writers/pgsql/diff_statements_test.go b/pkg/writers/pgsql/diff_statements_test.go new file mode 100644 index 0000000..eafd8b9 --- /dev/null +++ b/pkg/writers/pgsql/diff_statements_test.go @@ -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)) + } +} diff --git a/pkg/writers/pgsql/migration_writer.go b/pkg/writers/pgsql/migration_writer.go index 5fc3f7d..5f2ca5c 100644 --- a/pkg/writers/pgsql/migration_writer.go +++ b/pkg/writers/pgsql/migration_writer.go @@ -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 == "" { diff --git a/pkg/writers/pgsql/serial_sequence_test.go b/pkg/writers/pgsql/serial_sequence_test.go new file mode 100644 index 0000000..ad23d85 --- /dev/null +++ b/pkg/writers/pgsql/serial_sequence_test.go @@ -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)) + } +} diff --git a/pkg/writers/pgsql/templates/set_sequence_value.tmpl b/pkg/writers/pgsql/templates/set_sequence_value.tmpl index c94f79e..aaee5a0 100644 --- a/pkg/writers/pgsql/templates/set_sequence_value.tmpl +++ b/pkg/writers/pgsql/templates/set_sequence_value.tmpl @@ -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; $$; \ No newline at end of file diff --git a/pkg/writers/pgsql/writer.go b/pkg/writers/pgsql/writer.go index bcdeb7f..7e37ab6 100644 --- a/pkg/writers/pgsql/writer.go +++ b/pkg/writers/pgsql/writer.go @@ -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 != "" {