Table.Columns/Constraints/Indexes/Relationships are Go maps, and every writer, reader, diff, inspector, and merge code path that iterated them directly was subject to Go's randomized map order, so identical input could produce different output (or a different in-report violation/diff order) on every run. Most visibly this showed up as bun/gorm `unique:` struct tags changing order across consecutive `make models` runs with no source change. Fixed by sorting map iteration (by Sequence then Name, or alphabetically for string-keyed maps) everywhere the order affects generated output or first-match tie-break logic, across the bun, gorm, sqlite, dbml, drawdb, pgsql, prisma, graphql, typeorm, drizzle, and dctx writers; the dctx, prisma, and typeorm readers; the shared models.GetPrimaryKey/ GetForeignKeys helpers; pkg/diff, pkg/inspector, and pkg/merge; and the TUI column/relationship pickers in pkg/ui.
2034 lines
58 KiB
Go
2034 lines
58 KiB
Go
package pgsql
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
"sort"
|
|
"strings"
|
|
"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()))
|
|
}
|
|
|
|
if schemaRequiresPGTrgm(schema) {
|
|
statements = append(statements, `CREATE EXTENSION IF NOT EXISTS pg_trgm`)
|
|
}
|
|
|
|
// 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 for GIN indexes
|
|
columnExprs := make([]string, 0, len(index.Columns))
|
|
for _, colName := range index.Columns {
|
|
colExpr := colName
|
|
if col, ok := resolveIndexColumn(table, colName); ok {
|
|
if strings.EqualFold(indexType, "gin") {
|
|
if opClass := ginOperatorClassForColumn(col, index.Comment); opClass != "" {
|
|
colExpr = fmt.Sprintf("%s %s", colName, opClass)
|
|
}
|
|
}
|
|
}
|
|
columnExprs = append(columnExprs, colExpr)
|
|
}
|
|
|
|
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",
|
|
uniqueStr, quoteIdentifier(index.Name), w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), 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
|
|
}
|
|
|
|
// 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
|
|
}
|
|
|
|
// 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 {
|
|
if !schemaRequiresPGTrgm(schema) {
|
|
return nil
|
|
}
|
|
|
|
fmt.Fprintln(w.writer, "CREATE EXTENSION IF NOT EXISTS pg_trgm;")
|
|
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
|
|
}
|
|
|
|
// 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))
|
|
}
|
|
|
|
// Build column list with operator class support for GIN indexes
|
|
columnExprs := make([]string, 0, len(index.Columns))
|
|
for _, colName := range index.Columns {
|
|
if col, ok := resolveIndexColumn(table, colName); ok {
|
|
colExpr := col.SQLName()
|
|
if strings.EqualFold(index.Type, "gin") {
|
|
opClass := ginOperatorClassForColumn(col, index.Comment)
|
|
if opClass != "" {
|
|
colExpr = fmt.Sprintf("%s %s", col.SQLName(), opClass)
|
|
}
|
|
}
|
|
columnExprs = append(columnExprs, colExpr)
|
|
}
|
|
}
|
|
|
|
if len(columnExprs) == 0 {
|
|
continue
|
|
}
|
|
|
|
unique := ""
|
|
if index.Unique {
|
|
unique = "UNIQUE "
|
|
}
|
|
|
|
indexType := index.Type
|
|
if indexType == "" {
|
|
indexType = "btree"
|
|
}
|
|
|
|
whereClause := ""
|
|
if index.Where != "" {
|
|
whereClause = fmt.Sprintf(" WHERE %s", index.Where)
|
|
}
|
|
|
|
fmt.Fprintf(w.writer, "CREATE %sINDEX IF NOT EXISTS %s\n",
|
|
unique, indexName)
|
|
fmt.Fprintf(w.writer, " ON %s USING %s (%s)%s;\n\n",
|
|
w.qualTable(schema.SQLName(), table.SQLName()), indexType, strings.Join(columnExprs, ", "), 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")
|
|
}
|
|
|
|
func ginOperatorClassForColumn(col *models.Column, comment string) string {
|
|
if col == nil {
|
|
return ""
|
|
}
|
|
|
|
sqlType := effectiveColumnSQLType(col)
|
|
baseType := pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
|
|
isArray := pgsql.IsArrayType(sqlType)
|
|
requested := extractOperatorClass(comment)
|
|
|
|
if requested != "" && ginOperatorClassCompatible(baseType, isArray, requested) {
|
|
return requested
|
|
}
|
|
|
|
if isArray {
|
|
return "array_ops"
|
|
}
|
|
|
|
switch {
|
|
case isTextGinBaseType(baseType):
|
|
return "gin_trgm_ops"
|
|
case baseType == "jsonb":
|
|
return "jsonb_ops"
|
|
default:
|
|
return requested
|
|
}
|
|
}
|
|
|
|
func ginOperatorClassCompatible(baseType string, isArray bool, opClass string) bool {
|
|
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 isTextGinBaseType(baseType string) bool {
|
|
switch baseType {
|
|
case "text", "varchar", "character varying", "char", "character", "string", "citext", "bpchar":
|
|
return true
|
|
default:
|
|
return false
|
|
}
|
|
}
|
|
|
|
func schemaRequiresPGTrgm(schema *models.Schema) bool {
|
|
if schema == nil {
|
|
return false
|
|
}
|
|
for _, table := range schema.Tables {
|
|
if table == nil {
|
|
continue
|
|
}
|
|
for _, index := range table.Indexes {
|
|
if index == nil || !strings.EqualFold(index.Type, "gin") {
|
|
continue
|
|
}
|
|
for _, colName := range index.Columns {
|
|
col, ok := resolveIndexColumn(table, colName)
|
|
if !ok || col == nil {
|
|
continue
|
|
}
|
|
if ginOperatorClassForColumn(col, index.Comment) == "gin_trgm_ops" {
|
|
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.
|
|
func extractOperatorClass(comment string) string {
|
|
if comment == "" {
|
|
return ""
|
|
}
|
|
lowerComment := strings.ToLower(comment)
|
|
// Common GIN/GiST operator classes
|
|
opClasses := []string{"gin_trgm_ops", "gist_trgm_ops", "gin_bigm_ops", "jsonb_ops", "jsonb_path_ops", "array_ops"}
|
|
for _, op := range opClasses {
|
|
if strings.Contains(lowerComment, op) {
|
|
return op
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
// 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)
|
|
}
|