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:
+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