diff --git a/cmd/relspec/convert.go b/cmd/relspec/convert.go index 5efd0e5..b09f406 100644 --- a/cmd/relspec/convert.go +++ b/cmd/relspec/convert.go @@ -229,6 +229,7 @@ func runConvert(cmd *cobra.Command, args []string) error { if err != nil { return fmt.Errorf("failed to read source: %w", err) } + finalizeCommentedRefs(db, stderrWarn) fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", db.Name) fmt.Fprintf(os.Stderr, " Found: %d schema(s)\n", len(db.Schemas)) @@ -261,6 +262,18 @@ func runConvert(cmd *cobra.Command, args []string) error { return nil } +// finalizeCommentedRefs resolves DBML `// Ref:` comments against the fully +// loaded model and warns about refs whose target is not loaded. +func finalizeCommentedRefs(db *models.Database, warn func(string)) { + for _, w := range dbml.ResolveCommentedRefs(db, true) { + warn(w) + } +} + +func stderrWarn(msg string) { + fmt.Fprintf(os.Stderr, " ⚠ %s\n", msg) +} + func readDatabaseListForConvert(dbType string, files []string) (*models.Database, error) { if len(files) == 0 { return nil, fmt.Errorf("file list is empty") diff --git a/cmd/relspec/job.go b/cmd/relspec/job.go index cbacd6b..20c232e 100644 --- a/cmd/relspec/job.go +++ b/cmd/relspec/job.go @@ -578,6 +578,7 @@ func runMergeJob(rj *resolvedJob, lg *jobLogger) error { lg.logf("merging: %s", inputLabel(ri)) merge.MergeDatabases(base, db, opts) } + finalizeCommentedRefs(base, func(w string) { lg.logf("warning: %s", w) }) base.UpdateDate() return writeJobOutput(rj, base, lg) } @@ -814,6 +815,7 @@ func readJobInputs(rj *resolvedJob, lg *jobLogger) (*models.Database, error) { if base == nil { return nil, fmt.Errorf("no inputs produced a database") } + finalizeCommentedRefs(base, func(w string) { lg.logf("warning: %s", w) }) return base, nil } diff --git a/cmd/relspec/merge.go b/cmd/relspec/merge.go index 653f338..ccac526 100644 --- a/cmd/relspec/merge.go +++ b/cmd/relspec/merge.go @@ -248,6 +248,7 @@ func runMerge(cmd *cobra.Command, args []string) error { } result := merge.MergeDatabases(targetDB, sourceDB, opts) + finalizeCommentedRefs(targetDB, stderrWarn) // Update timestamp targetDB.UpdateDate() diff --git a/cmd/relspec/templ.go b/cmd/relspec/templ.go index b2ad631..66922af 100644 --- a/cmd/relspec/templ.go +++ b/cmd/relspec/templ.go @@ -114,6 +114,7 @@ func runTempl(cmd *cobra.Command, args []string) error { if err != nil { return fmt.Errorf("failed to read source: %w", err) } + finalizeCommentedRefs(db, stderrWarn) // Print database stats schemaCount := len(db.Schemas) diff --git a/pkg/merge/merge.go b/pkg/merge/merge.go index de0c84f..5029f34 100644 --- a/pkg/merge/merge.go +++ b/pkg/merge/merge.go @@ -91,6 +91,45 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions) if !opts.SkipDomains { r.mergeDomains(target, source) } + + mergeDatabaseMetadata(target, source) +} + +// mergeDatabaseMetadata adds missing metadata keys and unions []string values, +// so per-file reader state (e.g. pending DBML commented refs) survives a merge. +func mergeDatabaseMetadata(target, source *models.Database) { + if len(source.Metadata) == 0 { + return + } + if target.Metadata == nil { + target.Metadata = make(map[string]any, len(source.Metadata)) + } + for key, srcVal := range source.Metadata { + tgtVal, exists := target.Metadata[key] + if !exists { + if list, ok := srcVal.([]string); ok { + srcVal = append([]string(nil), list...) + } + target.Metadata[key] = srcVal + continue + } + tgtList, tgtOK := tgtVal.([]string) + srcList, srcOK := srcVal.([]string) + if !tgtOK || !srcOK { + continue + } + seen := make(map[string]bool, len(tgtList)) + for _, v := range tgtList { + seen[v] = true + } + for _, v := range srcList { + if !seen[v] { + tgtList = append(tgtList, v) + seen[v] = true + } + } + target.Metadata[key] = tgtList + } } func (r *MergeResult) mergeSchemaContents(target, source *models.Schema, opts *MergeOptions) { diff --git a/pkg/merge/merge_test.go b/pkg/merge/merge_test.go index 7b29805..bacf89a 100644 --- a/pkg/merge/merge_test.go +++ b/pkg/merge/merge_test.go @@ -721,3 +721,32 @@ func TestComplexMerge(t *testing.T) { t.Error("Expected ukey_users_guid constraint to exist") } } + +func TestMergeDatabases_Metadata(t *testing.T) { + target := &models.Database{Metadata: map[string]any{ + "refs": []string{"a", "b"}, + "name": "target", + }} + source := &models.Database{Metadata: map[string]any{ + "refs": []string{"b", "c"}, + "name": "source", + "extra": []string{"x"}, + }} + + MergeDatabases(target, source, nil) + + if got := target.Metadata["refs"].([]string); strings.Join(got, ",") != "a,b,c" { + t.Errorf("refs = %v, want [a b c]", got) + } + if got := target.Metadata["name"]; got != "target" { + t.Errorf("name = %v, want target (existing scalar keys are kept)", got) + } + extra := target.Metadata["extra"].([]string) + if strings.Join(extra, ",") != "x" { + t.Errorf("extra = %v, want [x]", extra) + } + source.Metadata["extra"].([]string)[0] = "changed" + if extra[0] != "x" { + t.Error("copied slice must not alias the source") + } +} diff --git a/pkg/readers/dbml/README.md b/pkg/readers/dbml/README.md index a35d8d7..ef0daa7 100644 --- a/pkg/readers/dbml/README.md +++ b/pkg/readers/dbml/README.md @@ -90,11 +90,28 @@ Ref: posts.user_id > users.id [delete: cascade] - Default values (`default`) - Inline references (`ref`) - Standalone `Ref` blocks +- Commented cross-file refs (`// Ref:` — see below) - Indexes and composite indexes - Table notes and column notes - Enums - Dialect directives (`@postgres:` / `@sqlite:` — see below) +## Commented cross-file refs + +`// Ref:` / `// ref:` lines (ignored by dbdiagram) become FKs + relationships once both ends are loaded. + +| Rule | Behaviour | +|---|---| +| When | Single file / directory: end of read. `--from-list`, `merge`, jobs: after all inputs are combined | +| Match | `schema.table.column` on both sides, case-insensitive | +| Operators | `>`, `<`, `-` (parsed like `Ref:`) | +| Duplicate of an FK on the same columns | Skipped silently | +| Target not loaded | Kept pending; skipped with a warning on the final pass | +| Column type mismatch | Warning, FK still created (`serial`≈`integer`, `bigserial`≈`bigint`) | +| Pending state | `Database.Metadata["dbml.commented_refs"]` (`[]string`) | + +API: `dbml.ResolveCommentedRefs(db, final bool) []string` returns warnings. + ## Dialect directives Lines of the form `@[()]: ` embed database-specific diff --git a/pkg/readers/dbml/commented_refs.go b/pkg/readers/dbml/commented_refs.go new file mode 100644 index 0000000..31cf59b --- /dev/null +++ b/pkg/readers/dbml/commented_refs.go @@ -0,0 +1,218 @@ +package dbml + +import ( + "fmt" + "regexp" + "strings" + + "git.warky.dev/wdevs/relspecgo/pkg/models" +) + +// CommentedRefsMetadataKey is the Database.Metadata key holding commented +// `// Ref:` lines not yet resolved against the model ([]string). +const CommentedRefsMetadataKey = "dbml.commented_refs" + +// commentedRefRegex matches `// Ref: ...` and `// ref: ...`. +var commentedRefRegex = regexp.MustCompile(`^//\s*[Rr]ef\s*:\s*(.+)$`) + +// commentedRef returns the ref body of a trimmed `// Ref:` comment line. +func commentedRef(line string) (string, bool) { + m := commentedRefRegex.FindStringSubmatch(line) + if m == nil { + return "", false + } + ref := strings.TrimSpace(m[1]) + return ref, ref != "" +} + +// PendingCommentedRefs returns the commented refs not yet resolved. +func PendingCommentedRefs(db *models.Database) []string { + if db == nil || db.Metadata == nil { + return nil + } + refs, _ := db.Metadata[CommentedRefsMetadataKey].([]string) + return refs +} + +func setPendingCommentedRefs(db *models.Database, refs []string) { + if len(refs) == 0 { + if db.Metadata != nil { + delete(db.Metadata, CommentedRefsMetadataKey) + } + return + } + if db.Metadata == nil { + db.Metadata = make(map[string]any) + } + db.Metadata[CommentedRefsMetadataKey] = refs +} + +// addPendingCommentedRef queues a ref, skipping exact repeats. +func addPendingCommentedRef(db *models.Database, ref string) { + refs := PendingCommentedRefs(db) + for _, existing := range refs { + if existing == ref { + return + } + } + setPendingCommentedRefs(db, append(refs, ref)) +} + +// ResolveCommentedRefs turns pending commented refs into foreign keys and +// relationships when both ends (schema.table.column) exist in db. A ref that +// matches an existing FK on the same columns is dropped as a duplicate. +// Unmatched refs stay pending unless final is set, in which case they are +// dropped with a warning. Returns human-readable warnings. +func ResolveCommentedRefs(db *models.Database, final bool) []string { + refs := PendingCommentedRefs(db) + if len(refs) == 0 { + return nil + } + + var warnings []string + var pending []string + parser := &Reader{} + + for _, ref := range refs { + fk := parser.parseRef(ref) + if fk == nil || len(fk.Columns) == 0 || len(fk.Columns) != len(fk.ReferencedColumns) { + warnings = append(warnings, fmt.Sprintf("skipping commented ref %q: cannot parse", ref)) + continue + } + + srcTable, srcCols, srcMissing := lookupColumns(db, fk.Schema, fk.Table, fk.Columns) + dstTable, dstCols, dstMissing := lookupColumns(db, fk.ReferencedSchema, fk.ReferencedTable, fk.ReferencedColumns) + if srcMissing != "" || dstMissing != "" { + if final { + missing := srcMissing + if missing == "" { + missing = dstMissing + } + warnings = append(warnings, fmt.Sprintf("skipping commented ref %q: %s not found", ref, missing)) + } else { + pending = append(pending, ref) + } + continue + } + + fk.Schema, fk.Table, fk.Columns = srcTable.Schema, srcTable.Name, columnNames(srcCols) + fk.ReferencedSchema, fk.ReferencedTable, fk.ReferencedColumns = dstTable.Schema, dstTable.Name, columnNames(dstCols) + + if hasFKOnColumns(srcTable, fk.Columns) { + continue // already declared by an uncommented Ref or inline ref + } + if _, taken := srcTable.Constraints[fk.Name]; taken { + warnings = append(warnings, fmt.Sprintf("skipping commented ref %q: constraint %s already exists", ref, fk.Name)) + continue + } + + for i := range srcCols { + if !compatibleFKTypes(srcCols[i].Type, dstCols[i].Type) { + warnings = append(warnings, fmt.Sprintf("commented ref %q: type mismatch %s.%s.%s (%s) -> %s.%s.%s (%s)", + ref, srcTable.Schema, srcTable.Name, srcCols[i].Name, srcCols[i].Type, + dstTable.Schema, dstTable.Name, dstCols[i].Name, dstCols[i].Type)) + } + } + + if srcTable.Constraints == nil { + srcTable.Constraints = make(map[string]*models.Constraint) + } + srcTable.Constraints[fk.Name] = fk + addFKRelationship(srcTable, fk) + } + + setPendingCommentedRefs(db, pending) + return warnings +} + +// lookupColumns finds a table and its columns, case-insensitively. missing +// names the first object not found, or is empty. +func lookupColumns(db *models.Database, schemaName, tableName string, cols []string) (*models.Table, []*models.Column, string) { + qualified := schemaName + "." + tableName + var table *models.Table + for _, schema := range db.Schemas { + if !strings.EqualFold(schema.Name, schemaName) { + continue + } + for _, t := range schema.Tables { + if strings.EqualFold(t.Name, tableName) { + table = t + break + } + } + } + if table == nil { + return nil, nil, "table " + qualified + } + + found := make([]*models.Column, 0, len(cols)) + for _, name := range cols { + col := table.Columns[name] + if col == nil { + for _, c := range table.Columns { + if strings.EqualFold(c.Name, name) { + col = c + break + } + } + } + if col == nil { + return nil, nil, "column " + qualified + "." + name + } + found = append(found, col) + } + return table, found, "" +} + +func columnNames(cols []*models.Column) []string { + names := make([]string, len(cols)) + for i, c := range cols { + names[i] = c.Name + } + return names +} + +// hasFKOnColumns reports whether table already has a foreign key over cols. +func hasFKOnColumns(table *models.Table, cols []string) bool { + for _, c := range table.Constraints { + if c.Type != models.ForeignKeyConstraint || len(c.Columns) != len(cols) { + continue + } + same := true + for i := range cols { + if !strings.EqualFold(c.Columns[i], cols[i]) { + same = false + break + } + } + if same { + return true + } + } + return false +} + +// fkTypeAliases maps serial and alias spellings to their storage type. +var fkTypeAliases = map[string]string{ + "smallserial": "smallint", "serial2": "smallint", "int2": "smallint", + "serial": "integer", "serial4": "integer", "int": "integer", "int4": "integer", + "bigserial": "bigint", "serial8": "bigint", "int8": "bigint", +} + +// compatibleFKTypes compares column types ignoring case, length and serial +// vs. integer spelling. Unknown (empty) types are treated as compatible. +func compatibleFKTypes(a, b string) bool { + na, nb := normalizeFKType(a), normalizeFKType(b) + return na == "" || nb == "" || na == nb +} + +func normalizeFKType(t string) string { + t = strings.ToLower(strings.TrimSpace(t)) + if i := strings.Index(t, "("); i >= 0 { + t = strings.TrimSpace(t[:i]) + } + if alias, ok := fkTypeAliases[t]; ok { + return alias + } + return t +} diff --git a/pkg/readers/dbml/commented_refs_test.go b/pkg/readers/dbml/commented_refs_test.go new file mode 100644 index 0000000..d6ab2c7 --- /dev/null +++ b/pkg/readers/dbml/commented_refs_test.go @@ -0,0 +1,260 @@ +package dbml + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "git.warky.dev/wdevs/relspecgo/pkg/models" + "git.warky.dev/wdevs/relspecgo/pkg/readers" +) + +func readDBMLString(t *testing.T, content string) *models.Database { + t.Helper() + path := filepath.Join(t.TempDir(), "in.dbml") + if err := os.WriteFile(path, []byte(content), 0o644); err != nil { + t.Fatalf("write fixture: %v", err) + } + db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase() + if err != nil { + t.Fatalf("ReadDatabase() error = %v", err) + } + return db +} + +func findTable(db *models.Database, schema, table string) *models.Table { + for _, s := range db.Schemas { + if s.Name != schema { + continue + } + for _, t := range s.Tables { + if t.Name == table { + return t + } + } + } + return nil +} + +func fkOn(table *models.Table, col string) *models.Constraint { + for _, c := range table.Constraints { + if c.Type == models.ForeignKeyConstraint && len(c.Columns) == 1 && c.Columns[0] == col { + return c + } + } + return nil +} + +func relFor(table *models.Table, fkName string) *models.Relationship { + for _, r := range table.Relationships { + if r.ForeignKey == fkName { + return r + } + } + return nil +} + +func TestCommentedRef(t *testing.T) { + tests := []struct { + line string + want string + ok bool + }{ + {`// Ref: a.b.c > d.e.f`, `a.b.c > d.e.f`, true}, + {`// ref: a.b.c - d.e.f`, `a.b.c - d.e.f`, true}, + {`//Ref:a.b.c > d.e.f`, `a.b.c > d.e.f`, true}, + {`// Reference notes`, "", false}, + {`// see Ref: a.b.c > d.e.f`, "", false}, + {`// Ref:`, "", false}, + } + for _, tt := range tests { + got, ok := commentedRef(tt.line) + if ok != tt.ok || got != tt.want { + t.Errorf("commentedRef(%q) = (%q, %v), want (%q, %v)", tt.line, got, ok, tt.want, tt.ok) + } + } +} + +// A commented ref whose tables are in the same file resolves on read. +func TestReader_CommentedRefSameFile(t *testing.T) { + db := readDBMLString(t, `Table "org"."department" { + "id_department" bigserial [pk] +} +Table "entity"."employee" { + "id_employee" bigserial [pk] + "rid_department" bigint +} +// Ref: "entity"."employee"."rid_department" > "org"."department"."id_department" [delete: restrict, update: restrict] +`) + emp := findTable(db, "entity", "employee") + fk := fkOn(emp, "rid_department") + if fk == nil { + t.Fatal("expected FK on rid_department") + } + if fk.ReferencedSchema != "org" || fk.ReferencedTable != "department" || fk.ReferencedColumns[0] != "id_department" { + t.Errorf("FK target = %s.%s.%v", fk.ReferencedSchema, fk.ReferencedTable, fk.ReferencedColumns) + } + if fk.OnDelete != "restrict" || fk.OnUpdate != "restrict" { + t.Errorf("FK actions = %q/%q, want restrict/restrict", fk.OnDelete, fk.OnUpdate) + } + if relFor(emp, fk.Name) == nil { + t.Error("expected relationship for FK") + } + if refs := PendingCommentedRefs(db); len(refs) != 0 { + t.Errorf("pending = %v, want none", refs) + } +} + +// A cross-file commented ref stays pending after a single-file read. +func TestReader_CommentedRefCrossFilePending(t *testing.T) { + db := readDBMLString(t, `Table "entity"."employee" { + "id_employee" bigserial [pk] + "rid_department" bigint +} +// Ref: "entity"."employee"."rid_department" > "org"."department"."id_department" +`) + if fk := fkOn(findTable(db, "entity", "employee"), "rid_department"); fk != nil { + t.Fatal("FK must not resolve without the target table") + } + if refs := PendingCommentedRefs(db); len(refs) != 1 { + t.Fatalf("pending = %v, want 1", refs) + } +} + +// Directory reads resolve commented refs after all files are merged. +func TestReader_CommentedRefDirectory(t *testing.T) { + db, err := NewReader(&readers.ReaderOptions{ + FilePath: filepath.Join("..", "..", "..", "tests", "assets", "dbml", "multifile"), + }).ReadDatabase() + if err != nil { + t.Fatalf("ReadDatabase() error = %v", err) + } + fk := fkOn(findTable(db, "public", "posts"), "user_id") + if fk == nil { + t.Fatal("expected FK posts.user_id from 9_refs.dbml commented ref") + } + if fk.ReferencedTable != "users" || fk.OnDelete != "CASCADE" { + t.Errorf("FK = %s ondelete %s, want users ondelete CASCADE", fk.ReferencedTable, fk.OnDelete) + } +} + +func crossFileDB(t *testing.T, refs ...string) *models.Database { + t.Helper() + db := readDBMLString(t, `Table "org"."department" { + "id_department" bigserial [pk] +} +Table "entity"."employee" { + "id_employee" bigserial [pk] + "rid_department" bigint + "rid_manager" integer + "rid_team" bigint +} +`) + setPendingCommentedRefs(db, refs) + return db +} + +func TestResolveCommentedRefs(t *testing.T) { + t.Run("lowercase ref and one-to-one", func(t *testing.T) { + db := crossFileDB(t, + `"entity"."employee"."rid_department" > "org"."department"."id_department"`, + `entity.employee.rid_team - org.department.id_department`, + ) + if w := ResolveCommentedRefs(db, true); len(w) != 0 { + t.Errorf("warnings = %v, want none", w) + } + emp := findTable(db, "entity", "employee") + for _, col := range []string{"rid_department", "rid_team"} { + fk := fkOn(emp, col) + if fk == nil { + t.Fatalf("expected FK on %s", col) + } + if relFor(emp, fk.Name) == nil { + t.Errorf("expected relationship for %s", fk.Name) + } + } + if len(emp.Relationships) != 2 { + t.Errorf("relationships = %d, want 2 (same target must not overwrite)", len(emp.Relationships)) + } + }) + + t.Run("type mismatch warns but resolves", func(t *testing.T) { + db := crossFileDB(t, `entity.employee.rid_manager > org.department.id_department`) + w := ResolveCommentedRefs(db, true) + if len(w) != 1 || !strings.Contains(w[0], "type mismatch") { + t.Errorf("warnings = %v, want one type mismatch", w) + } + if fkOn(findTable(db, "entity", "employee"), "rid_manager") == nil { + t.Error("expected FK on rid_manager") + } + }) + + t.Run("missing target stays pending until final", func(t *testing.T) { + db := crossFileDB(t, `entity.employee.rid_team > hr.team.id_team`) + if w := ResolveCommentedRefs(db, false); len(w) != 0 { + t.Errorf("non-final warnings = %v, want none", w) + } + if len(PendingCommentedRefs(db)) != 1 { + t.Fatal("ref should stay pending") + } + w := ResolveCommentedRefs(db, true) + if len(w) != 1 || !strings.Contains(w[0], "table hr.team not found") { + t.Errorf("final warnings = %v, want missing table", w) + } + if len(PendingCommentedRefs(db)) != 0 { + t.Error("final pass must clear pending refs") + } + }) + + t.Run("missing column", func(t *testing.T) { + db := crossFileDB(t, `entity.employee.rid_nope > org.department.id_department`) + w := ResolveCommentedRefs(db, true) + if len(w) != 1 || !strings.Contains(w[0], "column entity.employee.rid_nope not found") { + t.Errorf("warnings = %v, want missing column", w) + } + }) + + t.Run("deduplicates against uncommented ref", func(t *testing.T) { + db := readDBMLString(t, `Table "org"."department" { + "id_department" bigserial [pk] +} +Table "entity"."employee" { + "id_employee" bigserial [pk] + "rid_department" bigint +} +Ref: "entity"."employee"."rid_department" > "org"."department"."id_department" +// Ref: "entity"."employee"."rid_department" > "org"."department"."id_department" +`) + emp := findTable(db, "entity", "employee") + count := 0 + for _, c := range emp.Constraints { + if c.Type == models.ForeignKeyConstraint { + count++ + } + } + if count != 1 || len(emp.Relationships) != 1 { + t.Errorf("FKs = %d, relationships = %d, want 1 and 1", count, len(emp.Relationships)) + } + }) +} + +func TestCompatibleFKTypes(t *testing.T) { + tests := []struct { + a, b string + want bool + }{ + {"bigint", "bigserial", true}, + {"integer", "serial", true}, + {"INT4", "integer", true}, + {"varchar(10)", "varchar(20)", true}, + {"bigint", "serial", false}, + {"uuid", "bigint", false}, + {"", "bigint", true}, + } + for _, tt := range tests { + if got := compatibleFKTypes(tt.a, tt.b); got != tt.want { + t.Errorf("compatibleFKTypes(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.want) + } + } +} diff --git a/pkg/readers/dbml/reader.go b/pkg/readers/dbml/reader.go index 664399c..c2a1888 100644 --- a/pkg/readers/dbml/reader.go +++ b/pkg/readers/dbml/reader.go @@ -48,7 +48,22 @@ func (r *Reader) ReadDatabase() (*models.Database, error) { return nil, fmt.Errorf("failed to read file: %w", err) } - return r.parseDBML(string(content)) + db, err := r.parseDBML(string(content)) + if err != nil { + return nil, err + } + r.resolveCommentedRefs(db) + return db, nil +} + +// resolveCommentedRefs resolves the commented refs whose tables are loaded. +// Unmatched refs stay pending for a later pass over a combined model. +func (r *Reader) resolveCommentedRefs(db *models.Database) { + for _, w := range ResolveCommentedRefs(db, false) { + if r.options.Progress != nil { + r.options.Progress("warning: " + w) + } + } } // ReadSchema reads and parses DBML input, returning a Schema model @@ -125,6 +140,7 @@ func (r *Reader) readDirectoryDBML(dirPath string) (*models.Database, error) { } } + r.resolveCommentedRefs(db) return db, nil } @@ -440,6 +456,10 @@ func mergeDatabase(baseDB, fileDB *models.Database) { // Merge domains baseDB.Domains = append(baseDB.Domains, fileDB.Domains...) + for _, ref := range PendingCommentedRefs(fileDB) { + addPendingCommentedRef(baseDB, ref) + } + // Use first non-empty description if baseDB.Description == "" && fileDB.Description != "" { baseDB.Description = fileDB.Description @@ -493,8 +513,12 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) { continue } - // Skip empty lines and comments + // Skip empty lines and comments. A commented `// Ref:` is kept as a + // pending cross-file ref, resolved once every file is loaded. if line == "" || strings.HasPrefix(line, "//") { + if ref, ok := commentedRef(line); ok { + addPendingCommentedRef(db, ref) + } continue } @@ -581,7 +605,13 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) { index := r.parseIndex(line, currentTable.Name, currentSchema) if index != nil { - currentTable.Indexes[index.Name] = index + // Keep a reused name under a unique map key so the duplicate is + // not silently dropped; the inspector reports it. + key := index.Name + for n := 2; currentTable.Indexes[key] != nil; n++ { + key = fmt.Sprintf("%s#%d", index.Name, n) + } + currentTable.Indexes[key] = index lastIndex = index } continue @@ -659,20 +689,12 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) { // for DBML refs so diffing equivalent schemas compares the same model. for _, schema := range schemaMap { for _, table := range schema.Tables { - for _, constraint := range table.Constraints { + for _, name := range sortedConstraintNames(table.Constraints) { + constraint := table.Constraints[name] if constraint.Type != models.ForeignKeyConstraint { continue } - name := fmt.Sprintf("%s_to_%s", table.Name, constraint.ReferencedTable) - relationship := models.InitRelationship(name, models.OneToMany) - relationship.FromTable = table.Name - relationship.FromSchema = table.Schema - relationship.FromColumns = append([]string(nil), constraint.Columns...) - relationship.ToTable = constraint.ReferencedTable - relationship.ToSchema = constraint.ReferencedSchema - relationship.ToColumns = append([]string(nil), constraint.ReferencedColumns...) - relationship.ForeignKey = constraint.Name - table.Relationships[name] = relationship + addFKRelationship(table, constraint) } } } @@ -685,6 +707,39 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) { return db, nil } +// sortedConstraintNames returns constraint keys in sorted order so derived +// relationship names do not depend on map iteration order. +func sortedConstraintNames(constraints map[string]*models.Constraint) []string { + names := make([]string, 0, len(constraints)) + for name := range constraints { + names = append(names, name) + } + sort.Strings(names) + return names +} + +// addFKRelationship derives the relationship for a foreign key, matching how +// the PostgreSQL reader models FKs. +func addFKRelationship(table *models.Table, constraint *models.Constraint) { + name := fmt.Sprintf("%s_to_%s", table.Name, constraint.ReferencedTable) + if existing, taken := table.Relationships[name]; taken && existing.ForeignKey != constraint.Name { + // A second FK to the same table must not overwrite the first. + name = fmt.Sprintf("%s_%s", name, strings.Join(constraint.Columns, "_")) + } + relationship := models.InitRelationship(name, models.OneToMany) + relationship.FromTable = table.Name + relationship.FromSchema = table.Schema + relationship.FromColumns = append([]string(nil), constraint.Columns...) + relationship.ToTable = constraint.ReferencedTable + relationship.ToSchema = constraint.ReferencedSchema + relationship.ToColumns = append([]string(nil), constraint.ReferencedColumns...) + relationship.ForeignKey = constraint.Name + if table.Relationships == nil { + table.Relationships = make(map[string]*models.Relationship) + } + table.Relationships[name] = relationship +} + // setTableNote preserves multiple table notes. The first maps to Description // and the second to Comment, matching the model fields used by code writers. func setTableNote(table *models.Table, note string) { diff --git a/pkg/readers/dbml/reader_test.go b/pkg/readers/dbml/reader_test.go index be98ae4..eb60b7a 100644 --- a/pkg/readers/dbml/reader_test.go +++ b/pkg/readers/dbml/reader_test.go @@ -1085,3 +1085,39 @@ func TestReader_MultilineTableNote(t *testing.T) { t.Errorf("column note = %q, want %q", got, want) } } + +// A name reused inside one Indexes block must not silently drop the earlier +// index; both are kept so the inspector can report the duplicate. +func TestReader_DuplicateIndexNameInTableKept(t *testing.T) { + dbmlContent := `Table individual_actor { + id bigint [pk] + rid_actor bigint + kind text + + Indexes { + (rid_actor) [name: 'idx_actor', unique] + (rid_actor, kind) [name: 'idx_actor'] + } +} +` + dir := t.TempDir() + path := filepath.Join(dir, "dup_index.dbml") + if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil { + t.Fatalf("failed to write fixture: %v", err) + } + + db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase() + if err != nil { + t.Fatalf("ReadDatabase() error = %v", err) + } + + table := db.Schemas[0].Tables[0] + if len(table.Indexes) != 2 { + t.Fatalf("expected 2 indexes, got %d: %v", len(table.Indexes), table.Indexes) + } + for key, idx := range table.Indexes { + if idx.Name != "idx_actor" { + t.Errorf("index %q: Name = %q, want idx_actor", key, idx.Name) + } + } +} diff --git a/pkg/writers/pgsql/writer.go b/pkg/writers/pgsql/writer.go index 36f3d25..87c4aba 100644 --- a/pkg/writers/pgsql/writer.go +++ b/pkg/writers/pgsql/writer.go @@ -1259,13 +1259,18 @@ func (w *Writer) writeForeignKeys(schema *models.Schema) error { fkName = fmt.Sprintf("fk_%s_%s", table.SQLName(), rel.ToTable) } - // Find the foreign key constraint that matches this relationship + // Find the foreign key constraint that matches this relationship. + // Prefer the exact name: with several FKs to the same table, a + // referenced-table match alone picks an arbitrary one. var fkConstraint *models.Constraint - for _, constraint := range table.Constraints { - if constraint.Type == models.ForeignKeyConstraint && - (constraint.Name == fkName || constraint.ReferencedTable == rel.ToTable) { - fkConstraint = constraint - break + if c, ok := table.Constraints[fkName]; ok && c.Type == models.ForeignKeyConstraint { + fkConstraint = c + } else { + for _, constraint := range sortConstraints(table.Constraints) { + if constraint.Type == models.ForeignKeyConstraint && constraint.ReferencedTable == rel.ToTable { + fkConstraint = constraint + break + } } } diff --git a/pkg/writers/pgsql/writer_test.go b/pkg/writers/pgsql/writer_test.go index 5ac6179..39741d2 100644 --- a/pkg/writers/pgsql/writer_test.go +++ b/pkg/writers/pgsql/writer_test.go @@ -1549,3 +1549,52 @@ func TestGenerateColumnDefinition_IdentityColumnEmitsIdentityClauseNotDefault(t t.Fatalf("generateColumnDefinition() = %q, want %q", got, want) } } + +// Two FKs to the same table: each relationship must emit its own FK columns, +// not whichever FK map iteration returns first. +func TestWriteDatabase_MultipleForeignKeysToSameTable(t *testing.T) { + db := models.InitDatabase("testdb") + schema := models.InitSchema("public") + + dept := models.InitTable("department", "public") + id := models.InitColumn("id_department", "department", "public") + id.Type = "bigint" + id.IsPrimaryKey = true + dept.Columns["id_department"] = id + + emp := models.InitTable("employee", "public") + for _, name := range []string{"rid_department", "rid_manager"} { + col := models.InitColumn(name, "employee", "public") + col.Type = "bigint" + emp.Columns[name] = col + + fkName := "fk_employee_" + name + fk := models.InitConstraint(fkName, models.ForeignKeyConstraint) + fk.Schema, fk.Table, fk.Columns = "public", "employee", []string{name} + fk.ReferencedSchema, fk.ReferencedTable, fk.ReferencedColumns = "public", "department", []string{"id_department"} + emp.Constraints[fkName] = fk + + rel := models.InitRelationship("employee_to_department_"+name, models.OneToMany) + rel.FromTable, rel.ToTable, rel.ToSchema, rel.ForeignKey = "employee", "department", "public", fkName + emp.Relationships[rel.Name] = rel + } + + schema.Tables = append(schema.Tables, dept, emp) + db.Schemas = append(db.Schemas, schema) + + for i := 0; i < 20; i++ { + var buf bytes.Buffer + writer := NewWriter(&writers.WriterOptions{}) + writer.writer = &buf + if err := writer.WriteDatabase(db); err != nil { + t.Fatalf("WriteDatabase failed: %v", err) + } + output := buf.String() + for _, name := range []string{"rid_department", "rid_manager"} { + want := "ADD CONSTRAINT fk_employee_" + name + "\n FOREIGN KEY (" + name + ")" + if !strings.Contains(output, want) { + t.Fatalf("run %d: missing %q in output:\n%s", i, want, output) + } + } + } +}