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:
@@ -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)
|
||||
}
|
||||
|
||||
+73
-75
@@ -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
|
||||
}
|
||||
|
||||
|
||||
+49
-16
@@ -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)))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user