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:
@@ -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) {
|
||||||
|
|||||||
@@ -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"`
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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)))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
Reference in New Issue
Block a user