package models import ( "reflect" "testing" ) // viewFixture builds a two-schema database whose map contents would randomise output order. func viewFixture() *Database { db := InitDatabase("shop") db.Description = "desc" db.DatabaseType = PostgresqlDatabaseType db.DatabaseVersion = "16" for _, sn := range []string{"sales", "public"} { s := InitSchema(sn) s.Owner = "owner_" + sn s.Scripts = append(s.Scripts, InitScript("seed")) users := InitTable("users", sn) for _, cn := range []string{"id", "email", "name"} { c := InitColumn(cn, "users", sn) c.Type = "text" users.Columns[cn] = c } users.Columns["id"].IsPrimaryKey = true users.Columns["id"].NotNull = true pk := InitConstraint("users_pkey", PrimaryKeyConstraint) pk.Columns = []string{"id"} users.Constraints["users_pkey"] = pk ck := InitConstraint("users_ck", CheckConstraint) ck.Expression = "id > 0" users.Constraints["users_ck"] = ck users.Indexes["users_idx"] = InitIndex("users_idx", "users", sn) orders := InitTable("orders", sn) oid := InitColumn("id", "orders", sn) orders.Columns["id"] = oid uid := InitColumn("user_id", "orders", sn) orders.Columns["user_id"] = uid fk := InitConstraint("orders_user_fk", ForeignKeyConstraint) fk.Columns = []string{"user_id"} fk.ReferencedSchema = sn fk.ReferencedTable = "users" fk.ReferencedColumns = []string{"id"} fk.OnDelete = "CASCADE" orders.Constraints["orders_user_fk"] = fk rel := InitRelationship("orders_users", RelationType("one_to_many")) rel.FromTable, rel.FromSchema = "orders", sn rel.ToTable, rel.ToSchema = "users", sn rel.ForeignKey = "orders_user_fk" rel.ThroughTable, rel.ThroughSchema = "link", sn orders.Relationships["orders_users"] = rel plain := InitRelationship("plain", RelationType("one_to_one")) plain.FromTable, plain.FromSchema = "orders", sn plain.ToTable, plain.ToSchema = "users", sn orders.Relationships["plain"] = plain s.Tables = append(s.Tables, users, orders) db.Schemas = append(db.Schemas, s) } return db } func TestToFlatColumns(t *testing.T) { db := viewFixture() first := db.ToFlatColumns() if len(first) != 2*(3+2) { t.Fatalf("got %d columns", len(first)) } for i := 1; i < len(first); i++ { if first[i-1].FullyQualifiedName >= first[i].FullyQualifiedName { t.Fatalf("not sorted at %d: %s >= %s", i, first[i-1].FullyQualifiedName, first[i].FullyQualifiedName) } } if first[0].FullyQualifiedName != "shop.public.orders.id" { t.Errorf("first: %s", first[0].FullyQualifiedName) } var id *FlatColumn for _, c := range first { if c.FullyQualifiedName == "shop.sales.users.id" { id = c } } if id == nil || !id.IsPrimaryKey || !id.NotNull || id.Type != "text" || id.DatabaseName != "shop" || id.SchemaName != "sales" || id.TableName != "users" || id.ColumnName != "id" { t.Errorf("flat id column: %+v", id) } for i := 0; i < 20; i++ { if !reflect.DeepEqual(first, db.ToFlatColumns()) { t.Fatal("ToFlatColumns not deterministic") } } if got := InitDatabase("e").ToFlatColumns(); got == nil || len(got) != 0 { t.Errorf("empty db: %v", got) } } func TestToFlatTables(t *testing.T) { got := viewFixture().ToFlatTables() if len(got) != 4 { t.Fatalf("got %d tables", len(got)) } // schema order follows the database slice: sales first if got[0].FullyQualifiedName != "shop.sales.users" || got[0].ColumnCount != 3 || got[0].ConstraintCount != 2 || got[0].IndexCount != 1 { t.Errorf("first: %+v", got[0]) } if got[1].FullyQualifiedName != "shop.sales.orders" || got[1].ColumnCount != 2 || got[1].ConstraintCount != 1 { t.Errorf("second: %+v", got[1]) } if got := InitDatabase("e").ToFlatTables(); got == nil || len(got) != 0 { t.Errorf("empty db: %v", got) } } func TestToFlatConstraints(t *testing.T) { db := viewFixture() got := db.ToFlatConstraints() if len(got) != 6 { t.Fatalf("got %d constraints", len(got)) } for i := 1; i < len(got); i++ { if got[i-1].FullyQualifiedName >= got[i].FullyQualifiedName { t.Fatalf("not sorted: %s >= %s", got[i-1].FullyQualifiedName, got[i].FullyQualifiedName) } } var fk, ck *FlatConstraint for _, c := range got { switch c.FullyQualifiedName { case "shop.sales.orders.orders_user_fk": fk = c case "shop.sales.users.users_ck": ck = c } } if fk == nil || fk.ReferencedFQN != "shop.sales.users" || fk.OnDelete != "CASCADE" || fk.Type != ForeignKeyConstraint { t.Errorf("fk: %+v", fk) } if ck == nil || ck.ReferencedFQN != "" || ck.Expression != "id > 0" { t.Errorf("check: %+v", ck) } // FK without a referenced table gets no FQN. db2 := InitDatabase("d") s := InitSchema("s") tb := InitTable("t", "s") tb.Constraints["fk"] = InitConstraint("fk", ForeignKeyConstraint) s.Tables = append(s.Tables, tb) db2.Schemas = append(db2.Schemas, s) if out := db2.ToFlatConstraints(); len(out) != 1 || out[0].ReferencedFQN != "" { t.Errorf("unreferenced fk: %+v", out) } if got := InitDatabase("e").ToFlatConstraints(); got == nil || len(got) != 0 { t.Errorf("empty db: %v", got) } } func TestToFlatRelationships(t *testing.T) { db := viewFixture() got := db.ToFlatRelationships() if len(got) != 4 { t.Fatalf("got %d relationships", len(got)) } for i := 1; i < len(got); i++ { a, b := got[i-1], got[i] if a.FromFQN > b.FromFQN || (a.FromFQN == b.FromFQN && a.RelationshipName > b.RelationshipName) { t.Fatalf("not sorted at %d", i) } } var through, plain *FlatRelationship for _, r := range got { if r.FromSchema == "sales" && r.RelationshipName == "orders_users" { through = r } if r.FromSchema == "sales" && r.RelationshipName == "plain" { plain = r } } if through == nil || through.ThroughTableFQN != "shop.sales.link" || through.FromFQN != "shop.sales.orders" || through.ToFQN != "shop.sales.users" || through.ForeignKey != "orders_user_fk" { t.Errorf("through: %+v", through) } if plain == nil || plain.ThroughTableFQN != "" { t.Errorf("plain: %+v", plain) } for i := 0; i < 20; i++ { if !reflect.DeepEqual(got, db.ToFlatRelationships()) { t.Fatal("ToFlatRelationships not deterministic") } } if got := InitDatabase("e").ToFlatRelationships(); got == nil || len(got) != 0 { t.Errorf("empty db: %v", got) } } func TestSummaries(t *testing.T) { db := viewFixture() ds := db.ToSummary() if ds.Name != "shop" || ds.Description != "desc" || ds.DatabaseType != PostgresqlDatabaseType || ds.DatabaseVersion != "16" || ds.SchemaCount != 2 || ds.TotalTables != 4 || ds.TotalColumns != 10 { t.Errorf("database summary: %+v", ds) } if es := InitDatabase("e").ToSummary(); es.SchemaCount != 0 || es.TotalTables != 0 || es.TotalColumns != 0 { t.Errorf("empty summary: %+v", es) } ss := db.Schemas[0].ToSummary() if ss.Name != "sales" || ss.Owner != "owner_sales" || ss.TableCount != 2 || ss.ScriptCount != 1 || ss.TotalColumns != 5 || ss.TotalConstraints != 3 { t.Errorf("schema summary: %+v", ss) } users := db.Schemas[0].Tables[0].ToSummary() if users.Name != "users" || users.Schema != "sales" || users.ColumnCount != 3 || users.ConstraintCount != 2 || users.IndexCount != 1 || users.RelationshipCount != 0 || !users.HasPrimaryKey || users.ForeignKeyCount != 0 { t.Errorf("users summary: %+v", users) } orders := db.Schemas[0].Tables[1].ToSummary() if orders.HasPrimaryKey || orders.ForeignKeyCount != 1 || orders.RelationshipCount != 2 { t.Errorf("orders summary: %+v", orders) } } func TestDirectiveFromAny(t *testing.T) { want := Directive{Namespace: "postgres", Key: "partition", Args: "partition by RANGE (x)", Line: 7} tests := []struct { name string in any want Directive ok bool }{ {"directive", want, want, true}, {"string map int line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": 7}, want, true}, {"string map int64 line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": int64(7)}, want, true}, {"string map float64 line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": float64(7)}, want, true}, {"any map ignores non-string keys", map[any]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": 7, 5: "x"}, want, true}, {"key derived from args", map[string]any{"namespace": "sqlite", "args": "WITHOUT ROWID extra"}, Directive{Namespace: "sqlite", Key: "without", Args: "WITHOUT ROWID extra"}, true}, {"wrong field types ignored", map[string]any{"namespace": 1, "key": 2, "args": 3, "line": "x"}, Directive{}, true}, {"unsupported type", "nope", Directive{}, false}, {"nil", nil, Directive{}, false}, {"int", 5, Directive{}, false}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got, ok := directiveFromAny(tt.in) if ok != tt.ok || got != tt.want { t.Errorf("got (%+v,%v), want (%+v,%v)", got, ok, tt.want, tt.ok) } }) } }