From c3de1bf87d289472533cd3fdbfc299a9eb47f984 Mon Sep 17 00:00:00 2001 From: Divyam Talwar Date: Tue, 22 Sep 2026 06:18:03 +0530 Subject: [PATCH] fix: preserve table coordinates across ERD renderers Bare relation names and dot-concatenated schema keys collide for ordinary same-named cross-schema tables and valid identifiers containing periods. Use structured table coordinates for catalog assembly, index attachment, relationship resolution and every renderer. Keep display labels separate from identity and retain compact Mermaid IDs when unambiguous. Preserve source/target schema metadata on edges, use deterministic opaque IDs for unsafe or colliding Mermaid names, and avoid parsing FK display text as coordinates. Legacy unique-name edges remain compatible; ambiguous legacy edges do not silently guess a schema. Real renderer regressions fail on 9e41414 and the complete ERD race suite passes 20 repetitions. Add a disposable catalog fixture whose cleanup is limited to successfully created schemas. PostgreSQL/parser verification is required before publishing. No dependency or collector SQL changes. --- internal/erd/cross_schema_identity_test.go | 222 ++++++++++++ internal/erd/erd.go | 319 ++++++++++++++---- internal/erd/html.go | 46 +-- internal/erd/introspect.go | 30 +- internal/erd/introspect_integration_test.go | 123 +++++++ .../erd/parent_identity_regression_test.go | 17 + internal/erd/row.go | 64 ++-- 7 files changed, 683 insertions(+), 138 deletions(-) create mode 100644 internal/erd/cross_schema_identity_test.go create mode 100644 internal/erd/introspect_integration_test.go create mode 100644 internal/erd/parent_identity_regression_test.go diff --git a/internal/erd/cross_schema_identity_test.go b/internal/erd/cross_schema_identity_test.go new file mode 100644 index 0000000..28bd7be --- /dev/null +++ b/internal/erd/cross_schema_identity_test.go @@ -0,0 +1,222 @@ +package erd + +import ( + "html" + "regexp" + "strings" + "testing" +) + +func TestRenderersDistinguishQualifiedTableIdentity(t *testing.T) { + s := Schema{Tables: []Table{ + {Schema: "a.b", Name: "users", Columns: []Column{{Name: "id", Type: "bigint", PK: true}}}, + {Schema: "a", Name: "b.users", Columns: []Column{{Name: "id", Type: "bigint", PK: true}}}, + {Schema: "auth", Name: "users", Columns: []Column{{Name: "id", Type: "bigint", PK: true}}}, + {Schema: `销售 "区"`, Name: "订单 表", Columns: []Column{{Name: "id", Type: "bigint", PK: true}}}, + }} + + wantLabels := []string{ + `"a.b".users`, + `a."b.users"`, + `auth.users`, + `"销售 ""区"""."订单 表"`, + } + for name, out := range map[string]string{ + "column ASCII": RenderASCII(s, false), + "row ASCII": RenderASCIIRow(s), + } { + for _, label := range wantLabels { + if strings.Count(out, label) != 1 { + t.Errorf("%s must render qualified label %q exactly once:\n%s", name, label, out) + } + } + } + + htmlOut := html.UnescapeString(RenderHTML(s)) + for _, label := range wantLabels { + if strings.Count(htmlOut, ">"+label+"") != 1 { + t.Errorf("HTML must render qualified label %q exactly once", label) + } + } + + mermaid := RenderMermaid(s) + aliases := regexp.MustCompile(`(?m)^ (table_[0-9a-f]{64})\["([^"]*)"\] \{$`).FindAllStringSubmatch(mermaid, -1) + if len(aliases) != len(s.Tables) { + t.Fatalf("Mermaid must declare each unsafe or colliding table through an opaque ID and alias; got %d declarations:\n%s", len(aliases), mermaid) + } + ids := map[string]bool{} + for _, match := range aliases { + if ids[match[1]] { + t.Fatalf("Mermaid reused opaque entity ID %q:\n%s", match[1], mermaid) + } + ids[match[1]] = true + } + for _, alias := range []string{ + `#quot;a.b#quot;.users`, + `a.#quot;b.users#quot;`, + `auth.users`, + `#quot;销售 #quot;#quot;区#quot;#quot;#quot;.#quot;订单 表#quot;`, + } { + if !strings.Contains(mermaid, `[`+`"`+alias+`"`+`]`) { + t.Errorf("Mermaid missing safe alias %q:\n%s", alias, mermaid) + } + } +} + +func TestRenderersRouteSchemaQualifiedEdges(t *testing.T) { + s := Schema{ + Tables: []Table{ + {Schema: "auth", Name: "users", Columns: []Column{{Name: "id", Type: "bigint", PK: true}}}, + {Schema: "crm", Name: "accounts", Columns: []Column{ + {Name: "id", Type: "bigint", PK: true}, + {Name: "user_id", Type: "bigint", FKTarget: "crm.users.id"}, + }}, + {Schema: "crm", Name: "users", Columns: []Column{{Name: "id", Type: "bigint", PK: true}}}, + {Schema: "sales", Name: "orders", Columns: []Column{ + {Name: "id", Type: "bigint", PK: true}, + {Name: "buyer_id", Type: "bigint", FKTarget: "crm.users.id"}, + }}, + {Schema: "sales", Name: "refunds", Columns: []Column{ + {Name: "id", Type: "bigint", PK: true}, + {Name: "order_id", Type: "bigint", FKTarget: "orders.id"}, + }}, + }, + Edges: []Edge{ + {FromSchema: "crm", FromTable: "accounts", FromColumn: "user_id", ToSchema: "crm", ToTable: "users", ToColumn: "id"}, + {FromSchema: "sales", FromTable: "orders", FromColumn: "buyer_id", ToSchema: "crm", ToTable: "users", ToColumn: "id"}, + {FromSchema: "sales", FromTable: "refunds", FromColumn: "order_id", ToSchema: "sales", ToTable: "orders", ToColumn: "id"}, + }, + } + + column := RenderASCII(s, false) + columnDiagram := strings.Split(column, "\nRelationships\n")[0] + if got := strings.Count(columnDiagram, "▶"); got != 2 { + t.Errorf("column ASCII routed %d distinct target arrows, want 2:\n%s", got, column) + } + for _, want := range []string{ + "crm.users", + "├─< accounts (user_id)", + "└─< orders (buyer_id)", + "└─< refunds (order_id)", + } { + if !strings.Contains(column, want) { + t.Errorf("column ASCII/forest missing %q:\n%s", want, column) + } + } + + row := RenderASCIIRow(s) + rowDiagram := strings.Split(row, "\nRelationships\n")[0] + if got := strings.Count(rowDiagram, "<"); got != 2 { + t.Errorf("row ASCII routed %d distinct target arrows, want 2:\n%s", got, row) + } + + htmlOut := RenderHTML(s) + if got := strings.Count(htmlOut, `"+html.EscapeString(label)+"") + if labelAt < 0 { + return "" + } + groupAt := strings.LastIndex(rendered[:labelAt], " 0 { - parent = parent[:dot] - } - fkRows = append(fkRows, struct { - row int - target string - }{len(lines) + 1 + i, parent}) - } + columnRow[columnIdentity{Table: id, Column: c.Name}] = len(lines) + 1 + i } lines = append(lines, boxLines...) lines = append(lines, "") } - for _, fk := range fkRows { - if pr, ok := titleRow[fk.target]; ok { - conns = append(conns, conn{childRow: fk.row, parentRow: pr}) + for _, e := range resolveEdges(tables, s.Edges) { + child, childOK := columnRow[columnIdentity{Table: e.From, Column: e.Edge.FromColumn}] + parent, parentOK := titleRow[e.To] + if childOK && parentOK { + conns = append(conns, conn{childRow: child, parentRow: parent}) } } @@ -244,6 +238,171 @@ func maxInt(a, b int) int { return b } +type tableIdentity struct { + Schema string + Name string +} + +type columnIdentity struct { + Table tableIdentity + Column string +} + +type resolvedEdge struct { + Edge Edge + From tableIdentity + To tableIdentity + FromPresent bool + ToPresent bool +} + +func identityOf(t Table) tableIdentity { + return tableIdentity{Schema: t.Schema, Name: t.Name} +} + +func lessIdentity(a, b tableIdentity) bool { + if a.Schema != b.Schema { + return a.Schema < b.Schema + } + return a.Name < b.Name +} + +func resolveTableIdentity(tables []Table, schema, name string) (tableIdentity, bool, bool) { + var match tableIdentity + matches := 0 + for _, table := range tables { + if table.Name != name || (schema != "" && table.Schema != schema) { + continue + } + match = identityOf(table) + matches++ + } + if matches == 1 { + return match, true, true + } + if matches > 1 { + return tableIdentity{}, false, false + } + return tableIdentity{Schema: schema, Name: name}, false, true +} + +func resolveEdges(tables []Table, edges []Edge) []resolvedEdge { + resolved := make([]resolvedEdge, 0, len(edges)) + for _, edge := range edges { + from, fromPresent, fromOK := resolveTableIdentity(tables, edge.FromSchema, edge.FromTable) + to, toPresent, toOK := resolveTableIdentity(tables, edge.ToSchema, edge.ToTable) + if fromOK && toOK { + resolved = append(resolved, resolvedEdge{ + Edge: edge, From: from, To: to, + FromPresent: fromPresent, ToPresent: toPresent, + }) + } + } + sort.Slice(resolved, func(i, j int) bool { + if resolved[i].To != resolved[j].To { + return lessIdentity(resolved[i].To, resolved[j].To) + } + if resolved[i].From != resolved[j].From { + return lessIdentity(resolved[i].From, resolved[j].From) + } + if resolved[i].Edge.FromColumn != resolved[j].Edge.FromColumn { + return resolved[i].Edge.FromColumn < resolved[j].Edge.FromColumn + } + return resolved[i].Edge.ToColumn < resolved[j].Edge.ToColumn + }) + return resolved +} + +func drawableEdges(edges []resolvedEdge) []resolvedEdge { + drawable := make([]resolvedEdge, 0, len(edges)) + for _, edge := range edges { + if edge.FromPresent && edge.ToPresent { + drawable = append(drawable, edge) + } + } + return drawable +} + +func simpleIdentifier(s string) bool { + if s == "" || !((s[0] >= 'a' && s[0] <= 'z') || s[0] == '_') { + return false + } + for i := 1; i < len(s); i++ { + c := s[i] + if !((c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '_' || c == '$') { + return false + } + } + return true +} + +func displayIdentifier(s string) string { + if simpleIdentifier(s) { + return s + } + return `"` + strings.ReplaceAll(s, `"`, `""`) + `"` +} + +func qualifiedTableName(id tableIdentity) string { + if id.Schema == "" { + return displayIdentifier(id.Name) + } + return displayIdentifier(id.Schema) + "." + displayIdentifier(id.Name) +} + +func duplicateIdentityNames(tables []Table, edges []resolvedEdge) map[string]bool { + identities := map[tableIdentity]bool{} + for _, table := range tables { + identities[identityOf(table)] = true + } + for _, edge := range edges { + identities[edge.From] = true + identities[edge.To] = true + } + counts := map[string]int{} + for id := range identities { + counts[id.Name]++ + } + duplicates := map[string]bool{} + for name, count := range counts { + duplicates[name] = count > 1 + } + return duplicates +} + +func forestTableName(id tableIdentity, duplicates map[string]bool, present bool) string { + if duplicates[id.Name] || (!present && id.Schema != "") { + return qualifiedTableName(id) + } + return id.Name +} + +func mermaidEntityID(id tableIdentity) string { + payload := fmt.Sprintf("%d:%s%d:%s", len(id.Schema), id.Schema, len(id.Name), id.Name) + sum := sha256.Sum256([]byte(payload)) + return fmt.Sprintf("table_%x", sum) +} + +func mermaidSafeEntityName(name string) bool { + if !simpleIdentifier(name) || strings.Contains(name, "$") { + return false + } + switch strings.ToLower(name) { + case "class", "classdef", "direction", "end", "erdiagram", "many", "one", "only", "style", "subgraph", "to", "u", "zero": + return false + } + return true +} + +func mermaidText(s string) string { + return strings.NewReplacer( + "#", "#35;", + `"`, "#quot;", + "\r", "#13;", + "\n", "#10;", + ).Replace(s) +} + // writeTableBox renders one table: // // ┌─ public.orders ───────────────────┐ @@ -278,7 +437,7 @@ func writeTableBox(b *strings.Builder, t Table) { } idxRows = append(idxRows, row) } - title := t.Schema + "." + t.Name + title := qualifiedTableName(identityOf(t)) inner := len(title) + 4 for _, r := range append(append([]string(nil), rows...), idxRows...) { inner = max(inner, len(r)+2) @@ -305,40 +464,39 @@ func writeTableBox(b *strings.Builder, t Table) { // Each child appears once, under its first (alphabetical) parent; additional // parents show as a cross-link. Cycle-safe via a visited set. func writeForest(b *strings.Builder, s Schema) { - if len(s.Edges) == 0 { + edges := resolveEdges(s.Tables, s.Edges) + if len(edges) == 0 { return } b.WriteString("Relationships\n") - children := map[string][]Edge{} // parent → edges into it - firstParent := map[string]string{} - hasParent := map[string]bool{} - edges := append([]Edge(nil), s.Edges...) - sort.Slice(edges, func(i, j int) bool { - if edges[i].ToTable != edges[j].ToTable { - return edges[i].ToTable < edges[j].ToTable - } - return edges[i].FromTable < edges[j].FromTable - }) + children := map[tableIdentity][]resolvedEdge{} // parent → edges into it + firstParent := map[tableIdentity]tableIdentity{} + hasParent := map[tableIdentity]bool{} for _, e := range edges { - children[e.ToTable] = append(children[e.ToTable], e) - hasParent[e.FromTable] = true - if _, ok := firstParent[e.FromTable]; !ok { - firstParent[e.FromTable] = e.ToTable + children[e.To] = append(children[e.To], e) + hasParent[e.From] = true + if _, ok := firstParent[e.From]; !ok { + firstParent[e.From] = e.To } } + duplicates := duplicateIdentityNames(s.Tables, edges) + present := map[tableIdentity]bool{} + for _, table := range s.Tables { + present[identityOf(table)] = true + } - var roots []string + var roots []tableIdentity for parent := range children { if !hasParent[parent] { roots = append(roots, parent) } } - sort.Strings(roots) + sort.Slice(roots, func(i, j int) bool { return lessIdentity(roots[i], roots[j]) }) - visited := map[string]bool{} - var walk func(table, indent string) - walk = func(table, indent string) { + visited := map[tableIdentity]bool{} + var walk func(table tableIdentity, indent string) + walk = func(table tableIdentity, indent string) { if visited[table] { return } @@ -351,30 +509,30 @@ func writeForest(b *strings.Builder, s Schema) { branch = "└─<" childIndent = indent + " " } - line := fmt.Sprintf("%s%s %s (%s)", indent, branch, e.FromTable, e.FromColumn) - if firstParent[e.FromTable] != table { + line := fmt.Sprintf("%s%s %s (%s)", indent, branch, forestTableName(e.From, duplicates, present[e.From]), e.Edge.FromColumn) + if firstParent[e.From] != table { line += " · also above" fmt.Fprintln(b, line) continue } fmt.Fprintln(b, line) - walk(e.FromTable, childIndent) + walk(e.From, childIndent) } } for _, r := range roots { - fmt.Fprintln(b, r) + fmt.Fprintln(b, forestTableName(r, duplicates, present[r])) walk(r, " ") } // Cycles (every member has a parent) still deserve printing. - var leftovers []string + var leftovers []tableIdentity for parent := range children { if !visited[parent] { leftovers = append(leftovers, parent) } } - sort.Strings(leftovers) + sort.Slice(leftovers, func(i, j int) bool { return lessIdentity(leftovers[i], leftovers[j]) }) for _, r := range leftovers { - fmt.Fprintln(b, r+" (cycle)") + fmt.Fprintln(b, forestTableName(r, duplicates, present[r])+" (cycle)") walk(r, " ") } } @@ -384,20 +542,57 @@ func writeForest(b *strings.Builder, s Schema) { func RenderMermaid(s Schema) string { var b strings.Builder b.WriteString("erDiagram\n") - edges := append([]Edge(nil), s.Edges...) - sort.Slice(edges, func(i, j int) bool { - if edges[i].ToTable != edges[j].ToTable { - return edges[i].ToTable < edges[j].ToTable + tables := append([]Table(nil), s.Tables...) + sort.Slice(tables, func(i, j int) bool { return lessIdentity(identityOf(tables[i]), identityOf(tables[j])) }) + edges := resolveEdges(tables, s.Edges) + duplicates := duplicateIdentityNames(tables, edges) + present := map[tableIdentity]bool{} + identities := map[tableIdentity]bool{} + forceAlias := map[tableIdentity]bool{} + for _, table := range tables { + id := identityOf(table) + present[id] = true + identities[id] = true + } + for _, edge := range edges { + identities[edge.From] = true + identities[edge.To] = true + forceAlias[edge.From] = forceAlias[edge.From] || (!edge.FromPresent && edge.From.Schema != "") + forceAlias[edge.To] = forceAlias[edge.To] || (!edge.ToPresent && edge.To.Schema != "") + } + orderedIdentities := make([]tableIdentity, 0, len(identities)) + for id := range identities { + orderedIdentities = append(orderedIdentities, id) + } + sort.Slice(orderedIdentities, func(i, j int) bool { return lessIdentity(orderedIdentities[i], orderedIdentities[j]) }) + + tableNames := map[tableIdentity]string{} + declarations := map[tableIdentity]string{} + usedNames := map[string]bool{} + for _, id := range orderedIdentities { + aliased := duplicates[id.Name] || !mermaidSafeEntityName(id.Name) || forceAlias[id] + name := id.Name + if aliased { + name = mermaidEntityID(id) + } + if usedNames[name] { + aliased = true + base := mermaidEntityID(id) + name = base + for suffix := 2; usedNames[name]; suffix++ { + name = fmt.Sprintf("%s_%d", base, suffix) + } + } + usedNames[name] = true + tableNames[id] = name + declarations[id] = name + if aliased { + declarations[id] += `["` + mermaidText(qualifiedTableName(id)) + `"]` } - return edges[i].FromTable < edges[j].FromTable - }) - for _, e := range edges { - fmt.Fprintf(&b, " %s ||--o{ %s : %s\n", e.ToTable, e.FromTable, e.FromColumn) } - tables := append([]Table(nil), s.Tables...) - sort.Slice(tables, func(i, j int) bool { return tables[i].Name < tables[j].Name }) for _, t := range tables { - fmt.Fprintf(&b, " %s {\n", t.Name) + id := identityOf(t) + fmt.Fprintf(&b, " %s {\n", declarations[id]) for _, c := range t.Columns { marker := "" switch { @@ -414,6 +609,14 @@ func RenderMermaid(s Schema) string { } b.WriteString(" }\n") } + for _, id := range orderedIdentities { + if !present[id] && declarations[id] != tableNames[id] { + fmt.Fprintf(&b, " %s\n", declarations[id]) + } + } + for _, e := range edges { + fmt.Fprintf(&b, " %s ||--o{ %s : %s\n", tableNames[e.To], tableNames[e.From], e.Edge.FromColumn) + } return b.String() } diff --git a/internal/erd/html.go b/internal/erd/html.go index 249c31b..f4e06ce 100644 --- a/internal/erd/html.go +++ b/internal/erd/html.go @@ -21,19 +21,19 @@ func RenderHTML(s Schema) string { ) tables := append([]Table(nil), s.Tables...) - sort.Slice(tables, func(i, j int) bool { return tables[i].Name < tables[j].Name }) + sort.Slice(tables, func(i, j int) bool { return lessIdentity(identityOf(tables[i]), identityOf(tables[j])) }) + edges := drawableEdges(resolveEdges(tables, s.Edges)) // Same layered layout as --layout row: parents left, children right. - depth := map[string]int{} - known := map[string]bool{} + depth := map[tableIdentity]int{} for _, t := range tables { - depth[t.Name], known[t.Name] = 0, true + depth[identityOf(t)] = 0 } for range tables { changed := false - for _, e := range s.Edges { - if known[e.FromTable] && known[e.ToTable] && depth[e.FromTable] < depth[e.ToTable]+1 { - depth[e.FromTable] = depth[e.ToTable] + 1 + for _, e := range edges { + if depth[e.From] < depth[e.To]+1 { + depth[e.From] = depth[e.To] + 1 changed = true } } @@ -50,18 +50,18 @@ func RenderHTML(s Schema) string { x, y, w, h float64 fkY map[string]float64 // FK column name → row center y } - boxes := map[string]*box{} + boxes := map[tableIdentity]*box{} colX := 0.0 for d := 0; d <= maxDepth; d++ { colW, y := 0.0, 0.0 var col []*Table for i := range tables { - if depth[tables[i].Name] == d { + if depth[identityOf(tables[i])] == d { col = append(col, &tables[i]) } } for _, t := range col { - w := float64(len(t.Schema+"."+t.Name))*charW + 2*padX + w := float64(len(qualifiedTableName(identityOf(*t))))*charW + 2*padX for _, c := range t.Columns { lw := float64(len(c.Name)+2+len(c.Type)+8)*charW + 2*padX if c.FKTarget != "" { @@ -79,16 +79,14 @@ func RenderHTML(s Schema) string { } b := &box{x: colX, y: y, w: w, h: titleH + float64(len(t.Columns))*rowH + ih + padY, fkY: map[string]float64{}} for i, c := range t.Columns { - if c.FKTarget != "" { - b.fkY[c.Name] = y + titleH + float64(i)*rowH + rowH/2 - } + b.fkY[c.Name] = y + titleH + float64(i)*rowH + rowH/2 } - boxes[t.Name] = b + boxes[identityOf(*t)] = b colW = maxFloat(colW, w) y += b.h + boxGap } for _, t := range col { - boxes[t.Name].w = colW + boxes[identityOf(*t)].w = colW } colX += colW + colGap } @@ -102,19 +100,9 @@ func RenderHTML(s Schema) string { var svg strings.Builder esc := html.EscapeString // Edges first, under the boxes. - edges := append([]Edge(nil), s.Edges...) - sort.Slice(edges, func(i, j int) bool { - if edges[i].ToTable != edges[j].ToTable { - return edges[i].ToTable < edges[j].ToTable - } - return edges[i].FromTable < edges[j].FromTable - }) for _, e := range edges { - child, parent := boxes[e.FromTable], boxes[e.ToTable] - if child == nil || parent == nil { - continue - } - y1, ok := child.fkY[e.FromColumn] + child, parent := boxes[e.From], boxes[e.To] + y1, ok := child.fkY[e.Edge.FromColumn] if !ok { continue } @@ -126,9 +114,9 @@ func RenderHTML(s Schema) string { } for i := range tables { t := &tables[i] - b := boxes[t.Name] + b := boxes[identityOf(*t)] fmt.Fprintf(&svg, ``+"\n", b.x, b.y, b.w, b.h) - fmt.Fprintf(&svg, `%s`+"\n", b.x+padX, b.y+20, esc(t.Schema+"."+t.Name)) + fmt.Fprintf(&svg, `%s`+"\n", b.x+padX, b.y+20, esc(qualifiedTableName(identityOf(*t)))) fmt.Fprintf(&svg, ``+"\n", b.x, b.y+titleH-2, b.x+b.w, b.y+titleH-2) for ci, c := range t.Columns { y := b.y + titleH + float64(ci)*rowH + 15 diff --git a/internal/erd/introspect.go b/internal/erd/introspect.go index a2a5bef..3fbc1b4 100644 --- a/internal/erd/introspect.go +++ b/internal/erd/introspect.go @@ -62,15 +62,15 @@ func Introspect(ctx context.Context, q Querier, schemaFilter string) (Schema, er Schema, Table, Column, Type string PK bool } - byTable := map[string]*Table{} - var order []string + byTable := map[tableIdentity]*Table{} + var order []tableIdentity for rows.Next() { var r colRow if err := rows.Scan(&r.Schema, &r.Table, &r.Column, &r.Type, &r.PK); err != nil { rows.Close() return s, err } - key := r.Schema + "." + r.Table + key := tableIdentity{Schema: r.Schema, Name: r.Table} t := byTable[key] if t == nil { t = &Table{Schema: r.Schema, Name: r.Table} @@ -89,16 +89,28 @@ func Introspect(ctx context.Context, q Querier, schemaFilter string) (Schema, er return s, fmt.Errorf("introspect foreign keys: %w", err) } defer rows.Close() + nameCounts := map[string]int{} + for id := range byTable { + nameCounts[id.Name]++ + } for rows.Next() { var fs, ft, fc, ts, tt, tc string if err := rows.Scan(&fs, &ft, &fc, &ts, &tt, &tc); err != nil { return s, err } - s.Edges = append(s.Edges, Edge{FromTable: ft, FromColumn: fc, ToTable: tt, ToColumn: tc}) - if t := byTable[fs+"."+ft]; t != nil { + s.Edges = append(s.Edges, Edge{ + FromSchema: fs, FromTable: ft, FromColumn: fc, + ToSchema: ts, ToTable: tt, ToColumn: tc, + }) + if t := byTable[tableIdentity{Schema: fs, Name: ft}]; t != nil { for i := range t.Columns { if t.Columns[i].Name == fc { - t.Columns[i].FKTarget = tt + "." + tc + targetID := tableIdentity{Schema: ts, Name: tt} + target := displayIdentifier(tt) + if nameCounts[tt] != 1 || byTable[targetID] == nil { + target = qualifiedTableName(targetID) + } + t.Columns[i].FKTarget = target + "." + displayIdentifier(tc) } } } @@ -134,9 +146,9 @@ ORDER BY 1, 2, 3` // introspectExtras fills indexes and the database header info; both are // best-effort decoration — an error leaves the structural diagram intact. func introspectExtras(ctx context.Context, q Querier, schemaFilter string, s *Schema) { - byKey := map[string]*Table{} + byKey := map[tableIdentity]*Table{} for i := range s.Tables { - byKey[s.Tables[i].Schema+"."+s.Tables[i].Name] = &s.Tables[i] + byKey[identityOf(s.Tables[i])] = &s.Tables[i] } if rows, err := q.Query(ctx, indexesSQL, schemaFilter); err == nil { for rows.Next() { @@ -145,7 +157,7 @@ func introspectExtras(ctx context.Context, q Querier, schemaFilter string, s *Sc if rows.Scan(&sch, &tbl, &name, &def, &uniq) != nil { break } - if t := byKey[sch+"."+tbl]; t != nil { + if t := byKey[tableIdentity{Schema: sch, Name: tbl}]; t != nil { t.Indexes = append(t.Indexes, Index{Name: name, Def: def, Unique: uniq}) } } diff --git a/internal/erd/introspect_integration_test.go b/internal/erd/introspect_integration_test.go new file mode 100644 index 0000000..89f5ac9 --- /dev/null +++ b/internal/erd/introspect_integration_test.go @@ -0,0 +1,123 @@ +package erd + +import ( + "context" + "fmt" + "os" + "strings" + "testing" + "time" + + "github.com/jackc/pgx/v5" +) + +func TestIntegrationIntrospectPreservesQualifiedTableIdentity(t *testing.T) { + dsn := os.Getenv("PGBOT_TEST_SUPERUSER_DSN") + if dsn == "" { + t.Skip("set PGBOT_TEST_SUPERUSER_DSN to run the disposable cross-schema ERD introspection test") + } + ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) + defer cancel() + conn, err := pgx.Connect(ctx, dsn) + if err != nil { + t.Fatalf("connect: %v", err) + } + t.Cleanup(func() { _ = conn.Close(context.Background()) }) + + prefix := fmt.Sprintf("pgbot_erd_%d", time.Now().UnixNano()) + collisionA := tableIdentity{Schema: prefix + ".part", Name: "users"} + collisionB := tableIdentity{Schema: prefix, Name: "part.users"} + sameName := tableIdentity{Schema: prefix + ` "客户"`, Name: "users"} + child := tableIdentity{Schema: prefix + " sales", Name: "orders"} + schemas := []string{collisionA.Schema, collisionB.Schema, sameName.Schema, child.Schema} + var createdSchemas []string + t.Cleanup(func() { + cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cleanupCancel() + for i := len(createdSchemas) - 1; i >= 0; i-- { + if _, err := conn.Exec(cleanupCtx, `DROP SCHEMA IF EXISTS `+pgx.Identifier{createdSchemas[i]}.Sanitize()+` CASCADE`); err != nil { + t.Errorf("drop schema %q: %v", createdSchemas[i], err) + } + } + }) + for _, schema := range schemas { + if _, err := conn.Exec(ctx, `CREATE SCHEMA `+pgx.Identifier{schema}.Sanitize()); err != nil { + t.Fatalf("create schema %q: %v", schema, err) + } + createdSchemas = append(createdSchemas, schema) + } + + createTable := func(id tableIdentity, body string) { + t.Helper() + if _, err := conn.Exec(ctx, `CREATE TABLE `+pgx.Identifier{id.Schema, id.Name}.Sanitize()+` (`+body+`)`); err != nil { + t.Fatalf("create table %s: %v", qualifiedTableName(id), err) + } + } + createTable(collisionA, `id bigint PRIMARY KEY, a_marker text`) + createTable(collisionB, `id bigint PRIMARY KEY, b_marker text`) + createTable(sameName, `id bigint PRIMARY KEY`) + createTable(child, `id bigint PRIMARY KEY, buyer_id bigint REFERENCES `+pgx.Identifier{collisionA.Schema, collisionA.Name}.Sanitize()+` (id)`) + for id, index := range map[tableIdentity]string{collisionA: "a_marker_idx", collisionB: "b_marker_idx"} { + column := strings.TrimSuffix(index, "_idx") + if _, err := conn.Exec(ctx, `CREATE INDEX `+pgx.Identifier{index}.Sanitize()+` ON `+pgx.Identifier{id.Schema, id.Name}.Sanitize()+` (`+pgx.Identifier{column}.Sanitize()+`)`); err != nil { + t.Fatalf("create index %q: %v", index, err) + } + } + + schema, err := Introspect(ctx, conn, "") + if err != nil { + t.Fatalf("introspect: %v", err) + } + byID := map[tableIdentity]Table{} + for _, table := range schema.Tables { + byID[identityOf(table)] = table + } + for id, marker := range map[tableIdentity]string{collisionA: "a_marker", collisionB: "b_marker"} { + table, ok := byID[id] + if !ok { + t.Errorf("missing table %s", qualifiedTableName(id)) + continue + } + found := false + for _, column := range table.Columns { + found = found || column.Name == marker + } + if !found { + t.Errorf("table %s lost its distinct %q column: %+v", qualifiedTableName(id), marker, table.Columns) + } + wantIndex := marker + "_idx" + found = false + for _, index := range table.Indexes { + found = found || index.Name == wantIndex + } + if !found { + t.Errorf("table %s lost its distinct %q index: %+v", qualifiedTableName(id), wantIndex, table.Indexes) + } + } + + var relationship *Edge + for i := range schema.Edges { + edge := &schema.Edges[i] + if edge.FromSchema == child.Schema && edge.FromTable == child.Name && edge.FromColumn == "buyer_id" { + relationship = edge + break + } + } + if relationship == nil { + t.Fatal("missing cross-schema foreign-key edge") + } + if relationship.ToSchema != collisionA.Schema || relationship.ToTable != collisionA.Name || relationship.ToColumn != "id" { + t.Errorf("foreign-key target lost schema identity: %+v", *relationship) + } + childTable := byID[child] + for _, column := range childTable.Columns { + if column.Name == "buyer_id" { + want := qualifiedTableName(collisionA) + ".id" + if column.FKTarget != want { + t.Errorf("FK label = %q, want %q", column.FKTarget, want) + } + return + } + } + t.Fatal("child table lost buyer_id column") +} diff --git a/internal/erd/parent_identity_regression_test.go b/internal/erd/parent_identity_regression_test.go new file mode 100644 index 0000000..f65a844 --- /dev/null +++ b/internal/erd/parent_identity_regression_test.go @@ -0,0 +1,17 @@ +package erd + +import ( + "strings" + "testing" +) + +func TestTableIdentityParentRegression(t *testing.T) { + s := Schema{Tables: []Table{{Schema: "a.b", Name: "users", Columns: []Column{{Name: "id", Type: "bigint"}}}, {Schema: "a", Name: "b.users", Columns: []Column{{Name: "id", Type: "bigint"}}}}} + got := RenderASCII(s, false) + if strings.Contains(got, "a.b.users") { + t.Fatalf("distinct table coordinates collapse into identical labels: %s", got) + } + if !strings.Contains(got, `"a.b".users`) || !strings.Contains(got, `a."b.users"`) { + t.Fatalf("qualified labels missing: %s", got) + } +} diff --git a/internal/erd/row.go b/internal/erd/row.go index 733b41a..5a5e451 100644 --- a/internal/erd/row.go +++ b/internal/erd/row.go @@ -16,26 +16,20 @@ func RenderASCIIRow(s Schema) string { } tables := append([]Table(nil), s.Tables...) - sort.Slice(tables, func(i, j int) bool { return tables[i].Name < tables[j].Name }) - byName := map[string]*Table{} - for i := range tables { - byName[tables[i].Name] = &tables[i] - } + sort.Slice(tables, func(i, j int) bool { return lessIdentity(identityOf(tables[i]), identityOf(tables[j])) }) + edges := drawableEdges(resolveEdges(tables, s.Edges)) // Depth: roots (no FK out, or FK to unknown) at 0; a child sits one right // of its deepest parent. Iterate to fixpoint; cycles keep their first depth. - depth := map[string]int{} + depth := map[tableIdentity]int{} for _, t := range tables { - depth[t.Name] = 0 + depth[identityOf(t)] = 0 } for iter := 0; iter < len(tables); iter++ { changed := false - for _, e := range s.Edges { - if _, ok := byName[e.FromTable]; !ok { - continue - } - if d, ok := depth[e.ToTable]; ok && depth[e.FromTable] < d+1 { - depth[e.FromTable] = d + 1 + for _, e := range edges { + if depth[e.From] < depth[e.To]+1 { + depth[e.From] = depth[e.To] + 1 changed = true } } @@ -50,25 +44,24 @@ func RenderASCIIRow(s Schema) string { // Columns: render each box, compute per-column width and stacked heights. type placed struct { - lines []string - x0, x1, y0 int // global coordinates; y0 = title row - fkRowByColumn map[string]int + lines []string + x0, x1, y0 int // global coordinates; y0 = title row + rowByColumn map[string]int } cols := make([][]*placed, maxDepth+1) - pl := map[string]*placed{} + pl := map[tableIdentity]*placed{} for i := range tables { t := &tables[i] var b strings.Builder writeTableBox(&b, *t) p := &placed{lines: strings.Split(strings.TrimRight(b.String(), "\n"), "\n"), - fkRowByColumn: map[string]int{}} + rowByColumn: map[string]int{}} for ci, c := range t.Columns { - if c.FKTarget != "" { - p.fkRowByColumn[c.Name] = ci + 1 // relative to box top - } + p.rowByColumn[c.Name] = ci + 1 // relative to box top } - cols[depth[t.Name]] = append(cols[depth[t.Name]], p) - pl[t.Name] = p + id := identityOf(*t) + cols[depth[id]] = append(cols[depth[id]], p) + pl[id] = p } // Gutter lanes: one vertical track per edge in the gutter left of the @@ -76,11 +69,8 @@ func RenderASCIIRow(s Schema) string { const lanesPerGutter = 4 gutterW := make([]int, maxDepth+1) // gutter g sits left of column g (g>=1) edgesInGutter := make([]int, maxDepth+2) - for _, e := range s.Edges { - if pl[e.FromTable] == nil || pl[e.ToTable] == nil { - continue - } - edgesInGutter[depth[e.FromTable]]++ + for _, e := range edges { + edgesInGutter[depth[e.From]]++ } for g := 1; g <= maxDepth; g++ { gutterW[g] = 4 + 2*minInt(edgesInGutter[g], lanesPerGutter) @@ -135,25 +125,15 @@ func RenderASCIIRow(s Schema) string { // `<` into the parent's right border. Only adjacent-column edges get a // line; longer spans (and over-cap fan-ins) keep their textual FK marker. laneUsed := map[int]int{} // gutter → lanes taken - edges := append([]Edge(nil), s.Edges...) - sort.Slice(edges, func(i, j int) bool { - if edges[i].ToTable != edges[j].ToTable { - return edges[i].ToTable < edges[j].ToTable - } - return edges[i].FromTable < edges[j].FromTable - }) for _, e := range edges { - child, parent := pl[e.FromTable], pl[e.ToTable] - if child == nil || parent == nil { - continue - } - g := depth[e.FromTable] - if depth[e.ToTable] != g-1 || laneUsed[g] >= lanesPerGutter { + child, parent := pl[e.From], pl[e.To] + g := depth[e.From] + if depth[e.To] != g-1 || laneUsed[g] >= lanesPerGutter { continue } lane := laneUsed[g] laneUsed[g]++ - fkRel, ok := child.fkRowByColumn[e.FromColumn] + fkRel, ok := child.rowByColumn[e.Edge.FromColumn] if !ok { continue }