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
|
||||
}
|
||||
|
||||
func (c *SchemaChange) addChange(field string, source, target any) {
|
||||
if c.Changes == nil {
|
||||
c.Changes = make(map[string]any)
|
||||
}
|
||||
c.Changes[field] = map[string]any{"source": source, "target": target}
|
||||
}
|
||||
|
||||
func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
||||
change := &SchemaChange{
|
||||
Name: source.Name,
|
||||
@@ -78,6 +85,16 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
||||
|
||||
hasChanges := false
|
||||
|
||||
// Compare schema attributes
|
||||
if source.Description != target.Description {
|
||||
change.addChange("description", source.Description, target.Description)
|
||||
hasChanges = true
|
||||
}
|
||||
if source.Owner != target.Owner {
|
||||
change.addChange("owner", source.Owner, target.Owner)
|
||||
hasChanges = true
|
||||
}
|
||||
|
||||
// Compare tables
|
||||
tableDiff := compareTables(source.Tables, target.Tables)
|
||||
if !isEmpty(tableDiff) {
|
||||
|
||||
+6
-5
@@ -18,11 +18,12 @@ type SchemaDiff struct {
|
||||
|
||||
// SchemaChange represents changes within a schema
|
||||
type SchemaChange struct {
|
||||
Name string `json:"name"`
|
||||
Tables *TableDiff `json:"tables,omitempty"`
|
||||
Views *ViewDiff `json:"views,omitempty"`
|
||||
Sequences *SequenceDiff `json:"sequences,omitempty"`
|
||||
Scripts *ScriptDiff `json:"scripts,omitempty"`
|
||||
Name string `json:"name"`
|
||||
Changes map[string]any `json:"changes,omitempty"` // Schema attributes that differ (description, owner), keyed by field name
|
||||
Tables *TableDiff `json:"tables,omitempty"`
|
||||
Views *ViewDiff `json:"views,omitempty"`
|
||||
Sequences *SequenceDiff `json:"sequences,omitempty"`
|
||||
Scripts *ScriptDiff `json:"scripts,omitempty"`
|
||||
}
|
||||
|
||||
// TableDiff represents differences in tables
|
||||
|
||||
@@ -82,6 +82,15 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions)
|
||||
} else {
|
||||
// Schema doesn't exist, add it
|
||||
newSchema := cloneSchema(srcSchema)
|
||||
if len(opts.SkipTableNames) > 0 {
|
||||
kept := newSchema.Tables[:0]
|
||||
for _, t := range newSchema.Tables {
|
||||
if !opts.SkipTableNames[strings.ToLower(t.SQLName())] {
|
||||
kept = append(kept, t)
|
||||
}
|
||||
}
|
||||
newSchema.Tables = kept
|
||||
}
|
||||
target.Schemas = append(target.Schemas, newSchema)
|
||||
r.SchemasAdded++
|
||||
}
|
||||
@@ -440,6 +449,8 @@ func cloneTable(table *models.Table) *models.Table {
|
||||
Description: table.Description,
|
||||
Schema: table.Schema,
|
||||
Comment: table.Comment,
|
||||
Tablespace: table.Tablespace,
|
||||
GUID: table.GUID,
|
||||
Sequence: table.Sequence,
|
||||
UpdatedAt: table.UpdatedAt,
|
||||
Columns: make(map[string]*models.Column),
|
||||
@@ -469,6 +480,14 @@ func cloneTable(table *models.Table) *models.Table {
|
||||
newTable.Indexes[idxName] = cloneIndex(index)
|
||||
}
|
||||
|
||||
// Clone relationships
|
||||
if table.Relationships != nil {
|
||||
newTable.Relationships = make(map[string]*models.Relationship, len(table.Relationships))
|
||||
for relName, rel := range table.Relationships {
|
||||
newTable.Relationships[relName] = cloneRelation(rel)
|
||||
}
|
||||
}
|
||||
|
||||
return newTable
|
||||
}
|
||||
|
||||
|
||||
@@ -15,6 +15,9 @@ import (
|
||||
// Reader implements the readers.Reader interface for Drizzle schema format
|
||||
type Reader struct {
|
||||
options *readers.ReaderOptions
|
||||
// enumVars maps the constant a pgEnum() is assigned to (e.g. "role") to the
|
||||
// enum's SQL name (e.g. "Role"), so columns declared as role('col') resolve.
|
||||
enumVars map[string]string
|
||||
}
|
||||
|
||||
// NewReader creates a new Drizzle reader with the given options
|
||||
@@ -29,6 +32,7 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
if r.options.FilePath == "" {
|
||||
return nil, fmt.Errorf("file path is required for Drizzle reader")
|
||||
}
|
||||
r.enumVars = make(map[string]string)
|
||||
|
||||
// Check if it's a file or directory
|
||||
info, err := os.Stat(r.options.FilePath)
|
||||
@@ -100,6 +104,13 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
|
||||
return nil, fmt.Errorf("failed to glob directory: %w", err)
|
||||
}
|
||||
|
||||
// Enums may be declared in a different file than the tables using them
|
||||
for _, file := range files {
|
||||
if content, err := os.ReadFile(file); err == nil {
|
||||
r.collectEnumVars(string(content))
|
||||
}
|
||||
}
|
||||
|
||||
// Parse each file
|
||||
for _, file := range files {
|
||||
content, err := os.ReadFile(file)
|
||||
@@ -125,9 +136,22 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
|
||||
return db, nil
|
||||
}
|
||||
|
||||
var enumVarRegex = regexp.MustCompile(`export\s+const\s+(\w+)\s*=\s*pgEnum\s*\(\s*['"](\w+)['"]`)
|
||||
|
||||
// collectEnumVars records every pgEnum() constant declared in content.
|
||||
func (r *Reader) collectEnumVars(content string) {
|
||||
if r.enumVars == nil {
|
||||
r.enumVars = make(map[string]string)
|
||||
}
|
||||
for _, m := range enumVarRegex.FindAllStringSubmatch(content, -1) {
|
||||
r.enumVars[m[1]] = m[2]
|
||||
}
|
||||
}
|
||||
|
||||
// parseDrizzle parses Drizzle schema content and returns a Database model
|
||||
func (r *Reader) parseDrizzle(content string) (*models.Database, error) {
|
||||
db := models.InitDatabase("database")
|
||||
r.collectEnumVars(content)
|
||||
|
||||
if r.options.Metadata != nil {
|
||||
if name, ok := r.options.Metadata["name"].(string); ok {
|
||||
@@ -375,6 +399,9 @@ func (r *Reader) parseColumnDefinition(line, fieldName, drizzleType string, tabl
|
||||
|
||||
// Map Drizzle type to SQL type
|
||||
column.Type = r.drizzleTypeToSQL(drizzleType)
|
||||
if enumName, ok := r.enumVars[drizzleType]; ok {
|
||||
column.Type = enumName
|
||||
}
|
||||
|
||||
// Default: columns are nullable unless specified
|
||||
column.NotNull = false
|
||||
|
||||
@@ -13,7 +13,8 @@ import (
|
||||
|
||||
// Reader implements the readers.Reader interface for Prisma schema format
|
||||
type Reader struct {
|
||||
options *readers.ReaderOptions
|
||||
options *readers.ReaderOptions
|
||||
enumNames map[string]bool // enum names declared in the schema being parsed
|
||||
}
|
||||
|
||||
// NewReader creates a new Prisma reader with the given options
|
||||
@@ -82,6 +83,8 @@ func (r *Reader) parsePrisma(content string) (*models.Database, error) {
|
||||
schema := models.InitSchema("public")
|
||||
schema.Enums = make([]*models.Enum, 0)
|
||||
|
||||
r.enumNames = collectEnumNames(content)
|
||||
|
||||
scanner := bufio.NewScanner(strings.NewReader(content))
|
||||
|
||||
// State tracking
|
||||
@@ -600,20 +603,21 @@ func (r *Reader) isPrimitiveType(typeName string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// isEnumType checks if a type name might be an enum
|
||||
// Note: We can't definitively check against schema.Enums at parse time
|
||||
// because enums might be defined after the model, so we just check
|
||||
// if it starts with uppercase (Prisma convention for enums)
|
||||
func (r *Reader) isEnumType(typeName string, table *models.Table) bool {
|
||||
// Simple heuristic: enum types start with uppercase letter
|
||||
// and are not known model names (though we can't check that yet)
|
||||
if len(typeName) > 0 && typeName[0] >= 'A' && typeName[0] <= 'Z' {
|
||||
// Additional check: primitive types are already handled above
|
||||
// So if it's uppercase and not primitive, it's likely an enum or model
|
||||
// We'll assume it's an enum if it's a single word
|
||||
return !strings.Contains(typeName, "_")
|
||||
// isEnumType reports whether typeName is an enum declared in the schema.
|
||||
// Enum names are collected up front because enums may be declared after the
|
||||
// models that use them.
|
||||
func (r *Reader) isEnumType(typeName string, _ *models.Table) bool {
|
||||
return r.enumNames[typeName]
|
||||
}
|
||||
|
||||
var enumDeclRegex = regexp.MustCompile(`(?m)^\s*enum\s+(\w+)\s*{`)
|
||||
|
||||
func collectEnumNames(content string) map[string]bool {
|
||||
names := make(map[string]bool)
|
||||
for _, m := range enumDeclRegex.FindAllStringSubmatch(content, -1) {
|
||||
names[m[1]] = true
|
||||
}
|
||||
return false
|
||||
return names
|
||||
}
|
||||
|
||||
// createConstraintFromRelation creates a FK constraint from a @relation attribute
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -17,7 +17,7 @@ export const users = pgTable('users', {
|
||||
lastLoginAt: timestamp('last_login_at'),
|
||||
passwordHash: varchar('password_hash').notNull(),
|
||||
profile: jsonb('profile'),
|
||||
role: pgEnum('UserRole')('role').notNull(),
|
||||
role: userRole('role').notNull(),
|
||||
updatedAt: timestamp('updated_at').notNull().default(sql`now()`),
|
||||
username: varchar('username').notNull().unique(),
|
||||
});
|
||||
@@ -131,7 +131,7 @@ export const orders = pgTable('orders', {
|
||||
notes: text('notes'),
|
||||
orderNumber: varchar('order_number').notNull().unique(),
|
||||
shippingAddress: jsonb('shipping_address').notNull(),
|
||||
status: pgEnum('OrderStatus')('status').notNull().default('pending'),
|
||||
status: orderStatus('status').notNull().default('pending'),
|
||||
totalAmount: numeric('total_amount').notNull(),
|
||||
updatedAt: timestamp('updated_at').notNull().default(sql`now()`),
|
||||
userId: integer('user_id').notNull().references(() => users.id),
|
||||
|
||||
@@ -13,7 +13,6 @@ export interface User {
|
||||
id: number;
|
||||
email: string;
|
||||
name: string | null;
|
||||
profile: string | null;
|
||||
role: Role;
|
||||
}
|
||||
|
||||
@@ -21,8 +20,7 @@ export const user = pgTable('User', {
|
||||
id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
|
||||
email: text('email').notNull().unique(),
|
||||
name: text('name'),
|
||||
profile: text('profile'),
|
||||
role: pgEnum('Role')('role').notNull().default('USER'),
|
||||
role: role('role').notNull().default('USER'),
|
||||
});
|
||||
|
||||
export type NewUser = typeof user.$inferInsert;
|
||||
@@ -30,14 +28,12 @@ export type NewUser = typeof user.$inferInsert;
|
||||
export interface Profile {
|
||||
id: number;
|
||||
bio: string;
|
||||
user: string;
|
||||
userId: number;
|
||||
}
|
||||
|
||||
export const profile = pgTable('Profile', {
|
||||
id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
|
||||
bio: text('bio').notNull(),
|
||||
user: text('user').notNull(),
|
||||
userId: integer('userId').notNull().unique().references(() => user.id),
|
||||
});
|
||||
|
||||
@@ -45,7 +41,6 @@ export type NewProfile = typeof profile.$inferInsert;
|
||||
// Table: Post
|
||||
export interface Post {
|
||||
id: number;
|
||||
author: string;
|
||||
authorId: number;
|
||||
createdAt: Date;
|
||||
published: boolean;
|
||||
@@ -55,7 +50,6 @@ export interface Post {
|
||||
|
||||
export const post = pgTable('Post', {
|
||||
id: integer('id').primaryKey().generatedAlwaysAsIdentity(),
|
||||
author: text('author').notNull(),
|
||||
authorId: integer('authorId').notNull().references(() => user.id),
|
||||
createdAt: timestamp('createdAt').notNull().default(sql`now()`),
|
||||
published: boolean('published').notNull().default(false),
|
||||
|
||||
Reference in New Issue
Block a user