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:
+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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user