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:
2026-10-02 22:37:37 +02:00
parent 43849324ce
commit d7d1d99ebc
5 changed files with 975 additions and 77 deletions
+189 -39
View File
@@ -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 == "" {