From a32647ee1667ff546358e7ad4dc9b769aeeb3e11 Mon Sep 17 00:00:00 2001 From: Hein Date: Sat, 3 Oct 2026 21:33:59 +0200 Subject: [PATCH] fix: writer/reader correctness and determinism issues - merge: cloneTable keeps relationships; skip-tables applies to new schemas - diff: detect schema description/owner changes - prisma: reader no longer turns relation fields into columns, enum detection via declared names; writer type mapping is ordered - typeorm writer: keep explicit SQL types that cannot be inferred - drizzle: enum columns call the enum constant; reader resolves them - mysql/mssql/sqlite writers: honour OutputPath, deterministic column and constraint order; mssql live execute covers full schema - template: ToYAML recovers from panics; Merge nil-pointer loop - regenerate drizzle fixtures --- pkg/diff/diff.go | 17 +++ pkg/diff/types.go | 11 +- pkg/merge/merge.go | 19 ++++ pkg/readers/drizzle/reader.go | 27 +++++ pkg/readers/prisma/reader.go | 32 +++--- pkg/writers/drizzle/type_mapper.go | 4 +- pkg/writers/mssql/writer.go | 148 ++++++++++++------------- pkg/writers/mysql/writer.go | 65 ++++++++--- pkg/writers/prisma/writer.go | 56 +++++----- pkg/writers/sqlite/writer.go | 55 ++++++--- pkg/writers/template/formatters.go | 8 +- pkg/writers/template/safe_access.go | 7 +- pkg/writers/typeorm/writer.go | 20 ++-- tests/assets/drizzle/schema-updated.ts | 4 +- tests/assets/drizzle/schema.ts | 8 +- 15 files changed, 301 insertions(+), 180 deletions(-) diff --git a/pkg/diff/diff.go b/pkg/diff/diff.go index d278810..048e71f 100644 --- a/pkg/diff/diff.go +++ b/pkg/diff/diff.go @@ -71,6 +71,13 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff { return diff } +func (c *SchemaChange) addChange(field string, source, target any) { + if c.Changes == nil { + c.Changes = make(map[string]any) + } + c.Changes[field] = map[string]any{"source": source, "target": target} +} + func compareSchemaDetails(source, target *models.Schema) *SchemaChange { change := &SchemaChange{ Name: source.Name, @@ -78,6 +85,16 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange { hasChanges := false + // Compare schema attributes + if source.Description != target.Description { + change.addChange("description", source.Description, target.Description) + hasChanges = true + } + if source.Owner != target.Owner { + change.addChange("owner", source.Owner, target.Owner) + hasChanges = true + } + // Compare tables tableDiff := compareTables(source.Tables, target.Tables) if !isEmpty(tableDiff) { diff --git a/pkg/diff/types.go b/pkg/diff/types.go index 659ff85..73831ae 100644 --- a/pkg/diff/types.go +++ b/pkg/diff/types.go @@ -18,11 +18,12 @@ type SchemaDiff struct { // SchemaChange represents changes within a schema type SchemaChange struct { - Name string `json:"name"` - Tables *TableDiff `json:"tables,omitempty"` - Views *ViewDiff `json:"views,omitempty"` - Sequences *SequenceDiff `json:"sequences,omitempty"` - Scripts *ScriptDiff `json:"scripts,omitempty"` + Name string `json:"name"` + Changes map[string]any `json:"changes,omitempty"` // Schema attributes that differ (description, owner), keyed by field name + Tables *TableDiff `json:"tables,omitempty"` + Views *ViewDiff `json:"views,omitempty"` + Sequences *SequenceDiff `json:"sequences,omitempty"` + Scripts *ScriptDiff `json:"scripts,omitempty"` } // TableDiff represents differences in tables diff --git a/pkg/merge/merge.go b/pkg/merge/merge.go index 5029f34..e32f1df 100644 --- a/pkg/merge/merge.go +++ b/pkg/merge/merge.go @@ -82,6 +82,15 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions) } else { // Schema doesn't exist, add it newSchema := cloneSchema(srcSchema) + if len(opts.SkipTableNames) > 0 { + kept := newSchema.Tables[:0] + for _, t := range newSchema.Tables { + if !opts.SkipTableNames[strings.ToLower(t.SQLName())] { + kept = append(kept, t) + } + } + newSchema.Tables = kept + } target.Schemas = append(target.Schemas, newSchema) r.SchemasAdded++ } @@ -440,6 +449,8 @@ func cloneTable(table *models.Table) *models.Table { Description: table.Description, Schema: table.Schema, Comment: table.Comment, + Tablespace: table.Tablespace, + GUID: table.GUID, Sequence: table.Sequence, UpdatedAt: table.UpdatedAt, Columns: make(map[string]*models.Column), @@ -469,6 +480,14 @@ func cloneTable(table *models.Table) *models.Table { newTable.Indexes[idxName] = cloneIndex(index) } + // Clone relationships + if table.Relationships != nil { + newTable.Relationships = make(map[string]*models.Relationship, len(table.Relationships)) + for relName, rel := range table.Relationships { + newTable.Relationships[relName] = cloneRelation(rel) + } + } + return newTable } diff --git a/pkg/readers/drizzle/reader.go b/pkg/readers/drizzle/reader.go index 504e422..6287529 100644 --- a/pkg/readers/drizzle/reader.go +++ b/pkg/readers/drizzle/reader.go @@ -15,6 +15,9 @@ import ( // Reader implements the readers.Reader interface for Drizzle schema format type Reader struct { options *readers.ReaderOptions + // enumVars maps the constant a pgEnum() is assigned to (e.g. "role") to the + // enum's SQL name (e.g. "Role"), so columns declared as role('col') resolve. + enumVars map[string]string } // NewReader creates a new Drizzle reader with the given options @@ -29,6 +32,7 @@ func (r *Reader) ReadDatabase() (*models.Database, error) { if r.options.FilePath == "" { return nil, fmt.Errorf("file path is required for Drizzle reader") } + r.enumVars = make(map[string]string) // Check if it's a file or directory info, err := os.Stat(r.options.FilePath) @@ -100,6 +104,13 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) { return nil, fmt.Errorf("failed to glob directory: %w", err) } + // Enums may be declared in a different file than the tables using them + for _, file := range files { + if content, err := os.ReadFile(file); err == nil { + r.collectEnumVars(string(content)) + } + } + // Parse each file for _, file := range files { content, err := os.ReadFile(file) @@ -125,9 +136,22 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) { return db, nil } +var enumVarRegex = regexp.MustCompile(`export\s+const\s+(\w+)\s*=\s*pgEnum\s*\(\s*['"](\w+)['"]`) + +// collectEnumVars records every pgEnum() constant declared in content. +func (r *Reader) collectEnumVars(content string) { + if r.enumVars == nil { + r.enumVars = make(map[string]string) + } + for _, m := range enumVarRegex.FindAllStringSubmatch(content, -1) { + r.enumVars[m[1]] = m[2] + } +} + // parseDrizzle parses Drizzle schema content and returns a Database model func (r *Reader) parseDrizzle(content string) (*models.Database, error) { db := models.InitDatabase("database") + r.collectEnumVars(content) if r.options.Metadata != nil { if name, ok := r.options.Metadata["name"].(string); ok { @@ -375,6 +399,9 @@ func (r *Reader) parseColumnDefinition(line, fieldName, drizzleType string, tabl // Map Drizzle type to SQL type column.Type = r.drizzleTypeToSQL(drizzleType) + if enumName, ok := r.enumVars[drizzleType]; ok { + column.Type = enumName + } // Default: columns are nullable unless specified column.NotNull = false diff --git a/pkg/readers/prisma/reader.go b/pkg/readers/prisma/reader.go index 7de71d7..93d01a0 100644 --- a/pkg/readers/prisma/reader.go +++ b/pkg/readers/prisma/reader.go @@ -13,7 +13,8 @@ import ( // Reader implements the readers.Reader interface for Prisma schema format type Reader struct { - options *readers.ReaderOptions + options *readers.ReaderOptions + enumNames map[string]bool // enum names declared in the schema being parsed } // NewReader creates a new Prisma reader with the given options @@ -82,6 +83,8 @@ func (r *Reader) parsePrisma(content string) (*models.Database, error) { schema := models.InitSchema("public") schema.Enums = make([]*models.Enum, 0) + r.enumNames = collectEnumNames(content) + scanner := bufio.NewScanner(strings.NewReader(content)) // State tracking @@ -600,20 +603,21 @@ func (r *Reader) isPrimitiveType(typeName string) bool { return false } -// isEnumType checks if a type name might be an enum -// Note: We can't definitively check against schema.Enums at parse time -// because enums might be defined after the model, so we just check -// if it starts with uppercase (Prisma convention for enums) -func (r *Reader) isEnumType(typeName string, table *models.Table) bool { - // Simple heuristic: enum types start with uppercase letter - // and are not known model names (though we can't check that yet) - if len(typeName) > 0 && typeName[0] >= 'A' && typeName[0] <= 'Z' { - // Additional check: primitive types are already handled above - // So if it's uppercase and not primitive, it's likely an enum or model - // We'll assume it's an enum if it's a single word - return !strings.Contains(typeName, "_") +// isEnumType reports whether typeName is an enum declared in the schema. +// Enum names are collected up front because enums may be declared after the +// models that use them. +func (r *Reader) isEnumType(typeName string, _ *models.Table) bool { + return r.enumNames[typeName] +} + +var enumDeclRegex = regexp.MustCompile(`(?m)^\s*enum\s+(\w+)\s*{`) + +func collectEnumNames(content string) map[string]bool { + names := make(map[string]bool) + for _, m := range enumDeclRegex.FindAllStringSubmatch(content, -1) { + names[m[1]] = true } - return false + return names } // createConstraintFromRelation creates a FK constraint from a @relation attribute diff --git a/pkg/writers/drizzle/type_mapper.go b/pkg/writers/drizzle/type_mapper.go index dbe63e5..a47bc5d 100644 --- a/pkg/writers/drizzle/type_mapper.go +++ b/pkg/writers/drizzle/type_mapper.go @@ -100,8 +100,8 @@ func (tm *TypeMapper) BuildColumnChain(col *models.Column, table *models.Table, // Determine Drizzle column type var drizzleType string if isEnum { - // For enum types, use the type name directly - drizzleType = fmt.Sprintf("pgEnum('%s')", col.Type) + // Enum columns call the enum constant declared via pgEnum(...) + drizzleType = tm.ToCamelCase(col.Type) } else { drizzleType = tm.SQLTypeToDrizzle(col.Type) } diff --git a/pkg/writers/mssql/writer.go b/pkg/writers/mssql/writer.go index 78d84a6..c7fa264 100644 --- a/pkg/writers/mssql/writer.go +++ b/pkg/writers/mssql/writer.go @@ -44,26 +44,11 @@ func (w *Writer) WriteDatabase(db *models.Database) error { return w.executeDatabaseSQL(db, connString) } - var writer io.Writer - var file *os.File - var err error - - // Use existing writer if already set (for testing) - if w.writer != nil { - writer = w.writer - } else if w.options.OutputPath != "" { - // Determine output destination - file, err = os.Create(w.options.OutputPath) - if err != nil { - return fmt.Errorf("failed to create output file: %w", err) - } - defer file.Close() - writer = file - } else { - writer = os.Stdout + release, err := w.openOutput() + if err != nil { + return err } - - w.writer = writer + defer release() // Write header comment fmt.Fprintf(w.writer, "-- MSSQL Database Schema\n") @@ -80,11 +65,34 @@ func (w *Writer) WriteDatabase(db *models.Database) error { return nil } +// openOutput points w.writer at the configured destination (output file or +// stdout) when none is set, and returns a func that releases it again. +func (w *Writer) openOutput() (func(), error) { + if w.writer != nil { + return func() {}, nil + } + if w.options.OutputPath != "" { + file, err := os.Create(w.options.OutputPath) + if err != nil { + return nil, fmt.Errorf("failed to create output file: %w", err) + } + w.writer = file + return func() { + file.Close() + w.writer = nil + }, nil + } + w.writer = os.Stdout + return func() { w.writer = nil }, nil +} + // WriteSchema writes a single schema and all its tables func (w *Writer) WriteSchema(schema *models.Schema) error { - if w.writer == nil { - w.writer = os.Stdout + release, err := w.openOutput() + if err != nil { + return err } + defer release() // Phase 1: Create schema (skip dbo schema and when flattening) if schema.Name != "dbo" && !w.options.FlattenSchema { @@ -153,9 +161,11 @@ func (w *Writer) WriteSchema(schema *models.Schema) error { // WriteTable writes a single table with all its elements func (w *Writer) WriteTable(table *models.Table) error { - if w.writer == nil { - w.writer = os.Stdout + release, err := w.openOutput() + if err != nil { + return err } + defer release() // Create a temporary schema with just this table schema := models.InitSchema(table.Schema) @@ -481,18 +491,12 @@ func (w *Writer) writeComments(schema *models.Schema, table *models.Table) error return nil } -// executeDatabaseSQL executes SQL statements directly on an MSSQL database +// executeDatabaseSQL executes the full generated schema (tables, keys, +// indexes, constraints and comments) directly on an MSSQL database. func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) error { - // Generate SQL statements - statements := []string{} - statements = append(statements, "-- MSSQL Database Schema") - statements = append(statements, fmt.Sprintf("-- Database: %s", db.Name)) - statements = append(statements, "-- Generated by RelSpec") - - for _, schema := range db.Schemas { - if err := w.generateSchemaStatements(schema, &statements); err != nil { - return fmt.Errorf("failed to generate statements for schema %s: %w", schema.Name, err) - } + statements, err := w.generateStatements(db) + if err != nil { + return err } // Connect to database @@ -510,17 +514,9 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro // Execute statements executedCount := 0 for i, stmt := range statements { - stmtTrimmed := strings.TrimSpace(stmt) - - // Skip comments and empty statements - if strings.HasPrefix(stmtTrimmed, "--") || stmtTrimmed == "" { - continue - } - fmt.Fprintf(os.Stderr, "Executing statement %d/%d...\n", i+1, len(statements)) - _, execErr := dbConn.ExecContext(ctx, stmt) - if execErr != nil { + if _, execErr := dbConn.ExecContext(ctx, stmt); execErr != nil { fmt.Fprintf(os.Stderr, "⚠ Warning: Statement failed: %v\n", execErr) continue } @@ -532,49 +528,51 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro return nil } -// generateSchemaStatements generates SQL statements for a schema -func (w *Writer) generateSchemaStatements(schema *models.Schema, statements *[]string) error { - // Phase 1: Create schema - if schema.Name != "dbo" && !w.options.FlattenSchema { - *statements = append(*statements, fmt.Sprintf("-- Schema: %s", schema.Name)) - *statements = append(*statements, fmt.Sprintf("CREATE SCHEMA [%s];", schema.Name)) - } +// generateStatements renders the same script WriteDatabase would produce and +// splits it into individually executable statements (comments removed). +func (w *Writer) generateStatements(db *models.Database) ([]string, error) { + var buf strings.Builder + saved := w.writer + w.writer = &buf + defer func() { w.writer = saved }() - // Phase 2: Create tables - *statements = append(*statements, fmt.Sprintf("-- Tables for schema: %s", schema.Name)) - for _, table := range schema.Tables { - createTableSQL := fmt.Sprintf("CREATE TABLE %s (", w.qualTable(schema.Name, table.Name)) - columnDefs := make([]string, 0) - - columns := getSortedColumns(table.Columns) - for _, col := range columns { - def := w.generateColumnDefinition(col) - columnDefs = append(columnDefs, " "+def) + for _, schema := range db.Schemas { + if err := w.WriteSchema(schema); err != nil { + return nil, fmt.Errorf("failed to generate statements for schema %s: %w", schema.Name, err) } - - createTableSQL += "\n" + strings.Join(columnDefs, ",\n") + "\n)" - *statements = append(*statements, createTableSQL) } - // Phase 3-7: Constraints and indexes will be added by WriteSchema logic - // For now, just create tables - return nil + // Every statement the writer emits ends with ";\n\n". + var statements []string + for _, chunk := range strings.Split(buf.String(), ";\n\n") { + lines := make([]string, 0) + for _, line := range strings.Split(chunk, "\n") { + if !strings.HasPrefix(strings.TrimSpace(line), "--") { + lines = append(lines, line) + } + } + if stmt := strings.TrimSpace(strings.Join(lines, "\n")); stmt != "" { + statements = append(statements, stmt) + } + } + return statements, nil } // Helper functions -// getSortedColumns returns columns sorted by sequence +// getSortedColumns returns columns sorted by sequence, then by name so that +// columns without a sequence still come out in a stable order. func getSortedColumns(columns map[string]*models.Column) []*models.Column { - names := make([]string, 0, len(columns)) - for name := range columns { - names = append(names, name) - } - sort.Strings(names) - sorted := make([]*models.Column, 0, len(columns)) - for _, name := range names { - sorted = append(sorted, columns[name]) + for _, col := range columns { + sorted = append(sorted, col) } + sort.Slice(sorted, func(i, j int) bool { + if sorted[i].Sequence != sorted[j].Sequence { + return sorted[i].Sequence < sorted[j].Sequence + } + return sorted[i].Name < sorted[j].Name + }) return sorted } diff --git a/pkg/writers/mysql/writer.go b/pkg/writers/mysql/writer.go index a0fda45..186d491 100644 --- a/pkg/writers/mysql/writer.go +++ b/pkg/writers/mysql/writer.go @@ -29,18 +29,11 @@ func (w *Writer) WriteDatabase(db *models.Database) error { if conn, ok := w.options.Metadata["connection_string"].(string); ok && conn != "" { return w.execute(db, conn) } - if w.writer == nil { - if w.options.OutputPath != "" { - f, err := os.Create(w.options.OutputPath) - if err != nil { - return err - } - defer f.Close() - w.writer = f - } else { - w.writer = os.Stdout - } + release, err := w.openOutput() + if err != nil { + return err } + defer release() return w.writeDatabaseDDL(db) } @@ -56,7 +49,33 @@ func (w *Writer) writeDatabaseDDL(db *models.Database) error { return nil } +// openOutput points w.writer at the configured destination (output file or +// stdout) when none is set, and returns a func that releases it again. +func (w *Writer) openOutput() (func(), error) { + if w.writer != nil { + return func() {}, nil + } + if w.options != nil && w.options.OutputPath != "" { + f, err := os.Create(w.options.OutputPath) + if err != nil { + return nil, err + } + w.writer = f + return func() { + f.Close() + w.writer = nil + }, nil + } + w.writer = os.Stdout + return func() { w.writer = nil }, nil +} + func (w *Writer) WriteSchema(s *models.Schema) error { + release, err := w.openOutput() + if err != nil { + return err + } + defer release() for _, t := range s.Tables { if err := w.writeTable(s, t); err != nil { return err @@ -66,9 +85,11 @@ func (w *Writer) WriteSchema(s *models.Schema) error { } func (w *Writer) WriteTable(t *models.Table) error { - if w.writer == nil { - w.writer = os.Stdout + release, err := w.openOutput() + if err != nil { + return err } + defer release() return w.writeTable(nil, t) } @@ -83,7 +104,12 @@ func (w *Writer) writeTable(s *models.Schema, t *models.Table) error { for _, c := range t.Columns { cols = append(cols, c) } - sort.Slice(cols, func(i, j int) bool { return cols[i].Sequence < cols[j].Sequence }) + sort.Slice(cols, func(i, j int) bool { + if cols[i].Sequence != cols[j].Sequence { + return cols[i].Sequence < cols[j].Sequence + } + return cols[i].Name < cols[j].Name + }) defs := []string{} pk := []string{} for _, c := range cols { @@ -105,7 +131,13 @@ func (w *Writer) writeTable(s *models.Schema, t *models.Table) error { pk = append(pk, quote(c.Name)) } } - for _, c := range t.Constraints { + constraintNames := make([]string, 0, len(t.Constraints)) + for name := range t.Constraints { + constraintNames = append(constraintNames, name) + } + sort.Strings(constraintNames) + for _, name := range constraintNames { + c := t.Constraints[name] if c.Type == models.PrimaryKeyConstraint { pk = nil for _, n := range c.Columns { @@ -117,7 +149,8 @@ func (w *Writer) writeTable(s *models.Schema, t *models.Table) error { if len(pk) > 0 { defs = append(defs, " PRIMARY KEY ("+strings.Join(pk, ", ")+")") } - for _, c := range t.Constraints { + for _, name := range constraintNames { + c := t.Constraints[name] if c.Type == models.UniqueConstraint { defs = append(defs, fmt.Sprintf(" CONSTRAINT %s UNIQUE (%s)", quote(c.Name), quoted(c.Columns))) } diff --git a/pkg/writers/prisma/writer.go b/pkg/writers/prisma/writer.go index a7eb16a..a925908 100644 --- a/pkg/writers/prisma/writer.go +++ b/pkg/writers/prisma/writer.go @@ -266,35 +266,37 @@ func (w *Writer) sqlTypeToPrisma(sqlType string, schema *models.Schema) string { } } - // Standard type mapping - typeMap := map[string]string{ - "text": "String", - "varchar": "String", - "character varying": "String", - "char": "String", - "boolean": "Boolean", - "bool": "Boolean", - "integer": "Int", - "int": "Int", - "int4": "Int", - "bigint": "BigInt", - "int8": "BigInt", - "double precision": "Float", - "float": "Float", - "float8": "Float", - "decimal": "Decimal", - "numeric": "Decimal", - "timestamp": "DateTime", - "timestamptz": "DateTime", - "date": "DateTime", - "jsonb": "Json", - "json": "Json", - "bytea": "Bytes", + // Ordered so more specific patterns win (bigint/int8 before int); map + // iteration would make the result nondeterministic. + typeMap := []struct{ pattern, prismaType string }{ + {"bigint", "BigInt"}, + {"int8", "BigInt"}, + {"text", "String"}, + {"varchar", "String"}, + {"character varying", "String"}, + {"char", "String"}, + {"boolean", "Boolean"}, + {"bool", "Boolean"}, + {"integer", "Int"}, + {"int4", "Int"}, + {"int", "Int"}, + {"double precision", "Float"}, + {"float8", "Float"}, + {"float", "Float"}, + {"decimal", "Decimal"}, + {"numeric", "Decimal"}, + {"timestamptz", "DateTime"}, + {"timestamp", "DateTime"}, + {"date", "DateTime"}, + {"jsonb", "Json"}, + {"json", "Json"}, + {"bytea", "Bytes"}, } - for sqlPattern, prismaType := range typeMap { - if strings.Contains(strings.ToLower(sqlType), sqlPattern) { - return prismaType + lower := strings.ToLower(sqlType) + for _, m := range typeMap { + if strings.Contains(lower, m.pattern) { + return m.prismaType } } diff --git a/pkg/writers/sqlite/writer.go b/pkg/writers/sqlite/writer.go index 3e93663..2018c48 100644 --- a/pkg/writers/sqlite/writer.go +++ b/pkg/writers/sqlite/writer.go @@ -44,29 +44,37 @@ func (w *Writer) WriteDatabase(db *models.Database) error { return w.executeDatabaseSQL(db, dbPath) } - var writer io.Writer - var file *os.File - var err error - - // Use existing writer if already set (for testing) - if w.writer != nil { - writer = w.writer - } else if w.options.OutputPath != "" { - // Determine output destination - file, err = os.Create(w.options.OutputPath) - if err != nil { - return fmt.Errorf("failed to create output file: %w", err) - } - defer file.Close() - writer = file - } else { - writer = os.Stdout + release, err := w.openOutput() + if err != nil { + return err } + defer release() - w.writer = writer return w.writeContent(db) } +// openOutput points w.writer at the configured destination (output file or +// stdout) when none is set, and returns a func that releases it again so the +// writer can be reused. +func (w *Writer) openOutput() (func(), error) { + if w.writer != nil { + return func() {}, nil + } + if w.options.OutputPath != "" { + file, err := os.Create(w.options.OutputPath) + if err != nil { + return nil, fmt.Errorf("failed to create output file: %w", err) + } + w.writer = file + return func() { + file.Close() + w.writer = nil + }, nil + } + w.writer = os.Stdout + return func() { w.writer = nil }, nil +} + // writeContent writes the header, pragma, and every schema's DDL to w.writer. func (w *Writer) writeContent(db *models.Database) error { // Write header comment @@ -184,6 +192,12 @@ func tableSchemaName(schema string) string { // WriteSchema writes a single schema as SQLite SQL func (w *Writer) WriteSchema(schema *models.Schema) error { + release, err := w.openOutput() + if err != nil { + return err + } + defer release() + tableSchema := tableSchemaName(schema.Name) if err := w.checkDirectives(schema); err != nil { @@ -229,6 +243,11 @@ func (w *Writer) WriteSchema(schema *models.Schema) error { // WriteTable writes a single table as SQLite SQL func (w *Writer) WriteTable(table *models.Table) error { + release, err := w.openOutput() + if err != nil { + return err + } + defer release() return w.writeTable("", table) } diff --git a/pkg/writers/template/formatters.go b/pkg/writers/template/formatters.go index b048648..878303e 100644 --- a/pkg/writers/template/formatters.go +++ b/pkg/writers/template/formatters.go @@ -30,7 +30,13 @@ func ToJSONPretty(v interface{}, indent string) string { // ToYAML converts a value to YAML string // Usage: {{ .Database | toYAML }} -func ToYAML(v interface{}) string { +func ToYAML(v interface{}) (out string) { + // yaml.v3 panics (rather than returning an error) for unsupported types such as channels. + defer func() { + if r := recover(); r != nil { + out = fmt.Sprintf("error: failed to marshal: %v", r) + } + }() data, err := yaml.Marshal(v) if err != nil { return fmt.Sprintf("error: failed to marshal: %v", err) diff --git a/pkg/writers/template/safe_access.go b/pkg/writers/template/safe_access.go index 2c97a0d..c24ddb6 100644 --- a/pkg/writers/template/safe_access.go +++ b/pkg/writers/template/safe_access.go @@ -101,11 +101,8 @@ func Merge(maps ...interface{}) map[interface{}]interface{} { for _, m := range maps { v := reflect.ValueOf(m) - // Dereference pointers - for v.Kind() == reflect.Pointer { - if v.IsNil() { - continue - } + // Dereference pointers; a nil pointer contributes nothing + for v.Kind() == reflect.Pointer && !v.IsNil() { v = v.Elem() } diff --git a/pkg/writers/typeorm/writer.go b/pkg/writers/typeorm/writer.go index 3c4b980..db1289c 100644 --- a/pkg/writers/typeorm/writer.go +++ b/pkg/writers/typeorm/writer.go @@ -394,16 +394,20 @@ func escapeSingleQuoted(s string) string { return strings.ReplaceAll(strings.ReplaceAll(s, `\`, `\\`), `'`, `\'`) } -// needsExplicitType checks if a SQL type needs explicit type declaration +// needsExplicitType checks if a SQL type needs explicit type declaration. +// The reader infers a column type from the TypeScript type alone when no +// explicit type is given, so any type that differs from that inferred type +// (varchar(255), numeric(10,2), timestamptz, ...) must be written out. func (w *Writer) needsExplicitType(sqlType string) bool { - // Types that don't map cleanly to TypeScript types need explicit declaration - explicitTypes := []string{"text", "uuid", "jsonb", "bigint"} - for _, t := range explicitTypes { - if strings.Contains(sqlType, t) { - return true - } + inferred := map[string]string{ + "string": "text", + "number": "integer", + "boolean": "boolean", + "Date": "timestamp", + "any": "jsonb", } - return false + return inferred[w.sqlTypeToTypeScript(sqlType)] != strings.ToLower(strings.TrimSpace(sqlType)) || + strings.Contains(sqlType, "uuid") || strings.Contains(sqlType, "bigint") } // hasUniqueConstraint checks if a column has a unique constraint diff --git a/tests/assets/drizzle/schema-updated.ts b/tests/assets/drizzle/schema-updated.ts index 6f56da2..f088c79 100644 --- a/tests/assets/drizzle/schema-updated.ts +++ b/tests/assets/drizzle/schema-updated.ts @@ -17,7 +17,7 @@ export const users = pgTable('users', { lastLoginAt: timestamp('last_login_at'), passwordHash: varchar('password_hash').notNull(), profile: jsonb('profile'), - role: pgEnum('UserRole')('role').notNull(), + role: userRole('role').notNull(), updatedAt: timestamp('updated_at').notNull().default(sql`now()`), username: varchar('username').notNull().unique(), }); @@ -131,7 +131,7 @@ export const orders = pgTable('orders', { notes: text('notes'), orderNumber: varchar('order_number').notNull().unique(), shippingAddress: jsonb('shipping_address').notNull(), - status: pgEnum('OrderStatus')('status').notNull().default('pending'), + status: orderStatus('status').notNull().default('pending'), totalAmount: numeric('total_amount').notNull(), updatedAt: timestamp('updated_at').notNull().default(sql`now()`), userId: integer('user_id').notNull().references(() => users.id), diff --git a/tests/assets/drizzle/schema.ts b/tests/assets/drizzle/schema.ts index 04a8c5e..b510878 100644 --- a/tests/assets/drizzle/schema.ts +++ b/tests/assets/drizzle/schema.ts @@ -13,7 +13,6 @@ export interface User { id: number; email: string; name: string | null; - profile: string | null; role: Role; } @@ -21,8 +20,7 @@ export const user = pgTable('User', { id: integer('id').primaryKey().generatedAlwaysAsIdentity(), email: text('email').notNull().unique(), name: text('name'), - profile: text('profile'), - role: pgEnum('Role')('role').notNull().default('USER'), + role: role('role').notNull().default('USER'), }); export type NewUser = typeof user.$inferInsert; @@ -30,14 +28,12 @@ export type NewUser = typeof user.$inferInsert; export interface Profile { id: number; bio: string; - user: string; userId: number; } export const profile = pgTable('Profile', { id: integer('id').primaryKey().generatedAlwaysAsIdentity(), bio: text('bio').notNull(), - user: text('user').notNull(), userId: integer('userId').notNull().unique().references(() => user.id), }); @@ -45,7 +41,6 @@ export type NewProfile = typeof profile.$inferInsert; // Table: Post export interface Post { id: number; - author: string; authorId: number; createdAt: Date; published: boolean; @@ -55,7 +50,6 @@ export interface Post { export const post = pgTable('Post', { id: integer('id').primaryKey().generatedAlwaysAsIdentity(), - author: text('author').notNull(), authorId: integer('authorId').notNull().references(() => user.id), createdAt: timestamp('createdAt').notNull().default(sql`now()`), published: boolean('published').notNull().default(false),