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 }