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:
@@ -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 == "" {
|
||||
|
||||
Reference in New Issue
Block a user