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
+231 -36
View File
@@ -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 != "" {