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