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
This commit is contained in:
2026-10-03 21:33:59 +02:00
parent 08e1417393
commit a32647ee16
15 changed files with 301 additions and 180 deletions
+17
View File
@@ -71,6 +71,13 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
return diff 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 { func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
change := &SchemaChange{ change := &SchemaChange{
Name: source.Name, Name: source.Name,
@@ -78,6 +85,16 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
hasChanges := false 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 // Compare tables
tableDiff := compareTables(source.Tables, target.Tables) tableDiff := compareTables(source.Tables, target.Tables)
if !isEmpty(tableDiff) { if !isEmpty(tableDiff) {
+1
View File
@@ -19,6 +19,7 @@ type SchemaDiff struct {
// SchemaChange represents changes within a schema // SchemaChange represents changes within a schema
type SchemaChange struct { type SchemaChange struct {
Name string `json:"name"` 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"` Tables *TableDiff `json:"tables,omitempty"`
Views *ViewDiff `json:"views,omitempty"` Views *ViewDiff `json:"views,omitempty"`
Sequences *SequenceDiff `json:"sequences,omitempty"` Sequences *SequenceDiff `json:"sequences,omitempty"`
+19
View File
@@ -82,6 +82,15 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions)
} else { } else {
// Schema doesn't exist, add it // Schema doesn't exist, add it
newSchema := cloneSchema(srcSchema) 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) target.Schemas = append(target.Schemas, newSchema)
r.SchemasAdded++ r.SchemasAdded++
} }
@@ -440,6 +449,8 @@ func cloneTable(table *models.Table) *models.Table {
Description: table.Description, Description: table.Description,
Schema: table.Schema, Schema: table.Schema,
Comment: table.Comment, Comment: table.Comment,
Tablespace: table.Tablespace,
GUID: table.GUID,
Sequence: table.Sequence, Sequence: table.Sequence,
UpdatedAt: table.UpdatedAt, UpdatedAt: table.UpdatedAt,
Columns: make(map[string]*models.Column), Columns: make(map[string]*models.Column),
@@ -469,6 +480,14 @@ func cloneTable(table *models.Table) *models.Table {
newTable.Indexes[idxName] = cloneIndex(index) 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 return newTable
} }
+27
View File
@@ -15,6 +15,9 @@ import (
// Reader implements the readers.Reader interface for Drizzle schema format // Reader implements the readers.Reader interface for Drizzle schema format
type Reader struct { type Reader struct {
options *readers.ReaderOptions 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 // 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 == "" { if r.options.FilePath == "" {
return nil, fmt.Errorf("file path is required for Drizzle reader") 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 // Check if it's a file or directory
info, err := os.Stat(r.options.FilePath) 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) 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 // Parse each file
for _, file := range files { for _, file := range files {
content, err := os.ReadFile(file) content, err := os.ReadFile(file)
@@ -125,9 +136,22 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
return db, nil 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 // parseDrizzle parses Drizzle schema content and returns a Database model
func (r *Reader) parseDrizzle(content string) (*models.Database, error) { func (r *Reader) parseDrizzle(content string) (*models.Database, error) {
db := models.InitDatabase("database") db := models.InitDatabase("database")
r.collectEnumVars(content)
if r.options.Metadata != nil { if r.options.Metadata != nil {
if name, ok := r.options.Metadata["name"].(string); ok { 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 // Map Drizzle type to SQL type
column.Type = r.drizzleTypeToSQL(drizzleType) column.Type = r.drizzleTypeToSQL(drizzleType)
if enumName, ok := r.enumVars[drizzleType]; ok {
column.Type = enumName
}
// Default: columns are nullable unless specified // Default: columns are nullable unless specified
column.NotNull = false column.NotNull = false
+17 -13
View File
@@ -14,6 +14,7 @@ import (
// Reader implements the readers.Reader interface for Prisma schema format // Reader implements the readers.Reader interface for Prisma schema format
type Reader struct { 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 // 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 := models.InitSchema("public")
schema.Enums = make([]*models.Enum, 0) schema.Enums = make([]*models.Enum, 0)
r.enumNames = collectEnumNames(content)
scanner := bufio.NewScanner(strings.NewReader(content)) scanner := bufio.NewScanner(strings.NewReader(content))
// State tracking // State tracking
@@ -600,20 +603,21 @@ func (r *Reader) isPrimitiveType(typeName string) bool {
return false return false
} }
// isEnumType checks if a type name might be an enum // isEnumType reports whether typeName is an enum declared in the schema.
// Note: We can't definitively check against schema.Enums at parse time // Enum names are collected up front because enums may be declared after the
// because enums might be defined after the model, so we just check // models that use them.
// if it starts with uppercase (Prisma convention for enums) func (r *Reader) isEnumType(typeName string, _ *models.Table) bool {
func (r *Reader) isEnumType(typeName string, table *models.Table) bool { return r.enumNames[typeName]
// 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, "_")
} }
return false
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 names
} }
// createConstraintFromRelation creates a FK constraint from a @relation attribute // createConstraintFromRelation creates a FK constraint from a @relation attribute
+2 -2
View File
@@ -100,8 +100,8 @@ func (tm *TypeMapper) BuildColumnChain(col *models.Column, table *models.Table,
// Determine Drizzle column type // Determine Drizzle column type
var drizzleType string var drizzleType string
if isEnum { if isEnum {
// For enum types, use the type name directly // Enum columns call the enum constant declared via pgEnum(...)
drizzleType = fmt.Sprintf("pgEnum('%s')", col.Type) drizzleType = tm.ToCamelCase(col.Type)
} else { } else {
drizzleType = tm.SQLTypeToDrizzle(col.Type) drizzleType = tm.SQLTypeToDrizzle(col.Type)
} }
+72 -74
View File
@@ -44,26 +44,11 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
return w.executeDatabaseSQL(db, connString) return w.executeDatabaseSQL(db, connString)
} }
var writer io.Writer release, err := w.openOutput()
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 { if err != nil {
return fmt.Errorf("failed to create output file: %w", err) return err
} }
defer file.Close() defer release()
writer = file
} else {
writer = os.Stdout
}
w.writer = writer
// Write header comment // Write header comment
fmt.Fprintf(w.writer, "-- MSSQL Database Schema\n") fmt.Fprintf(w.writer, "-- MSSQL Database Schema\n")
@@ -80,11 +65,34 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
return nil 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 // WriteSchema writes a single schema and all its tables
func (w *Writer) WriteSchema(schema *models.Schema) error { func (w *Writer) WriteSchema(schema *models.Schema) error {
if w.writer == nil { release, err := w.openOutput()
w.writer = os.Stdout if err != nil {
return err
} }
defer release()
// Phase 1: Create schema (skip dbo schema and when flattening) // Phase 1: Create schema (skip dbo schema and when flattening)
if schema.Name != "dbo" && !w.options.FlattenSchema { 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 // WriteTable writes a single table with all its elements
func (w *Writer) WriteTable(table *models.Table) error { func (w *Writer) WriteTable(table *models.Table) error {
if w.writer == nil { release, err := w.openOutput()
w.writer = os.Stdout if err != nil {
return err
} }
defer release()
// Create a temporary schema with just this table // Create a temporary schema with just this table
schema := models.InitSchema(table.Schema) schema := models.InitSchema(table.Schema)
@@ -481,18 +491,12 @@ func (w *Writer) writeComments(schema *models.Schema, table *models.Table) error
return nil 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 { func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) error {
// Generate SQL statements statements, err := w.generateStatements(db)
statements := []string{} if err != nil {
statements = append(statements, "-- MSSQL Database Schema") return err
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)
}
} }
// Connect to database // Connect to database
@@ -510,17 +514,9 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
// Execute statements // Execute statements
executedCount := 0 executedCount := 0
for i, stmt := range statements { 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)) fmt.Fprintf(os.Stderr, "Executing statement %d/%d...\n", i+1, len(statements))
_, execErr := dbConn.ExecContext(ctx, stmt) if _, execErr := dbConn.ExecContext(ctx, stmt); execErr != nil {
if execErr != nil {
fmt.Fprintf(os.Stderr, "⚠ Warning: Statement failed: %v\n", execErr) fmt.Fprintf(os.Stderr, "⚠ Warning: Statement failed: %v\n", execErr)
continue continue
} }
@@ -532,49 +528,51 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
return nil return nil
} }
// generateSchemaStatements generates SQL statements for a schema // generateStatements renders the same script WriteDatabase would produce and
func (w *Writer) generateSchemaStatements(schema *models.Schema, statements *[]string) error { // splits it into individually executable statements (comments removed).
// Phase 1: Create schema func (w *Writer) generateStatements(db *models.Database) ([]string, error) {
if schema.Name != "dbo" && !w.options.FlattenSchema { var buf strings.Builder
*statements = append(*statements, fmt.Sprintf("-- Schema: %s", schema.Name)) saved := w.writer
*statements = append(*statements, fmt.Sprintf("CREATE SCHEMA [%s];", schema.Name)) w.writer = &buf
defer func() { w.writer = saved }()
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)
}
} }
// Phase 2: Create tables // Every statement the writer emits ends with ";\n\n".
*statements = append(*statements, fmt.Sprintf("-- Tables for schema: %s", schema.Name)) var statements []string
for _, table := range schema.Tables { for _, chunk := range strings.Split(buf.String(), ";\n\n") {
createTableSQL := fmt.Sprintf("CREATE TABLE %s (", w.qualTable(schema.Name, table.Name)) lines := make([]string, 0)
columnDefs := make([]string, 0) for _, line := range strings.Split(chunk, "\n") {
if !strings.HasPrefix(strings.TrimSpace(line), "--") {
columns := getSortedColumns(table.Columns) lines = append(lines, line)
for _, col := range columns {
def := w.generateColumnDefinition(col)
columnDefs = append(columnDefs, " "+def)
} }
createTableSQL += "\n" + strings.Join(columnDefs, ",\n") + "\n)"
*statements = append(*statements, createTableSQL)
} }
if stmt := strings.TrimSpace(strings.Join(lines, "\n")); stmt != "" {
// Phase 3-7: Constraints and indexes will be added by WriteSchema logic statements = append(statements, stmt)
// For now, just create tables }
return nil }
return statements, nil
} }
// Helper functions // 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 { 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)) sorted := make([]*models.Column, 0, len(columns))
for _, name := range names { for _, col := range columns {
sorted = append(sorted, columns[name]) 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 return sorted
} }
+47 -14
View File
@@ -29,18 +29,11 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
if conn, ok := w.options.Metadata["connection_string"].(string); ok && conn != "" { if conn, ok := w.options.Metadata["connection_string"].(string); ok && conn != "" {
return w.execute(db, conn) return w.execute(db, conn)
} }
if w.writer == nil { release, err := w.openOutput()
if w.options.OutputPath != "" {
f, err := os.Create(w.options.OutputPath)
if err != nil { if err != nil {
return err return err
} }
defer f.Close() defer release()
w.writer = f
} else {
w.writer = os.Stdout
}
}
return w.writeDatabaseDDL(db) return w.writeDatabaseDDL(db)
} }
@@ -56,7 +49,33 @@ func (w *Writer) writeDatabaseDDL(db *models.Database) error {
return nil 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 { func (w *Writer) WriteSchema(s *models.Schema) error {
release, err := w.openOutput()
if err != nil {
return err
}
defer release()
for _, t := range s.Tables { for _, t := range s.Tables {
if err := w.writeTable(s, t); err != nil { if err := w.writeTable(s, t); err != nil {
return err return err
@@ -66,9 +85,11 @@ func (w *Writer) WriteSchema(s *models.Schema) error {
} }
func (w *Writer) WriteTable(t *models.Table) error { func (w *Writer) WriteTable(t *models.Table) error {
if w.writer == nil { release, err := w.openOutput()
w.writer = os.Stdout if err != nil {
return err
} }
defer release()
return w.writeTable(nil, t) 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 { for _, c := range t.Columns {
cols = append(cols, c) 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{} defs := []string{}
pk := []string{} pk := []string{}
for _, c := range cols { 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)) 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 { if c.Type == models.PrimaryKeyConstraint {
pk = nil pk = nil
for _, n := range c.Columns { for _, n := range c.Columns {
@@ -117,7 +149,8 @@ func (w *Writer) writeTable(s *models.Schema, t *models.Table) error {
if len(pk) > 0 { if len(pk) > 0 {
defs = append(defs, " PRIMARY KEY ("+strings.Join(pk, ", ")+")") 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 { if c.Type == models.UniqueConstraint {
defs = append(defs, fmt.Sprintf(" CONSTRAINT %s UNIQUE (%s)", quote(c.Name), quoted(c.Columns))) defs = append(defs, fmt.Sprintf(" CONSTRAINT %s UNIQUE (%s)", quote(c.Name), quoted(c.Columns)))
} }
+29 -27
View File
@@ -266,35 +266,37 @@ func (w *Writer) sqlTypeToPrisma(sqlType string, schema *models.Schema) string {
} }
} }
// Standard type mapping // Ordered so more specific patterns win (bigint/int8 before int); map
typeMap := map[string]string{ // iteration would make the result nondeterministic.
"text": "String", typeMap := []struct{ pattern, prismaType string }{
"varchar": "String", {"bigint", "BigInt"},
"character varying": "String", {"int8", "BigInt"},
"char": "String", {"text", "String"},
"boolean": "Boolean", {"varchar", "String"},
"bool": "Boolean", {"character varying", "String"},
"integer": "Int", {"char", "String"},
"int": "Int", {"boolean", "Boolean"},
"int4": "Int", {"bool", "Boolean"},
"bigint": "BigInt", {"integer", "Int"},
"int8": "BigInt", {"int4", "Int"},
"double precision": "Float", {"int", "Int"},
"float": "Float", {"double precision", "Float"},
"float8": "Float", {"float8", "Float"},
"decimal": "Decimal", {"float", "Float"},
"numeric": "Decimal", {"decimal", "Decimal"},
"timestamp": "DateTime", {"numeric", "Decimal"},
"timestamptz": "DateTime", {"timestamptz", "DateTime"},
"date": "DateTime", {"timestamp", "DateTime"},
"jsonb": "Json", {"date", "DateTime"},
"json": "Json", {"jsonb", "Json"},
"bytea": "Bytes", {"json", "Json"},
{"bytea", "Bytes"},
} }
for sqlPattern, prismaType := range typeMap { lower := strings.ToLower(sqlType)
if strings.Contains(strings.ToLower(sqlType), sqlPattern) { for _, m := range typeMap {
return prismaType if strings.Contains(lower, m.pattern) {
return m.prismaType
} }
} }
+36 -17
View File
@@ -44,27 +44,35 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
return w.executeDatabaseSQL(db, dbPath) return w.executeDatabaseSQL(db, dbPath)
} }
var writer io.Writer release, err := w.openOutput()
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 { if err != nil {
return fmt.Errorf("failed to create output file: %w", err) return err
} }
defer file.Close() defer release()
writer = file
} else { return w.writeContent(db)
writer = os.Stdout
} }
w.writer = writer // openOutput points w.writer at the configured destination (output file or
return w.writeContent(db) // 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. // writeContent writes the header, pragma, and every schema's DDL to w.writer.
@@ -184,6 +192,12 @@ func tableSchemaName(schema string) string {
// WriteSchema writes a single schema as SQLite SQL // WriteSchema writes a single schema as SQLite SQL
func (w *Writer) WriteSchema(schema *models.Schema) error { func (w *Writer) WriteSchema(schema *models.Schema) error {
release, err := w.openOutput()
if err != nil {
return err
}
defer release()
tableSchema := tableSchemaName(schema.Name) tableSchema := tableSchemaName(schema.Name)
if err := w.checkDirectives(schema); err != nil { 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 // WriteTable writes a single table as SQLite SQL
func (w *Writer) WriteTable(table *models.Table) error { func (w *Writer) WriteTable(table *models.Table) error {
release, err := w.openOutput()
if err != nil {
return err
}
defer release()
return w.writeTable("", table) return w.writeTable("", table)
} }
+7 -1
View File
@@ -30,7 +30,13 @@ func ToJSONPretty(v interface{}, indent string) string {
// ToYAML converts a value to YAML string // ToYAML converts a value to YAML string
// Usage: {{ .Database | toYAML }} // 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) data, err := yaml.Marshal(v)
if err != nil { if err != nil {
return fmt.Sprintf("error: failed to marshal: %v", err) return fmt.Sprintf("error: failed to marshal: %v", err)
+2 -5
View File
@@ -101,11 +101,8 @@ func Merge(maps ...interface{}) map[interface{}]interface{} {
for _, m := range maps { for _, m := range maps {
v := reflect.ValueOf(m) v := reflect.ValueOf(m)
// Dereference pointers // Dereference pointers; a nil pointer contributes nothing
for v.Kind() == reflect.Pointer { for v.Kind() == reflect.Pointer && !v.IsNil() {
if v.IsNil() {
continue
}
v = v.Elem() v = v.Elem()
} }
+12 -8
View File
@@ -394,16 +394,20 @@ func escapeSingleQuoted(s string) string {
return strings.ReplaceAll(strings.ReplaceAll(s, `\`, `\\`), `'`, `\'`) 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 { func (w *Writer) needsExplicitType(sqlType string) bool {
// Types that don't map cleanly to TypeScript types need explicit declaration inferred := map[string]string{
explicitTypes := []string{"text", "uuid", "jsonb", "bigint"} "string": "text",
for _, t := range explicitTypes { "number": "integer",
if strings.Contains(sqlType, t) { "boolean": "boolean",
return true "Date": "timestamp",
"any": "jsonb",
} }
} return inferred[w.sqlTypeToTypeScript(sqlType)] != strings.ToLower(strings.TrimSpace(sqlType)) ||
return false strings.Contains(sqlType, "uuid") || strings.Contains(sqlType, "bigint")
} }
// hasUniqueConstraint checks if a column has a unique constraint // hasUniqueConstraint checks if a column has a unique constraint
+2 -2
View File
@@ -17,7 +17,7 @@ export const users = pgTable('users', {
lastLoginAt: timestamp('last_login_at'), lastLoginAt: timestamp('last_login_at'),
passwordHash: varchar('password_hash').notNull(), passwordHash: varchar('password_hash').notNull(),
profile: jsonb('profile'), profile: jsonb('profile'),
role: pgEnum('UserRole')('role').notNull(), role: userRole('role').notNull(),
updatedAt: timestamp('updated_at').notNull().default(sql`now()`), updatedAt: timestamp('updated_at').notNull().default(sql`now()`),
username: varchar('username').notNull().unique(), username: varchar('username').notNull().unique(),
}); });
@@ -131,7 +131,7 @@ export const orders = pgTable('orders', {
notes: text('notes'), notes: text('notes'),
orderNumber: varchar('order_number').notNull().unique(), orderNumber: varchar('order_number').notNull().unique(),
shippingAddress: jsonb('shipping_address').notNull(), shippingAddress: jsonb('shipping_address').notNull(),
status: pgEnum('OrderStatus')('status').notNull().default('pending'), status: orderStatus('status').notNull().default('pending'),
totalAmount: numeric('total_amount').notNull(), totalAmount: numeric('total_amount').notNull(),
updatedAt: timestamp('updated_at').notNull().default(sql`now()`), updatedAt: timestamp('updated_at').notNull().default(sql`now()`),
userId: integer('user_id').notNull().references(() => users.id), userId: integer('user_id').notNull().references(() => users.id),
+1 -7
View File
@@ -13,7 +13,6 @@ export interface User {
id: number; id: number;
email: string; email: string;
name: string | null; name: string | null;
profile: string | null;
role: Role; role: Role;
} }
@@ -21,8 +20,7 @@ export const user = pgTable('User', {
id: integer('id').primaryKey().generatedAlwaysAsIdentity(), id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
email: text('email').notNull().unique(), email: text('email').notNull().unique(),
name: text('name'), name: text('name'),
profile: text('profile'), role: role('role').notNull().default('USER'),
role: pgEnum('Role')('role').notNull().default('USER'),
}); });
export type NewUser = typeof user.$inferInsert; export type NewUser = typeof user.$inferInsert;
@@ -30,14 +28,12 @@ export type NewUser = typeof user.$inferInsert;
export interface Profile { export interface Profile {
id: number; id: number;
bio: string; bio: string;
user: string;
userId: number; userId: number;
} }
export const profile = pgTable('Profile', { export const profile = pgTable('Profile', {
id: integer('id').primaryKey().generatedAlwaysAsIdentity(), id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
bio: text('bio').notNull(), bio: text('bio').notNull(),
user: text('user').notNull(),
userId: integer('userId').notNull().unique().references(() => user.id), userId: integer('userId').notNull().unique().references(() => user.id),
}); });
@@ -45,7 +41,6 @@ export type NewProfile = typeof profile.$inferInsert;
// Table: Post // Table: Post
export interface Post { export interface Post {
id: number; id: number;
author: string;
authorId: number; authorId: number;
createdAt: Date; createdAt: Date;
published: boolean; published: boolean;
@@ -55,7 +50,6 @@ export interface Post {
export const post = pgTable('Post', { export const post = pgTable('Post', {
id: integer('id').primaryKey().generatedAlwaysAsIdentity(), id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
author: text('author').notNull(),
authorId: integer('authorId').notNull().references(() => user.id), authorId: integer('authorId').notNull().references(() => user.id),
createdAt: timestamp('createdAt').notNull().default(sql`now()`), createdAt: timestamp('createdAt').notNull().default(sql`now()`),
published: boolean('published').notNull().default(false), published: boolean('published').notNull().default(false),