package pgsql import ( "context" "encoding/json" "fmt" "io" "os" "regexp" "sort" "strings" "sync" "time" "git.warky.dev/wdevs/relspecgo/pkg/models" "git.warky.dev/wdevs/relspecgo/pkg/pgsql" "git.warky.dev/wdevs/relspecgo/pkg/writers" ) // Writer implements the Writer interface for PostgreSQL SQL output type Writer struct { options *writers.WriterOptions writer io.Writer executionReport *ExecutionReport executor *TemplateExecutor } // ExecutionReport tracks the execution status of SQL statements type ExecutionReport struct { TotalStatements int `json:"total_statements"` ExecutedStatements int `json:"executed_statements"` FailedStatements int `json:"failed_statements"` Schemas []SchemaReport `json:"schemas"` Errors []ExecutionError `json:"errors,omitempty"` StartTime string `json:"start_time"` EndTime string `json:"end_time"` } // SchemaReport tracks execution per schema type SchemaReport struct { Name string `json:"name"` Tables []TableReport `json:"tables"` } // TableReport tracks execution per table type TableReport struct { Name string `json:"name"` Created bool `json:"created"` Error string `json:"error,omitempty"` } // ExecutionError represents a failed statement type ExecutionError struct { StatementNumber int `json:"statement_number"` Statement string `json:"statement"` Error string `json:"error"` } // NewWriter creates a new PostgreSQL SQL writer func NewWriter(options *writers.WriterOptions) *Writer { executor, _ := NewTemplateExecutor(options.FlattenSchema) return &Writer{ options: options, executor: executor, } } // qualTable returns a schema-qualified name using the writer's FlattenSchema setting. func (w *Writer) qualTable(schema, name string) string { return writers.QualifiedTableName(schema, name, w.options.FlattenSchema) } // WriteDatabase writes the entire database schema as SQL func (w *Writer) WriteDatabase(db *models.Database) error { // Check if we should execute SQL directly on a database if connString, ok := w.options.Metadata["connection_string"].(string); ok && connString != "" { 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 } w.writer = writer // Write header comment fmt.Fprintf(w.writer, "-- PostgreSQL Database Schema\n") fmt.Fprintf(w.writer, "-- Database: %s\n", db.Name) fmt.Fprintf(w.writer, "-- Generated by RelSpec\n") if w.options.ContinueOnError { fmt.Fprintf(w.writer, "\\set ON_ERROR_STOP off\n") } fmt.Fprintf(w.writer, "\n") // Process each schema in the database for _, schema := range db.Schemas { if err := w.WriteSchema(schema); err != nil { return fmt.Errorf("failed to write schema %s: %w", schema.Name, err) } } return nil } // GenerateDatabaseStatements generates SQL statements as a list for the entire database // Returns a slice of SQL statements that can be executed independently func (w *Writer) GenerateDatabaseStatements(db *models.Database) ([]string, error) { statements := []string{} // Add header comment statements = append(statements, "-- PostgreSQL Database Schema") statements = append(statements, fmt.Sprintf("-- Database: %s", db.Name)) statements = append(statements, "-- Generated by RelSpec") // Process each schema in the database for _, schema := range db.Schemas { schemaStatements, err := w.GenerateSchemaStatements(schema) if err != nil { return nil, fmt.Errorf("failed to generate statements for schema %s: %w", schema.Name, err) } statements = append(statements, schemaStatements...) } return statements, nil } // GenerateSchemaStatements generates SQL statements as a list for a single schema func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, error) { statements := []string{} // Phase 1: Create schema (skip entirely when flattening) if schema.Name != "public" && !w.options.FlattenSchema { statements = append(statements, fmt.Sprintf("-- Schema: %s", schema.Name)) statements = append(statements, fmt.Sprintf("CREATE SCHEMA IF NOT EXISTS %s", schema.SQLName())) } for _, extension := range requiredExtensions(schema) { statements = append(statements, fmt.Sprintf("CREATE EXTENSION IF NOT EXISTS %s", pgsql.QuoteExtensionName(extension))) } // Phase 2: Create sequences for _, table := range schema.Tables { pk := table.GetPrimaryKey() if pk == nil || !isIntegerType(pk.Type) || pk.Default == "" { continue } defaultStr, ok := pk.Default.(string) if !ok || !strings.Contains(strings.ToLower(defaultStr), "nextval") { continue } seqName := extractSequenceName(defaultStr) if seqName == "" { continue } stmt := fmt.Sprintf("CREATE SEQUENCE IF NOT EXISTS %s\n INCREMENT 1\n MINVALUE 1\n MAXVALUE 9223372036854775807\n START 1\n CACHE 1", w.qualTable(schema.SQLName(), seqName)) statements = append(statements, stmt) } // Phase 3: Create tables for _, table := range schema.Tables { stmts, err := w.generateCreateTableStatement(schema, table) if err != nil { return nil, fmt.Errorf("failed to generate table %s: %w", table.Name, err) } statements = append(statements, stmts...) } // Phase 3.5: Add missing columns (for existing tables) addColStmts, err := w.GenerateAddColumnStatements(schema) if err != nil { return nil, fmt.Errorf("failed to generate add column statements: %w", err) } statements = append(statements, addColStmts...) alterTypeStmts, err := w.GenerateAlterColumnTypeStatements(schema) if err != nil { return nil, fmt.Errorf("failed to generate alter column type statements: %w", err) } statements = append(statements, alterTypeStmts...) // Phase 4: Primary keys for _, table := range schema.Tables { // First check for explicit PrimaryKeyConstraint var pkConstraint *models.Constraint for _, constraint := range sortConstraints(table.Constraints) { if constraint.Type == models.PrimaryKeyConstraint { pkConstraint = constraint break } } var pkColumns []string var pkName string if pkConstraint != nil { pkColumns = pkConstraint.Columns pkName = pkConstraint.Name } else { // No explicit constraint, check for columns with IsPrimaryKey = true pkCols := []string{} for _, col := range table.Columns { if col.IsPrimaryKey { pkCols = append(pkCols, col.SQLName()) } } if len(pkCols) > 0 { // Sort for consistent output sort.Strings(pkCols) pkColumns = pkCols pkName = fmt.Sprintf("pk_%s_%s", schema.SQLName(), table.SQLName()) } } if len(pkColumns) > 0 { // Auto-generated primary key names to check for and drop autoGenPKNames := []string{ fmt.Sprintf("%s_pkey", table.Name), fmt.Sprintf("%s_%s_pkey", schema.Name, table.Name), } // Use template to generate primary key statement data := CreatePrimaryKeyWithAutoGenCheckData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: pkName, AutoGenNames: formatStringList(autoGenPKNames), Columns: strings.Join(pkColumns, ", "), ColumnNames: formatStringList(pkColumns), } stmt, err := w.executor.ExecuteCreatePrimaryKeyWithAutoGenCheck(data) if err != nil { return nil, fmt.Errorf("failed to generate primary key for %s.%s: %w", schema.Name, table.Name, err) } statements = append(statements, stmt) } } // Phase 5: Indexes for _, table := range schema.Tables { for _, index := range sortIndexes(table.Indexes) { // Skip primary key indexes if strings.HasSuffix(index.Name, "_pkey") { continue } uniqueStr := "" if index.Unique { uniqueStr = "UNIQUE " } indexType := index.Type if indexType == "" { indexType = "btree" } // Build column expressions with operator class support (GIN, pgvector, PostGIS) columnExprs := buildIndexColumnExpressions(table, index, indexType) withClause := "" if params := indexStorageParameters(index.Comment); params != "" { withClause = fmt.Sprintf(" WITH (%s)", params) } whereClause := "" if index.Where != "" { whereClause = fmt.Sprintf(" WHERE %s", index.Where) } stmt := fmt.Sprintf("CREATE %sINDEX IF NOT EXISTS %s ON %s USING %s (%s)%s%s", uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause) statements = append(statements, stmt) } } // Phase 5.5: Unique constraints for _, table := range schema.Tables { for _, constraint := range sortConstraints(table.Constraints) { if constraint.Type != models.UniqueConstraint { continue } // Use template to generate unique constraint statement data := CreateUniqueConstraintData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: constraint.Name, Columns: strings.Join(constraint.Columns, ", "), } stmt, err := w.executor.ExecuteCreateUniqueConstraint(data) if err != nil { return nil, fmt.Errorf("failed to generate unique constraint for %s.%s: %w", schema.Name, table.Name, err) } statements = append(statements, stmt) } } // Phase 5.7: Check constraints for _, table := range schema.Tables { for _, constraint := range sortConstraints(table.Constraints) { if constraint.Type != models.CheckConstraint { continue } // Use template to generate check constraint statement data := CreateCheckConstraintData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: constraint.Name, Expression: constraint.Expression, } stmt, err := w.executor.ExecuteCreateCheckConstraint(data) if err != nil { return nil, fmt.Errorf("failed to generate check constraint for %s.%s: %w", schema.Name, table.Name, err) } statements = append(statements, stmt) } } // Phase 6: Foreign keys for _, table := range schema.Tables { for _, constraint := range sortConstraints(table.Constraints) { if constraint.Type != models.ForeignKeyConstraint { continue } refSchema := constraint.ReferencedSchema if refSchema == "" { refSchema = schema.Name } onDelete := constraint.OnDelete if onDelete == "" { onDelete = "NO ACTION" } onUpdate := constraint.OnUpdate if onUpdate == "" { onUpdate = "NO ACTION" } // Use template to generate foreign key statement data := CreateForeignKeyWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: constraint.Name, SourceColumns: strings.Join(constraint.Columns, ", "), TargetSchema: refSchema, TargetTable: constraint.ReferencedTable, TargetColumns: strings.Join(constraint.ReferencedColumns, ", "), OnDelete: onDelete, OnUpdate: onUpdate, Deferrable: false, } stmt, err := w.executor.ExecuteCreateForeignKeyWithCheck(data) if err != nil { return nil, fmt.Errorf("failed to generate foreign key for %s.%s: %w", schema.Name, table.Name, err) } statements = append(statements, stmt) } } // Phase 7: Comments for _, table := range schema.Tables { if table.Comment != "" { stmt := fmt.Sprintf("COMMENT ON TABLE %s IS '%s'", w.qualTable(schema.SQLName(), table.SQLName()), escapeQuote(table.Comment)) statements = append(statements, stmt) } for _, column := range sortColumns(table.Columns) { if column.Comment != "" { stmt := fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'", w.qualTable(schema.SQLName(), table.SQLName()), column.SQLName(), escapeQuote(column.Comment)) statements = append(statements, stmt) } } } return statements, nil } // GenerateAddColumnStatements generates ALTER TABLE ADD COLUMN statements for existing tables // This is useful for schema evolution when new columns are added to existing tables func (w *Writer) GenerateAddColumnStatements(schema *models.Schema) ([]string, error) { statements := []string{} statements = append(statements, fmt.Sprintf("-- Add missing columns for schema: %s", schema.Name)) for _, table := range schema.Tables { // Sort columns by sequence or name for consistent output columns := make([]*models.Column, 0, len(table.Columns)) for _, col := range table.Columns { columns = append(columns, col) } sort.Slice(columns, func(i, j int) bool { if columns[i].Sequence != columns[j].Sequence { return columns[i].Sequence < columns[j].Sequence } return columns[i].Name < columns[j].Name }) for _, col := range columns { colDef := w.generateColumnDefinition(col) // Use template to generate add column statement data := AddColumnWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ColumnName: col.Name, ColumnDefinition: colDef, } stmt, err := w.executor.ExecuteAddColumnWithCheck(data) if err != nil { return nil, fmt.Errorf("failed to generate add column for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err) } statements = append(statements, stmt) } } return statements, nil } func (w *Writer) GenerateAlterColumnTypeStatements(schema *models.Schema) ([]string, error) { statements := []string{} statements = append(statements, fmt.Sprintf("-- Alter column types for schema: %s", schema.Name)) for _, table := range schema.Tables { columns := getSortedColumns(table.Columns) for _, col := range columns { targetType := effectiveAlterColumnSQLType(col) stmt, err := w.executor.ExecuteAlterColumnTypeWithCheck(AlterColumnTypeWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ColumnName: col.Name, NewType: targetType, EquivalentTypes: equivalentTypeListSQL(targetType), UsingExpr: buildAlterColumnUsingExpression(col.Name, targetType), }) if err != nil { return nil, fmt.Errorf("failed to generate alter column type for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err) } statements = append(statements, stmt) } } return statements, nil } // GenerateAlterColumnDefaultStatements generates guarded ALTER TABLE // statements to bring existing columns' DEFAULT clause in line with the // model, safe to run against a database that already has the columns. func (w *Writer) GenerateAlterColumnDefaultStatements(schema *models.Schema) ([]string, error) { statements := []string{} statements = append(statements, fmt.Sprintf("-- Alter column defaults for schema: %s", schema.Name)) for _, table := range schema.Tables { columns := getSortedColumns(table.Columns) for _, col := range columns { setDefault, defaultVal := formatColumnDefaultSQL(col) stmt, err := w.executor.ExecuteAlterColumnDefaultWithCheck(AlterColumnDefaultWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ColumnName: col.Name, SetDefault: setDefault, DefaultValue: defaultVal, }) if err != nil { return nil, fmt.Errorf("failed to generate alter column default for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err) } statements = append(statements, stmt) } } return statements, nil } // formatColumnDefaultSQL renders a column's model-level default into the // SQL literal/expression used by ALTER COLUMN ... SET DEFAULT, shared by // the full-schema writer and the diff-based migration writer. func formatColumnDefaultSQL(col *models.Column) (setDefault bool, defaultVal string) { if col.Default == nil { return false, "" } if value, ok := col.Default.(string); ok { return true, writers.QuoteDefaultValue(stripBackticks(value), col.Type) } return true, fmt.Sprintf("%v", col.Default) } // GenerateAlterColumnNullabilityStatements generates guarded ALTER TABLE // statements to bring existing columns' NOT NULL state in line with the // model, safe to run against a database that already has the columns. func (w *Writer) GenerateAlterColumnNullabilityStatements(schema *models.Schema) ([]string, error) { statements := []string{} statements = append(statements, fmt.Sprintf("-- Alter column nullability for schema: %s", schema.Name)) for _, table := range schema.Tables { columns := getSortedColumns(table.Columns) for _, col := range columns { stmt, err := w.executor.ExecuteAlterColumnNullabilityWithCheck(AlterColumnNullabilityWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ColumnName: col.Name, NotNull: col.NotNull, }) if err != nil { return nil, fmt.Errorf("failed to generate alter column nullability for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err) } statements = append(statements, stmt) } } return statements, nil } // GenerateAddColumnsForDatabase generates ALTER TABLE ADD COLUMN statements for the entire database func (w *Writer) GenerateAddColumnsForDatabase(db *models.Database) ([]string, error) { statements := []string{} statements = append(statements, "-- Add missing columns to existing tables") statements = append(statements, fmt.Sprintf("-- Database: %s", db.Name)) statements = append(statements, "-- Generated by RelSpec") for _, schema := range db.Schemas { schemaStatements, err := w.GenerateAddColumnStatements(schema) if err != nil { return nil, fmt.Errorf("failed to generate add column statements for schema %s: %w", schema.Name, err) } statements = append(statements, schemaStatements...) } return statements, nil } // generateCreateTableStatement generates CREATE TABLE statement func (w *Writer) generateCreateTableStatement(schema *models.Schema, table *models.Table) ([]string, error) { statements := []string{} // Sort columns by sequence or name columns := make([]*models.Column, 0, len(table.Columns)) for _, col := range table.Columns { columns = append(columns, col) } sort.Slice(columns, func(i, j int) bool { if columns[i].Sequence != columns[j].Sequence { return columns[i].Sequence < columns[j].Sequence } return columns[i].Name < columns[j].Name }) columnDefs := []string{} for _, col := range columns { def := w.generateColumnDefinition(col) columnDefs = append(columnDefs, " "+def) } stmt := fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (\n%s\n)", w.qualTable(schema.SQLName(), table.SQLName()), strings.Join(columnDefs, ",\n")) statements = append(statements, stmt) return statements, nil } // generateColumnDefinition generates column definition func (w *Writer) generateColumnDefinition(col *models.Column) string { parts := []string{col.SQLName()} parts = append(parts, effectiveColumnSQLType(col)) // NOT NULL if col.NotNull { parts = append(parts, "NOT NULL") } // DEFAULT if col.Default != nil { switch v := col.Default.(type) { case string: parts = append(parts, fmt.Sprintf("DEFAULT %s", writers.QuoteDefaultValue(stripBackticks(v), col.Type))) case bool: parts = append(parts, fmt.Sprintf("DEFAULT %v", v)) default: parts = append(parts, fmt.Sprintf("DEFAULT %v", v)) } } return strings.Join(parts, " ") } func effectiveColumnSQLType(col *models.Column) string { if col == nil { return "" } baseType := pgsql.ConvertSQLType(col.Type) typeStr := baseType hasExplicitTypeModifier := pgsql.HasExplicitTypeModifier(baseType) if !hasExplicitTypeModifier && col.Length > 0 && col.Precision == 0 { if pgsql.SupportsLength(baseType) { typeStr = fmt.Sprintf("%s(%d)", baseType, col.Length) } else if isTextTypeWithoutLength(baseType) { typeStr = fmt.Sprintf("varchar(%d)", col.Length) } } else if !hasExplicitTypeModifier && col.Precision > 0 { if pgsql.SupportsPrecision(baseType) { if col.Scale > 0 { typeStr = fmt.Sprintf("%s(%d,%d)", baseType, col.Precision, col.Scale) } else { typeStr = fmt.Sprintf("%s(%d)", baseType, col.Precision) } } } return typeStr } func effectiveAlterColumnSQLType(col *models.Column) string { typeStr := effectiveColumnSQLType(col) switch strings.ToLower(strings.TrimSpace(typeStr)) { case "smallserial": return "smallint" case "serial": return "integer" case "bigserial": return "bigint" default: return typeStr } } func buildAlterColumnUsingExpression(columnName, targetType string) string { if strings.TrimSpace(columnName) == "" || strings.TrimSpace(targetType) == "" { return "" } return fmt.Sprintf("%s::%s", quoteIdent(columnName), targetType) } func equivalentTypeListSQL(sqlType string) string { variants := pgsql.EquivalentSQLTypeVariants(sqlType) quoted := make([]string, 0, len(variants)) for _, variant := range variants { quoted = append(quoted, fmt.Sprintf("'%s'", escapeQuote(variant))) } return strings.Join(quoted, ", ") } // 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 } // Phase 1: Create schema (priority 1) if err := w.writeCreateSchema(schema); err != nil { return err } if err := w.writeRequiredExtensions(schema); err != nil { return err } // Phase 2: Create sequences (priority 80) if err := w.writeSequences(schema); err != nil { return err } // Phase 3: Create tables with columns (priority 100) if err := w.writeCreateTables(schema); err != nil { return err } // Phase 3.5: Add missing columns (priority 120) if err := w.writeAddColumns(schema); err != nil { return err } if err := w.writeAlterColumnTypes(schema); err != nil { return err } if err := w.writeAlterColumnDefaults(schema); err != nil { return err } if err := w.writeAlterColumnNullability(schema); err != nil { return err } // Phase 4: Create primary keys (priority 160) if err := w.writePrimaryKeys(schema); err != nil { return err } // Phase 5: Create indexes (priority 180) if err := w.writeIndexes(schema); err != nil { return err } // Phase 5.5: Create unique constraints (priority 185) if err := w.writeUniqueConstraints(schema); err != nil { return err } // Phase 5.7: Create check constraints (priority 190) if err := w.writeCheckConstraints(schema); err != nil { return err } // Phase 6: Create foreign key constraints (priority 195) if err := w.writeForeignKeys(schema); err != nil { return err } // Phase 7: Set sequence values (priority 200) if err := w.writeSetSequenceValues(schema); err != nil { return err } // Phase 8: Add comments (priority 200+) if err := w.writeComments(schema); err != nil { return err } return nil } // 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 } // Create a temporary schema with just this table schema := models.InitSchema(table.Schema) schema.Tables = append(schema.Tables, table) return w.WriteSchema(schema) } // WriteAddColumnStatements writes ALTER TABLE ADD COLUMN statements for a database // This is used for schema evolution/migration when new columns are added func (w *Writer) WriteAddColumnStatements(db *models.Database) error { 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 } w.writer = writer // Generate statements statements, err := w.GenerateAddColumnsForDatabase(db) if err != nil { return err } // Write each statement for _, stmt := range statements { fmt.Fprintf(w.writer, "%s;\n\n", stmt) } return nil } // writeCreateSchema generates CREATE SCHEMA statement func (w *Writer) writeCreateSchema(schema *models.Schema) error { if schema.Name == "public" || w.options.FlattenSchema { return nil } fmt.Fprintf(w.writer, "-- Schema: %s\n", schema.Name) fmt.Fprintf(w.writer, "CREATE SCHEMA IF NOT EXISTS %s;\n\n", schema.SQLName()) return nil } func (w *Writer) writeRequiredExtensions(schema *models.Schema) error { extensions := requiredExtensions(schema) if len(extensions) == 0 { return nil } for _, extension := range extensions { fmt.Fprintf(w.writer, "CREATE EXTENSION IF NOT EXISTS %s;\n", pgsql.QuoteExtensionName(extension)) } fmt.Fprintln(w.writer) return nil } // writeSequences generates CREATE SEQUENCE statements for identity columns func (w *Writer) writeSequences(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Sequences for schema: %s\n", schema.Name) for _, table := range schema.Tables { pk := table.GetPrimaryKey() if pk == nil { continue } // Only create sequences for integer-type PKs with identity if !isIntegerType(pk.Type) { continue } seqName := fmt.Sprintf("identity_%s_%s", table.SQLName(), pk.SQLName()) data := CreateSequenceData{ SchemaName: schema.Name, SequenceName: seqName, Increment: 1, MinValue: 1, MaxValue: 9223372036854775807, StartValue: 1, CacheSize: 1, } sql, err := w.executor.ExecuteCreateSequence(data) if err != nil { return fmt.Errorf("failed to generate create sequence for %s.%s: %w", schema.Name, seqName, err) } fmt.Fprint(w.writer, sql) fmt.Fprint(w.writer, "\n") } return nil } // writeCreateTables generates CREATE TABLE statements func (w *Writer) writeCreateTables(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Tables for schema: %s\n", schema.Name) for _, table := range schema.Tables { fmt.Fprintf(w.writer, "CREATE TABLE IF NOT EXISTS %s (\n", w.qualTable(schema.SQLName(), table.SQLName())) // Write columns columns := getSortedColumns(table.Columns) columnDefs := make([]string, 0, len(columns)) for _, col := range columns { // Use generateColumnDefinition to properly handle type, length, precision, and defaults colDef := " " + w.generateColumnDefinition(col) columnDefs = append(columnDefs, colDef) } fmt.Fprintf(w.writer, "%s\n", strings.Join(columnDefs, ",\n")) fmt.Fprintf(w.writer, ");\n\n") } return nil } // writeAddColumns generates ALTER TABLE ADD COLUMN statements for missing columns func (w *Writer) writeAddColumns(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Add missing columns for schema: %s\n", schema.Name) for _, table := range schema.Tables { // Sort columns by sequence or name for consistent output columns := getSortedColumns(table.Columns) for _, col := range columns { colDef := w.generateColumnDefinition(col) data := AddColumnWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ColumnName: col.Name, ColumnDefinition: colDef, } sql, err := w.executor.ExecuteAddColumnWithCheck(data) if err != nil { return fmt.Errorf("failed to generate add column for %s.%s.%s: %w", schema.Name, table.Name, col.Name, err) } fmt.Fprint(w.writer, sql) fmt.Fprint(w.writer, "\n") } } return nil } func (w *Writer) writeAlterColumnTypes(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Alter column types for schema: %s\n", schema.Name) statements, err := w.GenerateAlterColumnTypeStatements(schema) if err != nil { return err } for _, stmt := range statements[1:] { fmt.Fprint(w.writer, stmt) fmt.Fprint(w.writer, "\n") } return nil } func (w *Writer) writeAlterColumnDefaults(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Alter column defaults for schema: %s\n", schema.Name) statements, err := w.GenerateAlterColumnDefaultStatements(schema) if err != nil { return err } for _, stmt := range statements[1:] { fmt.Fprint(w.writer, stmt) fmt.Fprint(w.writer, "\n") } return nil } func (w *Writer) writeAlterColumnNullability(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Alter column nullability for schema: %s\n", schema.Name) statements, err := w.GenerateAlterColumnNullabilityStatements(schema) if err != nil { return err } for _, stmt := range statements[1:] { fmt.Fprint(w.writer, stmt) fmt.Fprint(w.writer, "\n") } return nil } // writePrimaryKeys generates ALTER TABLE statements for primary keys func (w *Writer) writePrimaryKeys(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Primary keys for schema: %s\n", schema.Name) for _, table := range schema.Tables { // Find primary key constraint var pkConstraint *models.Constraint for _, constraint := range sortConstraints(table.Constraints) { if constraint.Type == models.PrimaryKeyConstraint { pkConstraint = constraint break } } var columnNames []string pkName := fmt.Sprintf("pk_%s_%s", schema.SQLName(), table.SQLName()) if pkConstraint != nil { // Build column list from explicit constraint columnNames = make([]string, 0, len(pkConstraint.Columns)) for _, colName := range pkConstraint.Columns { if col, ok := table.Columns[colName]; ok { columnNames = append(columnNames, col.SQLName()) } } } else { // No explicit PK constraint, check for columns with IsPrimaryKey = true for _, col := range table.Columns { if col.IsPrimaryKey { columnNames = append(columnNames, col.SQLName()) } } // Sort for consistent output sort.Strings(columnNames) } if len(columnNames) == 0 { continue } // Auto-generated primary key names to check for and drop autoGenPKNames := []string{ fmt.Sprintf("%s_pkey", table.Name), fmt.Sprintf("%s_%s_pkey", schema.Name, table.Name), } data := CreatePrimaryKeyWithAutoGenCheckData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: pkName, AutoGenNames: formatStringList(autoGenPKNames), Columns: strings.Join(columnNames, ", "), ColumnNames: formatStringList(columnNames), } sql, err := w.executor.ExecuteCreatePrimaryKeyWithAutoGenCheck(data) if err != nil { return fmt.Errorf("failed to generate primary key for %s.%s: %w", schema.Name, table.Name, err) } fmt.Fprint(w.writer, sql) fmt.Fprint(w.writer, "\n") } return nil } // writeIndexes generates CREATE INDEX statements func (w *Writer) writeIndexes(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Indexes for schema: %s\n", schema.Name) for _, table := range schema.Tables { // Sort indexes by name for consistent output indexNames := make([]string, 0, len(table.Indexes)) for name := range table.Indexes { indexNames = append(indexNames, name) } sort.Strings(indexNames) for _, name := range indexNames { index := table.Indexes[name] // Skip if it's a primary key index (based on name convention or columns) // Primary keys are handled separately if strings.HasPrefix(strings.ToLower(index.Name), "pk_") { continue } indexName := index.Name if indexName == "" { indexType := "idx" if index.Unique { indexType = "uidx" } columnSuffix := strings.Join(index.Columns, "_") indexName = fmt.Sprintf("%s_%s_%s", indexType, table.SQLName(), strings.ToLower(columnSuffix)) } indexType := index.Type if indexType == "" { indexType = "btree" } // Build column list with operator class support (GIN, pgvector, PostGIS) columnExprs := buildIndexColumnExpressionsFiltered(table, index, indexType, true) if len(columnExprs) == 0 { continue } unique := "" if index.Unique { unique = "UNIQUE " } withClause := "" if params := indexStorageParameters(index.Comment); params != "" { withClause = fmt.Sprintf(" WITH (%s)", params) } whereClause := "" if index.Where != "" { whereClause = fmt.Sprintf(" WHERE %s", index.Where) } concurrently := "" if index.Concurrent { concurrently = "CONCURRENTLY " } fmt.Fprintf(w.writer, "CREATE %sINDEX %sIF NOT EXISTS %s\n", unique, concurrently, indexName) fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s%s;\n\n", w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), withClause, whereClause) } } return nil } // writeUniqueConstraints generates ALTER TABLE statements for unique constraints func (w *Writer) writeUniqueConstraints(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Unique constraints for schema: %s\n", schema.Name) for _, table := range schema.Tables { // Sort constraints by name for consistent output constraintNames := make([]string, 0, len(table.Constraints)) for name, constraint := range table.Constraints { if constraint.Type == models.UniqueConstraint { constraintNames = append(constraintNames, name) } } sort.Strings(constraintNames) for _, name := range constraintNames { constraint := table.Constraints[name] // Build column list columnExprs := make([]string, 0, len(constraint.Columns)) for _, colName := range constraint.Columns { if col, ok := table.Columns[colName]; ok { columnExprs = append(columnExprs, col.SQLName()) } } if len(columnExprs) == 0 { continue } sql, err := w.executor.ExecuteCreateUniqueConstraint(CreateUniqueConstraintData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: constraint.Name, Columns: strings.Join(columnExprs, ", "), }) if err != nil { return fmt.Errorf("failed to generate unique constraint: %w", err) } fmt.Fprintf(w.writer, "%s\n\n", sql) } } return nil } // writeCheckConstraints generates ALTER TABLE statements for check constraints func (w *Writer) writeCheckConstraints(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Check constraints for schema: %s\n", schema.Name) for _, table := range schema.Tables { // Sort constraints by name for consistent output constraintNames := make([]string, 0, len(table.Constraints)) for name, constraint := range table.Constraints { if constraint.Type == models.CheckConstraint { constraintNames = append(constraintNames, name) } } sort.Strings(constraintNames) for _, name := range constraintNames { constraint := table.Constraints[name] // Skip if expression is empty if constraint.Expression == "" { continue } sql, err := w.executor.ExecuteCreateCheckConstraint(CreateCheckConstraintData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: constraint.Name, Expression: constraint.Expression, }) if err != nil { return fmt.Errorf("failed to generate check constraint: %w", err) } fmt.Fprintf(w.writer, "%s\n\n", sql) } } return nil } // writeForeignKeys generates ALTER TABLE statements for foreign keys func (w *Writer) writeForeignKeys(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Foreign keys for schema: %s\n", schema.Name) for _, table := range schema.Tables { // Sort relationships by name for consistent output relNames := make([]string, 0, len(table.Relationships)) for name := range table.Relationships { relNames = append(relNames, name) } sort.Strings(relNames) for _, name := range relNames { rel := table.Relationships[name] // For relationships, we need to look up the foreign key constraint // that defines the actual column mappings fkName := rel.ForeignKey if fkName == "" { fkName = name } if fkName == "" { fkName = fmt.Sprintf("fk_%s_%s", table.SQLName(), rel.ToTable) } // Find the foreign key constraint that matches this relationship var fkConstraint *models.Constraint for _, constraint := range table.Constraints { if constraint.Type == models.ForeignKeyConstraint && (constraint.Name == fkName || constraint.ReferencedTable == rel.ToTable) { fkConstraint = constraint break } } // If no constraint found, skip this relationship if fkConstraint == nil { continue } // Build column lists from the constraint sourceColumns := make([]string, 0, len(fkConstraint.Columns)) for _, colName := range fkConstraint.Columns { if col, ok := table.Columns[colName]; ok { sourceColumns = append(sourceColumns, col.SQLName()) } } targetColumns := make([]string, 0, len(fkConstraint.ReferencedColumns)) for _, colName := range fkConstraint.ReferencedColumns { targetColumns = append(targetColumns, strings.ToLower(colName)) } if len(sourceColumns) == 0 || len(targetColumns) == 0 { continue } onDelete := "NO ACTION" if fkConstraint.OnDelete != "" { onDelete = strings.ToUpper(fkConstraint.OnDelete) } onUpdate := "NO ACTION" if fkConstraint.OnUpdate != "" { onUpdate = strings.ToUpper(fkConstraint.OnUpdate) } // Use constraint's referenced schema/table or relationship's ToSchema/ToTable refSchema := fkConstraint.ReferencedSchema if refSchema == "" { refSchema = rel.ToSchema } refTable := fkConstraint.ReferencedTable if refTable == "" { refTable = rel.ToTable } // Use template executor to generate foreign key with existence check data := CreateForeignKeyWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: fkName, SourceColumns: strings.Join(sourceColumns, ", "), TargetSchema: refSchema, TargetTable: refTable, TargetColumns: strings.Join(targetColumns, ", "), OnDelete: onDelete, OnUpdate: onUpdate, Deferrable: true, } sql, err := w.executor.ExecuteCreateForeignKeyWithCheck(data) if err != nil { return fmt.Errorf("failed to generate foreign key for %s.%s: %w", schema.Name, table.Name, err) } fmt.Fprint(w.writer, sql) } // Also process any foreign key constraints that don't have a relationship processedConstraints := make(map[string]bool) for _, rel := range table.Relationships { fkName := rel.ForeignKey if fkName == "" { fkName = rel.Name } if fkName != "" { processedConstraints[fkName] = true } } // Find unprocessed foreign key constraints constraintNames := make([]string, 0) for name, constraint := range table.Constraints { if constraint.Type == models.ForeignKeyConstraint && !processedConstraints[name] { constraintNames = append(constraintNames, name) } } sort.Strings(constraintNames) for _, name := range constraintNames { constraint := table.Constraints[name] // Build column lists sourceColumns := make([]string, 0, len(constraint.Columns)) for _, colName := range constraint.Columns { if col, ok := table.Columns[colName]; ok { sourceColumns = append(sourceColumns, col.SQLName()) } else { sourceColumns = append(sourceColumns, colName) } } targetColumns := make([]string, 0, len(constraint.ReferencedColumns)) for _, colName := range constraint.ReferencedColumns { targetColumns = append(targetColumns, strings.ToLower(colName)) } if len(sourceColumns) == 0 || len(targetColumns) == 0 { continue } onDelete := "NO ACTION" if constraint.OnDelete != "" { onDelete = strings.ToUpper(constraint.OnDelete) } onUpdate := "NO ACTION" if constraint.OnUpdate != "" { onUpdate = strings.ToUpper(constraint.OnUpdate) } refSchema := constraint.ReferencedSchema if refSchema == "" { refSchema = schema.Name } refTable := constraint.ReferencedTable // Use template executor to generate foreign key with existence check data := CreateForeignKeyWithCheckData{ SchemaName: schema.Name, TableName: table.Name, ConstraintName: constraint.Name, SourceColumns: strings.Join(sourceColumns, ", "), TargetSchema: refSchema, TargetTable: refTable, TargetColumns: strings.Join(targetColumns, ", "), OnDelete: onDelete, OnUpdate: onUpdate, Deferrable: false, } sql, err := w.executor.ExecuteCreateForeignKeyWithCheck(data) if err != nil { return fmt.Errorf("failed to generate foreign key for %s.%s: %w", schema.Name, table.Name, err) } fmt.Fprint(w.writer, sql) } } return nil } // writeSetSequenceValues generates statements to set sequence current values func (w *Writer) writeSetSequenceValues(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Set sequence values for schema: %s\n", schema.Name) for _, table := range schema.Tables { pk := table.GetPrimaryKey() if pk == nil || !isIntegerType(pk.Type) { continue } seqName := fmt.Sprintf("identity_%s_%s", table.SQLName(), pk.SQLName()) // Use template executor to generate set sequence value statement data := SetSequenceValueData{ SchemaName: schema.Name, TableName: table.Name, SequenceName: seqName, ColumnName: pk.Name, } sql, err := w.executor.ExecuteSetSequenceValue(data) if err != nil { return fmt.Errorf("failed to generate set sequence value for %s.%s: %w", schema.Name, table.Name, err) } fmt.Fprint(w.writer, sql) fmt.Fprint(w.writer, "\n") } return nil } // writeComments generates COMMENT statements for tables and columns func (w *Writer) writeComments(schema *models.Schema) error { fmt.Fprintf(w.writer, "-- Comments for schema: %s\n", schema.Name) for _, table := range schema.Tables { // Table comment if table.Description != "" { fmt.Fprintf(w.writer, "COMMENT ON TABLE %s IS '%s';\n", w.qualTable(schema.SQLName(), table.SQLName()), escapeQuote(table.Description)) } // Column comments for _, col := range getSortedColumns(table.Columns) { if col.Description != "" { fmt.Fprintf(w.writer, "COMMENT ON COLUMN %s.%s IS '%s';\n", w.qualTable(schema.SQLName(), table.SQLName()), col.SQLName(), escapeQuote(col.Description)) } } fmt.Fprintf(w.writer, "\n") } return nil } // Helper functions // getSortedColumns returns columns sorted by name 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]) } return sorted } // isIntegerType checks if a column type is an integer type func isIntegerType(colType string) bool { intTypes := []string{"integer", "int", "bigint", "smallint", "serial", "bigserial"} lowerType := strings.ToLower(colType) for _, t := range intTypes { if strings.HasPrefix(lowerType, t) { return true } } return false } // isTextType checks if a column type is a text type (for GIN index operator class) // func isTextType(colType string) bool { // textTypes := []string{"text", "varchar", "character varying", "char", "character", "string"} // lowerType := strings.ToLower(colType) // if strings.HasSuffix(lowerType, "[]") { // return false // } // for _, t := range textTypes { // if strings.HasPrefix(lowerType, t) { // return true // } // } // return false // } // isTextTypeWithoutLength checks if type is text (which should convert to varchar when length is specified) func isTextTypeWithoutLength(colType string) bool { return strings.EqualFold(colType, "text") } // vectorOperatorClasses maps pgvector operator classes to the column base type they // apply to. pgvector defines no default operator class, so an hnsw/ivfflat index must // always name one explicitly. var vectorOperatorClasses = map[string]string{ "vector_l2_ops": "vector", "vector_ip_ops": "vector", "vector_cosine_ops": "vector", "vector_l1_ops": "vector", "halfvec_l2_ops": "halfvec", "halfvec_ip_ops": "halfvec", "halfvec_cosine_ops": "halfvec", "halfvec_l1_ops": "halfvec", "sparsevec_l2_ops": "sparsevec", "sparsevec_ip_ops": "sparsevec", "sparsevec_cosine_ops": "sparsevec", "sparsevec_l1_ops": "sparsevec", "bit_hamming_ops": "bit", "bit_jaccard_ops": "bit", } // defaultVectorOperatorClasses is the operator class used for an hnsw/ivfflat index when // the index comment does not request one. Cosine distance is the common default for // embedding columns; override it with an "opclass" hint in the index comment. var defaultVectorOperatorClasses = map[string]string{ "vector": "vector_cosine_ops", "halfvec": "halfvec_cosine_ops", "sparsevec": "sparsevec_cosine_ops", "bit": "bit_hamming_ops", } // spatialOperatorClasses are the PostGIS operator classes recognized in index comments. // PostGIS installs default operator classes for gist/spgist/brin, so these are only // emitted when explicitly requested (e.g. the 3D/nD variants). var spatialOperatorClasses = map[string]bool{ "gist_geometry_ops_2d": true, "gist_geometry_ops_nd": true, "gist_geography_ops": true, "spgist_geometry_ops_2d": true, "spgist_geometry_ops_3d": true, "spgist_geometry_ops_nd": true, "brin_geometry_inclusion_ops_2d": true, "brin_geometry_inclusion_ops_3d": true, "brin_geometry_inclusion_ops_4d": true, "brin_geography_inclusion_ops_2d": true, "btree_geometry_ops": true, "btree_geography_ops": true, } // isVectorIndexMethod reports whether the access method indexes pgvector types, which // covers both pgvector itself (hnsw, ivfflat) and VectorChord (vchordrq, vchordg). func isVectorIndexMethod(method string) bool { switch strings.ToLower(strings.TrimSpace(method)) { case "hnsw", "ivfflat", "vchordrq", "vchordg": return true default: return false } } // indexOperatorClassForColumn returns the operator class to emit for a column in an index // of the given access method, honouring an explicit request from the index comment when it // is compatible with the column type. func indexOperatorClassForColumn(col *models.Column, indexType, comment string) string { if col == nil { return "" } sqlType := effectiveColumnSQLType(col) baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType)) isArray := pgsql.IsArrayType(sqlType) requested := extractOperatorClass(comment) method := strings.ToLower(strings.TrimSpace(indexType)) if method == "" { method = "btree" } if requested != "" && operatorClassCompatible(method, baseType, isArray, requested) { return requested } switch { case method == "gin": if isArray { return "array_ops" } switch { case isTextGinBaseType(baseType): return "gin_trgm_ops" case baseType == "jsonb": return "jsonb_ops" default: return requested } case isVectorIndexMethod(method): if isArray { return "" } return defaultVectorOperatorClasses[baseType] default: // gist/spgist/brin/btree have default operator classes (PostGIS included), // so nothing is emitted unless the comment requested a compatible class. return "" } } // ginOperatorClassForColumn is the GIN-specific form of indexOperatorClassForColumn. func ginOperatorClassForColumn(col *models.Column, comment string) string { return indexOperatorClassForColumn(col, "gin", comment) } func operatorClassCompatible(method, baseType string, isArray bool, opClass string) bool { if vectorType, ok := vectorOperatorClasses[opClass]; ok { return !isArray && baseType == vectorType && isVectorIndexMethod(method) } if spatialOperatorClasses[opClass] { return !isArray && pgsql.IsSpatialType(baseType) } switch opClass { case "gin_trgm_ops", "gin_bigm_ops": return !isArray && isTextGinBaseType(baseType) case "jsonb_ops", "jsonb_path_ops": return !isArray && baseType == "jsonb" case "array_ops": return isArray default: return true } } func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool { return operatorClassCompatible("gin", baseType, isArray, opClass) } func isTextGinBaseType(baseType string) bool { switch baseType { case "text", "varchar", "character varying", "char", "character", "string", "citext", "bpchar": return true default: return false } } // requiredExtensions returns the PostgreSQL extensions a schema depends on, ordered so // that dependencies are created first (postgis before postgis_topology, vector before // vchord). Extensions are detected from column types, index access methods, resolved // operator classes, and function calls in defaults, check constraints, partial index // predicates and view definitions. Extensions that leave no trace in the model (pg_cron, // timescaledb, postgres_fdw, …) can be declared in schema.Metadata["extensions"]. func requiredExtensions(schema *models.Schema) []string { if schema == nil { return nil } required := make(map[string]bool) add := func(names ...string) { for _, name := range names { if name != "" { required[name] = true } } } add(declaredExtensions(schema)...) for _, view := range schema.Views { if view == nil { continue } add(pgsql.ExtensionsForExpression(view.Definition)...) } for _, table := range schema.Tables { if table == nil { continue } for _, col := range table.Columns { if col == nil { continue } add(pgsql.TypeExtension(effectiveColumnSQLType(col))) if def, ok := col.Default.(string); ok { add(pgsql.ExtensionsForExpression(def)...) } } for _, constraint := range table.Constraints { if constraint == nil { continue } add(pgsql.ExtensionsForExpression(constraint.Expression)...) } for _, index := range table.Indexes { if index == nil { continue } add(pgsql.IndexMethodExtension(index.Type)) add(pgsql.ExtensionsForExpression(index.Where)...) for _, colName := range index.Columns { col, ok := resolveIndexColumn(table, colName) if !ok || col == nil { continue } opClass := indexOperatorClassForColumn(col, index.Type, index.Comment) add(pgsql.OperatorClassExtension(opClass)) add(btreeCompanionExtension(index.Type, col, opClass)) } } } extensions := make([]string, 0, len(required)) for ext := range required { extensions = append(extensions, ext) } // Pull in dependencies, so a declared postgis_topology also creates postgis. for i := 0; i < len(extensions); i++ { for _, dependency := range pgsql.ExtensionDependencies(extensions[i]) { if !required[dependency] { required[dependency] = true extensions = append(extensions, dependency) } } } return pgsql.SortExtensions(extensions) } // declaredExtensions reads schema.Metadata["extensions"], which accepts either a list or a // comma-separated string. Unknown names are kept: the metadata is an explicit instruction. func declaredExtensions(schema *models.Schema) []string { value, ok := schema.Metadata["extensions"] if !ok { return nil } var names []string switch declared := value.(type) { case string: names = strings.Split(declared, ",") case []string: names = declared case []any: for _, item := range declared { if name, ok := item.(string); ok { names = append(names, name) } } default: return nil } cleaned := make([]string, 0, len(names)) for _, name := range names { if name = strings.TrimSpace(name); name != "" { cleaned = append(cleaned, name) } } return cleaned } // btreeCompanionExtension returns btree_gin or btree_gist when a GIN/GiST index covers a // scalar type that neither access method has a built-in operator class for. Without the // companion extension PostgreSQL rejects the CREATE INDEX outright. func btreeCompanionExtension(indexType string, col *models.Column, opClass string) string { if opClass != "" { return "" } method := strings.ToLower(strings.TrimSpace(indexType)) if method != "gin" && method != "gist" { return "" } sqlType := effectiveColumnSQLType(col) if pgsql.IsArrayType(sqlType) { return "" } baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType)) if pgsql.TypeExtension(baseType) != "" { // Extension types (geometry, vector, citext, …) ship their own operator classes. return "" } if method == "gin" { if nativeGinBaseType(baseType) { return "" } return "btree_gin" } if nativeGistBaseType(baseType) { return "" } return "btree_gist" } // nativeGinBaseType reports whether core PostgreSQL provides a GIN operator class. func nativeGinBaseType(baseType string) bool { switch baseType { case "jsonb", "json", "tsvector", "tsquery": return true default: return false } } // nativeGistBaseType reports whether core PostgreSQL provides a GiST operator class. func nativeGistBaseType(baseType string) bool { switch baseType { case "tsvector", "tsquery", "point", "box", "circle", "polygon", "line", "lseg", "path", "inet", "cidr": return true } return strings.HasSuffix(baseType, "range") || strings.HasSuffix(baseType, "multirange") } func schemaRequiresPGTrgm(schema *models.Schema) bool { for _, ext := range requiredExtensions(schema) { if ext == "pg_trgm" { return true } } return false } func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bool) { if table == nil { return nil, false } if col, ok := table.Columns[colName]; ok && col != nil { return col, true } normalized := strings.ToLower(strings.Trim(colName, `"`)) for key, col := range table.Columns { if col == nil { continue } if strings.ToLower(strings.Trim(key, `"`)) == normalized { return col, true } if strings.ToLower(strings.Trim(col.Name, `"`)) == normalized { return col, true } if strings.ToLower(strings.Trim(col.SQLName(), `"`)) == normalized { return col, true } } return nil, false } // sortColumns returns columns sorted by Sequence then Name for deterministic output. func sortColumns(columns map[string]*models.Column) []*models.Column { result := make([]*models.Column, 0, len(columns)) for _, col := range columns { result = append(result, col) } sort.Slice(result, func(i, j int) bool { if result[i].Sequence > 0 && result[j].Sequence > 0 { return result[i].Sequence < result[j].Sequence } return result[i].Name < result[j].Name }) return result } // sortConstraints returns constraints sorted by Sequence then Name for deterministic output. func sortConstraints(constraints map[string]*models.Constraint) []*models.Constraint { result := make([]*models.Constraint, 0, len(constraints)) for _, c := range constraints { result = append(result, c) } sort.Slice(result, func(i, j int) bool { if result[i].Sequence > 0 && result[j].Sequence > 0 { return result[i].Sequence < result[j].Sequence } return result[i].Name < result[j].Name }) return result } // sortIndexes returns indexes sorted by Sequence then Name for deterministic output. func sortIndexes(indexes map[string]*models.Index) []*models.Index { result := make([]*models.Index, 0, len(indexes)) for _, idx := range indexes { result = append(result, idx) } sort.Slice(result, func(i, j int) bool { if result[i].Sequence > 0 && result[j].Sequence > 0 { return result[i].Sequence < result[j].Sequence } return result[i].Name < result[j].Name }) return result } // formatStringList formats a list of strings as a SQL-safe comma-separated quoted list func formatStringList(items []string) string { quoted := make([]string, len(items)) for i, item := range items { quoted[i] = fmt.Sprintf("'%s'", escapeQuote(item)) } return strings.Join(quoted, ", ") } // extractOperatorClass extracts operator class from index comment/note // Looks for common operator classes like gin_trgm_ops, gist_trgm_ops, etc. // explicitOperatorClassPattern matches an "opclass=" hint, the form the PostgreSQL // reader uses to carry an index's operator class through the model. var explicitOperatorClassPattern = regexp.MustCompile(`(?i)\bopclass\s*=\s*([a-z_][a-z0-9_]*)\b`) func extractOperatorClass(comment string) string { if comment == "" { return "" } lowerComment := strings.ToLower(comment) if matches := explicitOperatorClassPattern.FindStringSubmatch(lowerComment); len(matches) > 1 { return matches[1] } for _, op := range knownOperatorClasses() { if strings.Contains(lowerComment, op) { return op } } return "" } // knownOperatorClasses lists every operator class recognized in an index comment, // longest name first so that e.g. gist_geometry_ops_nd wins over a shorter prefix. var knownOperatorClasses = sync.OnceValue(func() []string { names := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"} for name := range vectorOperatorClasses { names = append(names, name) } for name := range spatialOperatorClasses { names = append(names, name) } sort.Slice(names, func(i, j int) bool { if len(names[i]) != len(names[j]) { return len(names[i]) > len(names[j]) } return names[i] < names[j] }) return names }) // indexStorageParameters extracts access-method storage parameters from an index comment. // Only well-formed "key = value" pairs are kept, so comment prose cannot leak into DDL. // Example: "opclass=vector_cosine_ops with (m=16, ef_construction=64)" -> "m = 16, ef_construction = 64". func indexStorageParameters(comment string) string { if comment == "" { return "" } return pgsql.FormatStorageParameters(pgsql.ExtractWithClause(comment)) } // escapeQuote escapes single quotes in strings for SQL func escapeQuote(s string) string { return strings.ReplaceAll(s, "'", "''") } // stripBackticks removes backticks from SQL expressions // DBML uses backticks for SQL expressions like `now()`, but PostgreSQL doesn't use backticks func stripBackticks(s string) string { return strings.ReplaceAll(s, "`", "") } // extractSequenceName extracts sequence name from nextval() expression // Example: "nextval('public.users_id_seq'::regclass)" returns "users_id_seq" func extractSequenceName(defaultExpr string) string { // Look for nextval('schema.sequence_name'::regclass) pattern start := strings.Index(defaultExpr, "'") if start == -1 { return "" } end := strings.Index(defaultExpr[start+1:], "'") if end == -1 { return "" } fullName := defaultExpr[start+1 : start+1+end] // Remove schema prefix if present parts := strings.Split(fullName, ".") if len(parts) > 1 { return parts[len(parts)-1] } return fullName } // executeDatabaseSQL executes SQL statements directly on a PostgreSQL database func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) error { // Initialize execution report w.executionReport = &ExecutionReport{ StartTime: getCurrentTimestamp(), Schemas: make([]SchemaReport, 0), Errors: make([]ExecutionError, 0), } // Generate SQL statements statements, err := w.GenerateDatabaseStatements(db) if err != nil { return fmt.Errorf("failed to generate SQL statements: %w", err) } w.executionReport.TotalStatements = len(statements) // Connect to database ctx := context.Background() conn, err := pgsql.Connect(ctx, connString, "writer-pgsql") if err != nil { return fmt.Errorf("failed to connect to database: %w", err) } defer conn.Close(ctx) // Track schemas and tables schemaMap := make(map[string]*SchemaReport) currentSchema := "" // Execute each statement for i, stmt := range statements { stmtTrimmed := strings.TrimSpace(stmt) // Skip comments if strings.HasPrefix(stmtTrimmed, "--") { // Check if this is a schema comment to track schema changes if strings.Contains(stmtTrimmed, "Schema:") { parts := strings.Split(stmtTrimmed, "Schema:") if len(parts) > 1 { currentSchema = strings.TrimSpace(parts[1]) if _, exists := schemaMap[currentSchema]; !exists { schemaReport := SchemaReport{ Name: currentSchema, Tables: make([]TableReport, 0), } schemaMap[currentSchema] = &schemaReport } } } continue } // Skip empty statements if stmtTrimmed == "" { continue } stmtType := detectStatementType(stmtTrimmed) stmtCtx := extractStatementContext(stmtTrimmed) if stmtCtx != "" { fmt.Fprintf(os.Stderr, "Executing statement %d/%d [%s] %s...\n", i+1, len(statements), stmtType, stmtCtx) } else { fmt.Fprintf(os.Stderr, "Executing statement %d/%d [%s]...\n", i+1, len(statements), stmtType) } _, execErr := conn.Exec(ctx, stmt) if execErr != nil { w.executionReport.FailedStatements++ execError := ExecutionError{ StatementNumber: i + 1, Statement: truncateStatement(stmt), Error: execErr.Error(), } w.executionReport.Errors = append(w.executionReport.Errors, execError) // Track table creation failure if strings.Contains(strings.ToUpper(stmtTrimmed), "CREATE TABLE") && currentSchema != "" { tableName := extractTableNameFromCreate(stmtTrimmed) if tableName != "" && schemaMap[currentSchema] != nil { schemaMap[currentSchema].Tables = append(schemaMap[currentSchema].Tables, TableReport{ Name: tableName, Created: false, Error: execErr.Error(), }) } } // Continue with next statement instead of failing completely fmt.Fprintf(os.Stderr, "⚠ Warning: Statement %d failed: %v\n", i+1, execErr) continue } w.executionReport.ExecutedStatements++ // Track successful table creation if strings.Contains(strings.ToUpper(stmtTrimmed), "CREATE TABLE") && currentSchema != "" { tableName := extractTableNameFromCreate(stmtTrimmed) if tableName != "" && schemaMap[currentSchema] != nil { schemaMap[currentSchema].Tables = append(schemaMap[currentSchema].Tables, TableReport{ Name: tableName, Created: true, }) } } } // Convert schema map to slice for _, schemaReport := range schemaMap { w.executionReport.Schemas = append(w.executionReport.Schemas, *schemaReport) } w.executionReport.EndTime = getCurrentTimestamp() // Write report if path is specified if reportPath, ok := w.options.Metadata["report_path"].(string); ok && reportPath != "" { if err := w.writeReport(reportPath); err != nil { fmt.Fprintf(os.Stderr, "⚠ Warning: Failed to write report: %v\n", err) } else { fmt.Fprintf(os.Stderr, "✓ Report written to: %s\n", reportPath) } } if w.executionReport.FailedStatements > 0 { fmt.Fprintf(os.Stderr, "⚠ Completed with %d errors out of %d statements\n", w.executionReport.FailedStatements, w.executionReport.TotalStatements) } else { fmt.Fprintf(os.Stderr, "✓ Successfully executed %d statements\n", w.executionReport.ExecutedStatements) } return nil } // writeReport writes the execution report to a JSON file func (w *Writer) writeReport(reportPath string) error { file, err := os.Create(reportPath) if err != nil { return fmt.Errorf("failed to create report file: %w", err) } defer file.Close() encoder := json.NewEncoder(file) encoder.SetIndent("", " ") if err := encoder.Encode(w.executionReport); err != nil { return fmt.Errorf("failed to encode report: %w", err) } return nil } // extractTableNameFromCreate extracts table name from CREATE TABLE statement func extractTableNameFromCreate(stmt string) string { // Match: CREATE TABLE [IF NOT EXISTS] schema.table_name or table_name upper := strings.ToUpper(stmt) idx := strings.Index(upper, "CREATE TABLE") if idx == -1 { return "" } rest := strings.TrimSpace(stmt[idx+12:]) // Skip "CREATE TABLE" // Skip "IF NOT EXISTS" if strings.HasPrefix(strings.ToUpper(rest), "IF NOT EXISTS") { rest = strings.TrimSpace(rest[13:]) } // Get the table name (first token before '(' or whitespace) tokens := strings.FieldsFunc(rest, func(r rune) bool { return r == '(' || r == ' ' || r == '\n' || r == '\t' }) if len(tokens) == 0 { return "" } // Handle schema.table format fullName := tokens[0] parts := strings.Split(fullName, ".") if len(parts) > 1 { return parts[len(parts)-1] } return fullName } // truncateStatement truncates long SQL statements for error messages func truncateStatement(stmt string) string { const maxLen = 200 if len(stmt) <= maxLen { return stmt } return stmt[:maxLen] + "..." } // getCurrentTimestamp returns the current timestamp in a readable format func getCurrentTimestamp() string { return time.Now().Format("2006-01-02 15:04:05") } // extractStatementContext returns a human-readable schema/table/column context string for a SQL statement. func extractStatementContext(stmt string) string { upper := strings.ToUpper(stmt) // DO $$ blocks: extract identifiers from information_schema WHERE clauses if strings.HasPrefix(upper, "DO $$") || strings.HasPrefix(upper, "DO $") { schema := extractSQLStringValue(stmt, "table_schema") table := extractSQLStringValue(stmt, "table_name") column := extractSQLStringValue(stmt, "column_name") constraint := extractSQLStringValue(stmt, "constraint_name") return buildStmtContext(schema, table, column, constraint) } // ALTER TABLE [schema.]table ... if strings.HasPrefix(upper, "ALTER TABLE") { schema, table := parseQualifiedIdent(strings.TrimSpace(stmt[11:])) if strings.Contains(upper, "ADD COLUMN") { col := firstIdentAfterKeyword(stmt, upper, "ADD COLUMN") return buildStmtContext(schema, table, col, "") } if strings.Contains(upper, "ALTER COLUMN") { col := firstIdentAfterKeyword(stmt, upper, "ALTER COLUMN") return buildStmtContext(schema, table, col, "") } if strings.Contains(upper, "ADD CONSTRAINT") { con := firstIdentAfterKeyword(stmt, upper, "ADD CONSTRAINT") return buildStmtContext(schema, table, "", con) } if strings.Contains(upper, "DROP CONSTRAINT") { con := firstIdentAfterKeyword(stmt, upper, "DROP CONSTRAINT") return buildStmtContext(schema, table, "", con) } return buildStmtContext(schema, table, "", "") } // CREATE TABLE [IF NOT EXISTS] [schema.]table if strings.HasPrefix(upper, "CREATE TABLE") { rest := strings.TrimSpace(stmt[12:]) if strings.HasPrefix(strings.ToUpper(rest), "IF NOT EXISTS") { rest = strings.TrimSpace(rest[13:]) } schema, table := parseQualifiedIdent(rest) return buildStmtContext(schema, table, "", "") } // CREATE SCHEMA name if strings.HasPrefix(upper, "CREATE SCHEMA") { name := firstBareIdent(strings.TrimSpace(stmt[13:])) return name } // CREATE [UNIQUE] INDEX ... ON [schema.]table if strings.HasPrefix(upper, "CREATE INDEX") || strings.HasPrefix(upper, "CREATE UNIQUE INDEX") { onIdx := strings.Index(upper, " ON ") if onIdx != -1 { schema, table := parseQualifiedIdent(strings.TrimSpace(stmt[onIdx+4:])) return buildStmtContext(schema, table, "", "") } } // COMMENT ON TABLE [schema.]table if strings.HasPrefix(upper, "COMMENT ON TABLE") { schema, table := parseQualifiedIdent(strings.TrimSpace(stmt[16:])) return buildStmtContext(schema, table, "", "") } // COMMENT ON COLUMN [schema.]table.col if strings.HasPrefix(upper, "COMMENT ON COLUMN") { return firstBareIdent(strings.TrimSpace(stmt[17:])) } return "" } // extractSQLStringValue extracts 'value' from patterns like: key = 'value' (case-insensitive key match). func extractSQLStringValue(stmt, key string) string { lower := strings.ToLower(stmt) idx := strings.Index(lower, strings.ToLower(key)) if idx == -1 { return "" } rest := strings.TrimSpace(stmt[idx+len(key):]) eqIdx := strings.IndexByte(rest, '=') if eqIdx == -1 || eqIdx > 5 { return "" } rest = strings.TrimSpace(rest[eqIdx+1:]) if len(rest) == 0 || rest[0] != '\'' { return "" } rest = rest[1:] end := strings.IndexByte(rest, '\'') if end == -1 { return "" } return rest[:end] } // parseQualifiedIdent extracts (schema, name) from the start of s, handling optional "schema"."table" quoting. func parseQualifiedIdent(s string) (schema, name string) { token := firstBareIdent(s) parts := strings.SplitN(token, ".", 2) if len(parts) == 2 { return stripQuotes(parts[0]), stripQuotes(parts[1]) } return "", stripQuotes(token) } // firstBareIdent returns the first whitespace/delimiter-terminated token, stripping outer quotes. func firstBareIdent(s string) string { s = strings.TrimSpace(s) if s == "" { return "" } end := strings.IndexAny(s, " \t\n(,;") if end == -1 { return s } return s[:end] } // firstIdentAfterKeyword returns the first identifier token after keyword (matched case-insensitively via upperStmt). func firstIdentAfterKeyword(stmt, upperStmt, keyword string) string { idx := strings.Index(upperStmt, keyword) if idx == -1 { return "" } return stripQuotes(firstBareIdent(strings.TrimSpace(stmt[idx+len(keyword):]))) } // stripQuotes removes surrounding double-quotes from an identifier. func stripQuotes(s string) string { return strings.Trim(s, "\"") } // buildStmtContext assembles a display string from available identifiers. func buildStmtContext(schema, table, column, constraint string) string { var b strings.Builder if schema != "" && table != "" { b.WriteString(schema) b.WriteByte('.') b.WriteString(table) } else if table != "" { b.WriteString(table) } if column != "" { if b.Len() > 0 { b.WriteByte(' ') } b.WriteByte('(') b.WriteString(column) b.WriteByte(')') } if constraint != "" { if b.Len() > 0 { b.WriteByte(' ') } b.WriteByte('[') b.WriteString(constraint) b.WriteByte(']') } return b.String() } // detectStatementType detects the type of SQL statement for logging func detectStatementType(stmt string) string { upperStmt := strings.ToUpper(stmt) // Check for DO blocks (used for conditional DDL) if strings.HasPrefix(upperStmt, "DO $$") || strings.HasPrefix(upperStmt, "DO $") { // Look inside the DO block for the actual operation if strings.Contains(upperStmt, "ALTER TABLE") && strings.Contains(upperStmt, "ADD CONSTRAINT") { if strings.Contains(upperStmt, "UNIQUE") { return "ADD UNIQUE CONSTRAINT" } else if strings.Contains(upperStmt, "FOREIGN KEY") { return "ADD FOREIGN KEY" } else if strings.Contains(upperStmt, "PRIMARY KEY") { return "ADD PRIMARY KEY" } else if strings.Contains(upperStmt, "CHECK") { return "ADD CHECK CONSTRAINT" } return "ADD CONSTRAINT" } if strings.Contains(upperStmt, "ALTER TABLE") && strings.Contains(upperStmt, "ADD COLUMN") { return "ADD COLUMN" } if strings.Contains(upperStmt, "DROP CONSTRAINT") { return "DROP CONSTRAINT" } return "DO BLOCK" } // Direct DDL statements if strings.HasPrefix(upperStmt, "CREATE SCHEMA") { return "CREATE SCHEMA" } if strings.HasPrefix(upperStmt, "CREATE SEQUENCE") { return "CREATE SEQUENCE" } if strings.HasPrefix(upperStmt, "CREATE TABLE") { return "CREATE TABLE" } if strings.HasPrefix(upperStmt, "CREATE INDEX") { return "CREATE INDEX" } if strings.HasPrefix(upperStmt, "CREATE UNIQUE INDEX") { return "CREATE UNIQUE INDEX" } if strings.HasPrefix(upperStmt, "ALTER TABLE") { if strings.Contains(upperStmt, "ADD CONSTRAINT") { if strings.Contains(upperStmt, "FOREIGN KEY") { return "ADD FOREIGN KEY" } else if strings.Contains(upperStmt, "PRIMARY KEY") { return "ADD PRIMARY KEY" } else if strings.Contains(upperStmt, "UNIQUE") { return "ADD UNIQUE CONSTRAINT" } else if strings.Contains(upperStmt, "CHECK") { return "ADD CHECK CONSTRAINT" } return "ADD CONSTRAINT" } if strings.Contains(upperStmt, "ADD COLUMN") { return "ADD COLUMN" } if strings.Contains(upperStmt, "DROP CONSTRAINT") { return "DROP CONSTRAINT" } if strings.Contains(upperStmt, "ALTER COLUMN") { return "ALTER COLUMN" } return "ALTER TABLE" } if strings.HasPrefix(upperStmt, "COMMENT ON TABLE") { return "COMMENT ON TABLE" } if strings.HasPrefix(upperStmt, "COMMENT ON COLUMN") { return "COMMENT ON COLUMN" } if strings.HasPrefix(upperStmt, "DROP TABLE") { return "DROP TABLE" } if strings.HasPrefix(upperStmt, "DROP INDEX") { return "DROP INDEX" } // Default return "SQL" } // quoteIdentifier wraps an identifier in double quotes if necessary // This is needed for identifiers that start with numbers or contain special characters func quoteIdentifier(s string) string { return quoteIdent(s) }