package diff import ( "reflect" "testing" "git.warky.dev/wdevs/relspecgo/pkg/models" ) func TestCompareSchemaDetails(t *testing.T) { mk := func() *models.Schema { s := models.InitSchema("public") s.Tables = []*models.Table{models.InitTable("t", "public")} return s } if got := compareSchemaDetails(mk(), mk()); got != nil { t.Errorf("identical schemas must yield nil, got %+v", got) } tests := []struct { name string mutate func(*models.Schema) check func(*SchemaChange) bool }{ {"table added", func(s *models.Schema) { s.Tables = append(s.Tables, models.InitTable("u", "public")) }, func(c *SchemaChange) bool { return c.Tables != nil && len(c.Tables.Extra) == 1 }}, {"view added", func(s *models.Schema) { s.Views = []*models.View{models.InitView("v", "public")} }, func(c *SchemaChange) bool { return c.Views != nil && len(c.Views.Extra) == 1 }}, {"sequence added", func(s *models.Schema) { s.Sequences = []*models.Sequence{models.InitSequence("sq", "public")} }, func(c *SchemaChange) bool { return c.Sequences != nil && len(c.Sequences.Extra) == 1 }}, {"script added", func(s *models.Schema) { s.Scripts = []*models.Script{models.InitScript("sc")} }, func(c *SchemaChange) bool { return c.Scripts != nil && len(c.Scripts.Extra) == 1 }}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { target := mk() tt.mutate(target) got := compareSchemaDetails(mk(), target) if got == nil || got.Name != "public" || !tt.check(got) { t.Errorf("unexpected change: %+v", got) } }) } } func TestCompareConstraintDetails(t *testing.T) { base := func() *models.Constraint { c := models.InitConstraint("fk", models.ForeignKeyConstraint) c.Columns = []string{"a"} c.ReferencedTable = "users" c.ReferencedColumns = []string{"id"} c.OnDelete = "CASCADE" c.OnUpdate = "NO ACTION" return c } if got := compareConstraintDetails(base(), base()); len(got) != 0 { t.Errorf("identical: %v", got) } tests := []struct { name string mutate func(*models.Constraint) wantKey string }{ {"type", func(c *models.Constraint) { c.Type = models.UniqueConstraint }, "type"}, {"columns", func(c *models.Constraint) { c.Columns = []string{"b"} }, "columns"}, {"referenced table", func(c *models.Constraint) { c.ReferencedTable = "other" }, "referenced_table"}, {"referenced columns", func(c *models.Constraint) { c.ReferencedColumns = []string{"x"} }, "referenced_columns"}, {"on delete", func(c *models.Constraint) { c.OnDelete = "SET NULL" }, "on_delete"}, {"on update", func(c *models.Constraint) { c.OnUpdate = "CASCADE" }, "on_update"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { target := base() tt.mutate(target) got := compareConstraintDetails(base(), target) if _, ok := got[tt.wantKey]; !ok || len(got) != 1 { t.Errorf("got %v, want only %q", got, tt.wantKey) } }) } // Action spelling variants that mean the same thing are not changes. a, b := base(), base() a.OnDelete, b.OnDelete = "cascade", " CASCADE " a.OnUpdate, b.OnUpdate = "", "no action" if got := compareConstraintDetails(a, b); len(got) != 0 { t.Errorf("equivalent actions reported as changes: %v", got) } } func TestNormalizeConstraintAction(t *testing.T) { tests := []struct{ in, want string }{ {"", ""}, {"NO ACTION", ""}, {"no action", ""}, {" No Action ", ""}, {"cascade", "CASCADE"}, {" set null ", "SET NULL"}, {"RESTRICT", "RESTRICT"}, } for _, tt := range tests { if got := normalizeConstraintAction(tt.in); got != tt.want { t.Errorf("normalizeConstraintAction(%q) = %q, want %q", tt.in, got, tt.want) } } } func TestConstraintCompareKey(t *testing.T) { uq := &models.Constraint{Name: "UQ_Name", Type: models.UniqueConstraint} if got := constraintCompareKey(uq); got != "uq_name" { t.Errorf("non-FK key: %q", got) } fk := func(name string) *models.Constraint { return &models.Constraint{ Name: name, Type: models.ForeignKeyConstraint, Schema: "Public", Table: "Orders", Columns: []string{"user_id"}, ReferencedSchema: "Public", ReferencedTable: "Users", ReferencedColumns: []string{"id"}, } } if constraintCompareKey(fk("a")) != constraintCompareKey(fk("b")) { t.Error("FK key must ignore the constraint name") } other := fk("a") other.ReferencedColumns = []string{"uid"} if constraintCompareKey(fk("a")) == constraintCompareKey(other) { t.Error("FK key must include referenced columns") } } func TestFilterPrimaryKeyConstraints(t *testing.T) { in := map[string]*models.Constraint{ "pk": {Name: "pk", Type: models.PrimaryKeyConstraint}, "uq": {Name: "uq", Type: models.UniqueConstraint}, "fk": {Name: "fk", Type: models.ForeignKeyConstraint}, } got := filterPrimaryKeyConstraints(in) if len(got) != 2 || got["pk"] != nil || got["uq"] == nil || got["fk"] == nil { t.Errorf("got %v", got) } if len(in) != 3 { t.Error("input must not be modified") } if got := filterPrimaryKeyConstraints(nil); got == nil || len(got) != 0 { t.Errorf("nil: %v", got) } } func TestCompareRelationshipDetails(t *testing.T) { base := func() *models.Relationship { r := models.InitRelationship("r", models.RelationType("one_to_many")) r.FromTable, r.ToTable = "orders", "users" r.FromColumns, r.ToColumns = []string{"user_id"}, []string{"id"} return r } if got := compareRelationshipDetails(base(), base()); len(got) != 0 { t.Errorf("identical: %v", got) } tests := []struct { name string mutate func(*models.Relationship) wantKey string }{ {"type", func(r *models.Relationship) { r.Type = "one_to_one" }, "type"}, {"from table", func(r *models.Relationship) { r.FromTable = "x" }, "from_table"}, {"to table", func(r *models.Relationship) { r.ToTable = "x" }, "to_table"}, {"from columns", func(r *models.Relationship) { r.FromColumns = []string{"x"} }, "from_columns"}, {"to columns", func(r *models.Relationship) { r.ToColumns = []string{"x"} }, "to_columns"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { target := base() tt.mutate(target) got := compareRelationshipDetails(base(), target) if _, ok := got[tt.wantKey]; !ok || len(got) != 1 { t.Errorf("got %v", got) } }) } } func TestCompareRelationshipsModified(t *testing.T) { src := map[string]*models.Relationship{ "same": {Name: "same", Type: "one_to_many"}, "changed": {Name: "changed", Type: "one_to_many"}, "missing": {Name: "missing"}, } tgt := map[string]*models.Relationship{ "same": {Name: "same", Type: "one_to_many"}, "changed": {Name: "changed", Type: "many_to_many"}, "extra": {Name: "extra"}, } d := compareRelationships(src, tgt) if len(d.Missing) != 1 || d.Missing[0].Name != "missing" || len(d.Extra) != 1 || d.Extra[0].Name != "extra" || len(d.Modified) != 1 || d.Modified[0].Name != "changed" { t.Errorf("got %+v", d) } if _, ok := d.Modified[0].Changes["type"]; !ok { t.Errorf("changes: %v", d.Modified[0].Changes) } } func TestCompareViews(t *testing.T) { v := func(name, def string) *models.View { return &models.View{Name: name, Definition: def} } src := []*models.View{v("Keep", "select 1"), v("Changed", "select 1"), v("Gone", "select 1")} tgt := []*models.View{v("keep", "select 1"), v("changed", "select 2"), v("New", "select 1")} d := compareViews(src, tgt) if len(d.Missing) != 1 || d.Missing[0].Name != "Gone" { t.Errorf("missing: %+v", d.Missing) } if len(d.Extra) != 1 || d.Extra[0].Name != "New" { t.Errorf("extra: %+v", d.Extra) } if len(d.Modified) != 1 || d.Modified[0].Name != "changed" || d.Modified[0].Source.Definition != "select 1" || d.Modified[0].Target.Definition != "select 2" { t.Errorf("modified: %+v", d.Modified) } want := map[string]any{"definition": map[string]string{"source": "select 1", "target": "select 2"}} if !reflect.DeepEqual(d.Modified[0].Changes, want) { t.Errorf("changes: %v", d.Modified[0].Changes) } if !isEmpty(compareViews(nil, nil)) { t.Error("nil views must be empty") } if got := compareViewDetails(v("a", "x"), v("a", "x")); len(got) != 0 { t.Errorf("same definition: %v", got) } } func TestCompareSequences(t *testing.T) { seq := func(name string, start, inc, min, max int64, cycle bool) *models.Sequence { return &models.Sequence{Name: name, StartValue: start, IncrementBy: inc, MinValue: min, MaxValue: max, Cycle: cycle} } src := []*models.Sequence{seq("Same", 1, 1, 1, 100, false), seq("Diff", 1, 1, 1, 100, false), seq("Gone", 1, 1, 1, 1, false)} tgt := []*models.Sequence{seq("same", 1, 1, 1, 100, false), seq("diff", 5, 2, 3, 200, true), seq("New", 1, 1, 1, 1, false)} d := compareSequences(src, tgt) if len(d.Missing) != 1 || d.Missing[0].Name != "Gone" || len(d.Extra) != 1 || d.Extra[0].Name != "New" || len(d.Modified) != 1 { t.Fatalf("got %+v", d) } ch := d.Modified[0].Changes for _, key := range []string{"start_value", "increment_by", "min_value", "max_value", "cycle"} { if _, ok := ch[key]; !ok { t.Errorf("missing change key %q in %v", key, ch) } } if got := ch["increment_by"].(map[string]int64); got["source"] != 1 || got["target"] != 2 { t.Errorf("increment_by: %v", got) } if got := ch["cycle"].(map[string]bool); got["source"] || !got["target"] { t.Errorf("cycle: %v", got) } if got := compareSequenceDetails(seq("a", 1, 1, 1, 1, false), seq("a", 1, 1, 1, 1, false)); len(got) != 0 { t.Errorf("identical: %v", got) } } func TestCompareScriptDetailsAllFields(t *testing.T) { a := &models.Script{Name: "s", SQL: "a", Rollback: "ra", RunAfter: []string{"x"}, Schema: "p", Version: "1", Priority: 1, Sequence: 1} b := &models.Script{Name: "s", SQL: "b", Rollback: "rb", RunAfter: []string{"y"}, Schema: "q", Version: "2", Priority: 2, Sequence: 2} got := compareScriptDetails(a, b) for _, key := range []string{"sql", "rollback", "run_after", "schema", "version", "priority", "sequence"} { if _, ok := got[key]; !ok { t.Errorf("missing %q in %v", key, got) } } if got := compareScriptDetails(a, a); len(got) != 0 { t.Errorf("identical: %v", got) } } func TestIsEmptyAllTypes(t *testing.T) { if !isEmpty(&ViewDiff{}) || !isEmpty(&SequenceDiff{}) { t.Error("empty view/sequence diffs must be empty") } if isEmpty(&ViewDiff{Extra: []*models.View{{Name: "v"}}}) || isEmpty(&SequenceDiff{Modified: []*SequenceChange{{Name: "s"}}}) { t.Error("non-empty diffs reported as empty") } if isEmpty(&ConstraintDiff{Modified: []*ConstraintChange{{Name: "c"}}}) || isEmpty(&RelationshipDiff{Missing: []*models.Relationship{{Name: "r"}}}) { t.Error("non-empty diffs reported as empty") } if isEmpty(&IndexDiff{Modified: []*IndexChange{{Name: "i"}}}) || isEmpty(&TableDiff{Modified: []*TableChange{{Name: "t"}}}) { t.Error("non-empty diffs reported as empty") } if isEmpty("something else") || isEmpty(nil) { t.Error("unknown types must not be treated as empty") } } func TestComputeSummaryFullTree(t *testing.T) { res := &DiffResult{Schemas: &SchemaDiff{ Missing: []*models.Schema{{Name: "m"}}, Extra: []*models.Schema{{Name: "e"}}, Modified: []*SchemaChange{{ Name: "public", Tables: &TableDiff{ Missing: []*models.Table{{Name: "a"}}, Extra: []*models.Table{{Name: "b"}, {Name: "c"}}, Modified: []*TableChange{{ Name: "t", Columns: &ColumnDiff{Missing: []*models.Column{{}}, Extra: []*models.Column{{}, {}}, Modified: []*ColumnChange{{}}}, Indexes: &IndexDiff{Missing: []*models.Index{{}}, Extra: []*models.Index{{}}, Modified: []*IndexChange{{}, {}}}, Constraints: &ConstraintDiff{Missing: []*models.Constraint{{}}, Modified: []*ConstraintChange{{}}}, Relationships: &RelationshipDiff{Extra: []*models.Relationship{{}}}, }}, }, Views: &ViewDiff{Missing: []*models.View{{}}, Extra: []*models.View{{}}, Modified: []*ViewChange{{}}}, Sequences: &SequenceDiff{Missing: []*models.Sequence{{}}, Extra: []*models.Sequence{{}, {}}}, Scripts: &ScriptDiff{Modified: []*ScriptChange{{}}}, }}, }} s := ComputeSummary(res) checks := []struct { name string got [3]int want [3]int }{ {"schemas", [3]int{s.Schemas.Missing, s.Schemas.Extra, s.Schemas.Modified}, [3]int{1, 1, 1}}, {"tables", [3]int{s.Tables.Missing, s.Tables.Extra, s.Tables.Modified}, [3]int{1, 2, 1}}, {"columns", [3]int{s.Columns.Missing, s.Columns.Extra, s.Columns.Modified}, [3]int{1, 2, 1}}, {"indexes", [3]int{s.Indexes.Missing, s.Indexes.Extra, s.Indexes.Modified}, [3]int{1, 1, 2}}, {"constraints", [3]int{s.Constraints.Missing, s.Constraints.Extra, s.Constraints.Modified}, [3]int{1, 0, 1}}, {"relationships", [3]int{s.Relationships.Missing, s.Relationships.Extra, s.Relationships.Modified}, [3]int{0, 1, 0}}, {"views", [3]int{s.Views.Missing, s.Views.Extra, s.Views.Modified}, [3]int{1, 1, 1}}, {"sequences", [3]int{s.Sequences.Missing, s.Sequences.Extra, s.Sequences.Modified}, [3]int{1, 2, 0}}, {"scripts", [3]int{s.Scripts.Missing, s.Scripts.Extra, s.Scripts.Modified}, [3]int{0, 0, 1}}, } for _, c := range checks { if c.got != c.want { t.Errorf("%s: got %v, want %v", c.name, c.got, c.want) } } if got := ComputeSummary(&DiffResult{}); got == nil || got.Schemas != (SchemaSummary{}) { t.Errorf("nil Schemas: %+v", got) } }