package merge import ( "strings" "testing" "git.warky.dev/wdevs/relspecgo/pkg/models" ) func TestMergeSequences(t *testing.T) { target := models.InitSchema("public") target.Sequences = []*models.Sequence{{Name: "Existing", StartValue: 1, IncrementBy: 1}} source := models.InitSchema("public") source.Sequences = []*models.Sequence{ {Name: "existing", StartValue: 100, IncrementBy: 10}, // conflicting: must not overwrite {Name: "fresh", StartValue: 5, IncrementBy: 2, MinValue: 1, MaxValue: 99, CacheSize: 3, Cycle: true, OwnedByTable: "t", OwnedByColumn: "id", Comment: "c", Description: "d"}, } res := &MergeResult{} res.mergeSequences(target, source) if res.SequencesAdded != 1 || len(target.Sequences) != 2 { t.Fatalf("added=%d len=%d", res.SequencesAdded, len(target.Sequences)) } if target.Sequences[0].StartValue != 1 || target.Sequences[0].IncrementBy != 1 { t.Errorf("existing sequence was modified: %+v", target.Sequences[0]) } added := target.Sequences[1] if added.Name != "fresh" || added.StartValue != 5 || added.IncrementBy != 2 || added.MinValue != 1 || added.MaxValue != 99 || added.CacheSize != 3 || !added.Cycle || added.OwnedByTable != "t" || added.OwnedByColumn != "id" || added.Comment != "c" || added.Description != "d" { t.Errorf("clone lost fields: %+v", added) } if added == source.Sequences[1] { t.Error("sequence must be cloned, not shared") } source.Sequences[1].StartValue = 777 if added.StartValue != 5 { t.Error("clone must be independent of source") } if cloneSequence(nil) != nil { t.Error("cloneSequence(nil) must be nil") } } func TestCloneSchemaIsIndependent(t *testing.T) { src := models.InitSchema("public") src.Description, src.Owner, src.Comment, src.Sequence = "d", "o", "c", 4 src.Permissions["r"] = "all" src.Metadata["k"] = "v" src.Scripts = []*models.Script{{Name: "s"}} tbl := models.InitTable("t", "public") col := models.InitColumn("id", "t", "public") col.Type = "integer" tbl.Columns["id"] = col tbl.Constraints["pk"] = &models.Constraint{Name: "pk", Type: models.PrimaryKeyConstraint, Columns: []string{"id"}} tbl.Indexes["i"] = &models.Index{Name: "i", Columns: []string{"id"}, Include: []string{"x"}} tbl.Metadata["tm"] = 1 src.Tables = []*models.Table{tbl} v := models.InitView("v", "public") v.Definition = "select 1" v.Columns["c"] = &models.Column{Name: "c"} v.Metadata["vm"] = 1 src.Views = []*models.View{v} src.Sequences = []*models.Sequence{{Name: "sq", StartValue: 3}} src.Enums = []*models.Enum{{Name: "e", Values: []string{"a", "b"}}} src.Relations = []*models.Relationship{{Name: "r", FromColumns: []string{"a"}, ToColumns: []string{"b"}, Properties: map[string]string{"p": "q"}}} got := cloneSchema(src) if got == src || got.Name != "public" || got.Description != "d" || got.Owner != "o" || got.Comment != "c" || got.Sequence != 4 { t.Fatalf("scalar fields: %+v", got) } if got.Permissions["r"] != "all" || got.Metadata["k"] != "v" || len(got.Scripts) != 1 { t.Errorf("maps/scripts: %+v", got) } if len(got.Tables) != 1 || got.Tables[0] == tbl || got.Tables[0].Columns["id"] == col || got.Tables[0].Columns["id"].Type != "integer" { t.Errorf("tables not deep cloned: %+v", got.Tables) } if len(got.Views) != 1 || got.Views[0] == v || got.Views[0].Definition != "select 1" || got.Views[0].Columns["c"] == v.Columns["c"] || got.Views[0].Metadata["vm"] != 1 { t.Errorf("views not deep cloned: %+v", got.Views) } if len(got.Sequences) != 1 || got.Sequences[0] == src.Sequences[0] || got.Sequences[0].StartValue != 3 { t.Errorf("sequences: %+v", got.Sequences) } if len(got.Enums) != 1 || got.Enums[0] == src.Enums[0] || strings.Join(got.Enums[0].Values, ",") != "a,b" { t.Errorf("enums: %+v", got.Enums) } if len(got.Relations) != 1 || got.Relations[0] == src.Relations[0] || got.Relations[0].Properties["p"] != "q" { t.Errorf("relations: %+v", got.Relations) } // Mutating the clone must not touch the source. got.Permissions["r"] = "none" got.Metadata["k"] = "changed" got.Tables[0].Columns["id"].Type = "text" got.Tables[0].Constraints["pk"].Columns[0] = "zzz" got.Tables[0].Indexes["i"].Columns[0] = "zzz" got.Tables[0].Metadata["tm"] = 2 got.Enums[0].Values[0] = "zzz" got.Relations[0].FromColumns[0] = "zzz" got.Relations[0].Properties["p"] = "zzz" got.Views[0].Columns["c"].Name = "zzz" if src.Permissions["r"] != "all" || src.Metadata["k"] != "v" || col.Type != "integer" || tbl.Constraints["pk"].Columns[0] != "id" || tbl.Indexes["i"].Columns[0] != "id" || tbl.Metadata["tm"] != 1 || src.Enums[0].Values[0] != "a" || src.Relations[0].FromColumns[0] != "a" || src.Relations[0].Properties["p"] != "q" || v.Columns["c"].Name != "c" { t.Error("clone shares state with the source") } if cloneSchema(nil) != nil { t.Error("cloneSchema(nil) must be nil") } bare := cloneSchema(&models.Schema{Name: "bare"}) if bare.Permissions != nil || bare.Metadata != nil { t.Errorf("nil maps must stay nil: %+v", bare) } } func TestCloneNilInputs(t *testing.T) { if cloneTable(nil) != nil || cloneColumn(nil) != nil || cloneConstraint(nil) != nil || cloneIndex(nil) != nil || cloneView(nil) != nil || cloneEnum(nil) != nil || cloneRelation(nil) != nil || cloneDomain(nil) != nil { t.Error("clone of nil must be nil") } } func TestCloneDomainAndRelation(t *testing.T) { d := &models.Domain{Name: "d", Description: "x", Comment: "c", Sequence: 2, Metadata: map[string]any{"k": 1}, Tables: []*models.DomainTable{{TableName: "t", SchemaName: "s"}}} cd := cloneDomain(d) if cd == d || cd.Name != "d" || cd.Description != "x" || cd.Comment != "c" || cd.Sequence != 2 || cd.Metadata["k"] != 1 || len(cd.Tables) != 1 { t.Errorf("domain clone: %+v", cd) } cd.Metadata["k"] = 2 if d.Metadata["k"] != 1 { t.Error("domain metadata shared") } r := &models.Relationship{Name: "r", Type: "one_to_many", FromTable: "a", FromSchema: "s", ToTable: "b", ToSchema: "s", ForeignKey: "fk", ThroughTable: "l", ThroughSchema: "s", Description: "d", Sequence: 3} cr := cloneRelation(r) if cr == r || cr.Name != "r" || cr.Type != "one_to_many" || cr.FromTable != "a" || cr.ToTable != "b" || cr.ForeignKey != "fk" || cr.ThroughTable != "l" || cr.Description != "d" || cr.Sequence != 3 { t.Errorf("relation clone: %+v", cr) } if cr.Properties != nil { t.Errorf("nil properties must stay nil") } } func TestExtractTypeParts(t *testing.T) { tests := []struct { name string col models.Column wantType string wantLen, wantPrec, wantScale int }{ {"plain", models.Column{Type: "TEXT"}, "text", 0, 0, 0}, {"trim and lower", models.Column{Type: " Integer "}, "integer", 0, 0, 0}, {"embedded length", models.Column{Type: "varchar(50)"}, "varchar", 50, 0, 0}, {"embedded precision and scale", models.Column{Type: "numeric(10,2)"}, "numeric", 0, 10, 2}, {"embedded with spaces", models.Column{Type: "numeric( 10 , 2 )"}, "numeric", 0, 10, 2}, {"fields win over embedded precision", models.Column{Type: "numeric(10,2)", Precision: 12, Scale: 4}, "numeric", 0, 12, 4}, {"fields win over embedded length", models.Column{Type: "varchar(50)", Length: 80}, "varchar", 80, 0, 0}, {"precision field blocks embedded length", models.Column{Type: "varchar(50)", Precision: 5}, "varchar", 0, 5, 0}, {"non-numeric modifier", models.Column{Type: "varchar(max)"}, "varchar", 0, 0, 0}, {"zero modifier", models.Column{Type: "char(0)"}, "char", 0, 0, 0}, {"serial sugar", models.Column{Type: "bigserial"}, "bigint", 0, 0, 0}, {"smallserial sugar", models.Column{Type: "smallserial"}, "smallint", 0, 0, 0}, {"three modifiers ignored", models.Column{Type: "x(1,2,3)"}, "x", 0, 0, 0}, {"empty", models.Column{}, "", 0, 0, 0}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { col := tt.col gt, gl, gp, gs := extractTypeParts(&col) if gt != tt.wantType || gl != tt.wantLen || gp != tt.wantPrec || gs != tt.wantScale { t.Errorf("got (%q,%d,%d,%d), want (%q,%d,%d,%d)", gt, gl, gp, gs, tt.wantType, tt.wantLen, tt.wantPrec, tt.wantScale) } }) } } func TestColumnTypeConflict(t *testing.T) { c := func(typ string, l, p, s int) *models.Column { return &models.Column{Type: typ, Length: l, Precision: p, Scale: s} } tests := []struct { name string a, b *models.Column want bool }{ {"nil target", nil, c("text", 0, 0, 0), false}, {"nil source", c("text", 0, 0, 0), nil, false}, {"same", c("text", 0, 0, 0), c("TEXT", 0, 0, 0), false}, {"different base", c("text", 0, 0, 0), c("integer", 0, 0, 0), true}, {"embedded equals field", c("varchar(50)", 0, 0, 0), c("varchar", 50, 0, 0), false}, {"different length", c("varchar", 50, 0, 0), c("varchar", 80, 0, 0), true}, {"different scale", c("numeric", 0, 10, 2), c("numeric", 0, 10, 3), true}, {"serial vs int", c("bigserial", 0, 0, 0), c("bigint", 0, 0, 0), false}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { if got := columnTypeConflict(tt.a, tt.b); got != tt.want { t.Errorf("got %v, want %v", got, tt.want) } }) } } func TestDescribeColumnType(t *testing.T) { tests := []struct { col *models.Column want string }{ {nil, ""}, {&models.Column{}, ""}, {&models.Column{Type: " "}, ""}, {&models.Column{Type: "text"}, "text"}, {&models.Column{Type: " numeric ", Precision: 10, Scale: 2}, "numeric(10,2)"}, {&models.Column{Type: "numeric", Precision: 10}, "numeric(10)"}, {&models.Column{Type: "varchar", Length: 50}, "varchar(50)"}, {&models.Column{Type: "varchar", Length: 50, Precision: 7}, "varchar(7)"}, } for _, tt := range tests { if got := describeColumnType(tt.col); got != tt.want { t.Errorf("describeColumnType(%+v) = %q, want %q", tt.col, got, tt.want) } } } func TestFirstNonEmpty(t *testing.T) { if got := firstNonEmpty("", " ", "x", "y"); got != "x" { t.Errorf("got %q", got) } if got := firstNonEmpty(); got != "" { t.Errorf("none: %q", got) } if got := firstNonEmpty("", " "); got != "" { t.Errorf("all blank: %q", got) } } func TestGetColumnTypeConflictSummary(t *testing.T) { conflicts := []ColumnTypeConflict{ {Schema: "s", Table: "t", Column: "a", TargetType: "text", SourceType: "integer"}, {Schema: "s", Table: "t", Column: "b", TargetType: "int", SourceType: "text"}, {Schema: "s", Table: "u", Column: "c", TargetType: "x", SourceType: "y"}, } res := &MergeResult{TypeConflicts: conflicts} if GetColumnTypeConflictSummary(nil, 5) != "" || GetColumnTypeConflictSummary(&MergeResult{}, 5) != "" { t.Error("no conflicts must yield empty summary") } all := GetColumnTypeConflictSummary(res, 0) if !strings.Contains(all, "column type conflicts detected:") || !strings.Contains(all, "s.t.a: target=text source=integer") || !strings.Contains(all, "s.u.c: target=x source=y") || strings.Contains(all, "more") { t.Errorf("unlimited summary:\n%s", all) } if neg := GetColumnTypeConflictSummary(res, -1); neg != all { t.Error("negative limit must behave as unlimited") } limited := GetColumnTypeConflictSummary(res, 2) if !strings.Contains(limited, "s.t.b") || strings.Contains(limited, "s.u.c") || !strings.HasSuffix(limited, "... and 1 more") { t.Errorf("limited summary:\n%s", limited) } exact := GetColumnTypeConflictSummary(res, 3) if strings.Contains(exact, "more") { t.Errorf("limit == len must not truncate:\n%s", exact) } } func TestMinHelper(t *testing.T) { if min(1, 2) != 1 || min(2, 1) != 1 || min(3, 3) != 3 { t.Error("min") } }