Files
relspecgo/pkg/writers/pgsql/serial_sequence_test.go
T
warkanum d7d1d99ebc 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
2026-10-02 22:37:37 +02:00

164 lines
5.3 KiB
Go

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))
}
}