- merge: cloneTable keeps relationships; skip-tables applies to new schemas - diff: detect schema description/owner changes - prisma: reader no longer turns relation fields into columns, enum detection via declared names; writer type mapping is ordered - typeorm writer: keep explicit SQL types that cannot be inferred - drizzle: enum columns call the enum constant; reader resolves them - mysql/mssql/sqlite writers: honour OutputPath, deterministic column and constraint order; mssql live execute covers full schema - template: ToYAML recovers from panics; Merge nil-pointer loop - regenerate drizzle fixtures
210 lines
5.0 KiB
Go
210 lines
5.0 KiB
Go
package mysql
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
"sort"
|
|
"strings"
|
|
|
|
_ "github.com/go-sql-driver/mysql"
|
|
|
|
"git.warky.dev/wdevs/relspecgo/pkg/mariadb"
|
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
|
)
|
|
|
|
type Writer struct {
|
|
options *writers.WriterOptions
|
|
writer io.Writer
|
|
}
|
|
|
|
func NewWriter(options *writers.WriterOptions) *Writer { return &Writer{options: options} }
|
|
func (w *Writer) WriteDatabase(db *models.Database) error {
|
|
if w.options == nil {
|
|
return fmt.Errorf("writer options are required")
|
|
}
|
|
if conn, ok := w.options.Metadata["connection_string"].(string); ok && conn != "" {
|
|
return w.execute(db, conn)
|
|
}
|
|
release, err := w.openOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer release()
|
|
return w.writeDatabaseDDL(db)
|
|
}
|
|
|
|
func (w *Writer) writeDatabaseDDL(db *models.Database) error {
|
|
if _, err := fmt.Fprintf(w.writer, "-- MySQL Database Schema\n-- Database: %s\n-- Generated by RelSpec\n\n", db.Name); err != nil {
|
|
return err
|
|
}
|
|
for _, s := range db.Schemas {
|
|
if err := w.WriteSchema(s); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// openOutput points w.writer at the configured destination (output file or
|
|
// stdout) when none is set, and returns a func that releases it again.
|
|
func (w *Writer) openOutput() (func(), error) {
|
|
if w.writer != nil {
|
|
return func() {}, nil
|
|
}
|
|
if w.options != nil && w.options.OutputPath != "" {
|
|
f, err := os.Create(w.options.OutputPath)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
w.writer = f
|
|
return func() {
|
|
f.Close()
|
|
w.writer = nil
|
|
}, nil
|
|
}
|
|
w.writer = os.Stdout
|
|
return func() { w.writer = nil }, nil
|
|
}
|
|
|
|
func (w *Writer) WriteSchema(s *models.Schema) error {
|
|
release, err := w.openOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer release()
|
|
for _, t := range s.Tables {
|
|
if err := w.writeTable(s, t); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (w *Writer) WriteTable(t *models.Table) error {
|
|
release, err := w.openOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer release()
|
|
return w.writeTable(nil, t)
|
|
}
|
|
|
|
func (w *Writer) writeTable(s *models.Schema, t *models.Table) error {
|
|
name := t.Name
|
|
if s != nil {
|
|
name = fmt.Sprintf("%s.%s", quote(s.Name), quote(t.Name))
|
|
} else {
|
|
name = quote(name)
|
|
}
|
|
cols := make([]*models.Column, 0, len(t.Columns))
|
|
for _, c := range t.Columns {
|
|
cols = append(cols, c)
|
|
}
|
|
sort.Slice(cols, func(i, j int) bool {
|
|
if cols[i].Sequence != cols[j].Sequence {
|
|
return cols[i].Sequence < cols[j].Sequence
|
|
}
|
|
return cols[i].Name < cols[j].Name
|
|
})
|
|
defs := []string{}
|
|
pk := []string{}
|
|
for _, c := range cols {
|
|
d := fmt.Sprintf(" %s %s", quote(c.Name), mariadb.ConvertCanonicalToMariaDB(c.Type))
|
|
if c.Length > 0 && strings.EqualFold(c.Type, "string") {
|
|
d = fmt.Sprintf(" %s VARCHAR(%d)", quote(c.Name), c.Length)
|
|
}
|
|
if c.NotNull {
|
|
d += " NOT NULL"
|
|
}
|
|
if c.AutoIncrement {
|
|
d += " AUTO_INCREMENT"
|
|
}
|
|
if c.Default != nil {
|
|
d += fmt.Sprintf(" DEFAULT %s", writers.QuoteDefaultValue(fmt.Sprint(c.Default), c.Type))
|
|
}
|
|
defs = append(defs, d)
|
|
if c.IsPrimaryKey {
|
|
pk = append(pk, quote(c.Name))
|
|
}
|
|
}
|
|
constraintNames := make([]string, 0, len(t.Constraints))
|
|
for name := range t.Constraints {
|
|
constraintNames = append(constraintNames, name)
|
|
}
|
|
sort.Strings(constraintNames)
|
|
for _, name := range constraintNames {
|
|
c := t.Constraints[name]
|
|
if c.Type == models.PrimaryKeyConstraint {
|
|
pk = nil
|
|
for _, n := range c.Columns {
|
|
pk = append(pk, quote(n))
|
|
}
|
|
break
|
|
}
|
|
}
|
|
if len(pk) > 0 {
|
|
defs = append(defs, " PRIMARY KEY ("+strings.Join(pk, ", ")+")")
|
|
}
|
|
for _, name := range constraintNames {
|
|
c := t.Constraints[name]
|
|
if c.Type == models.UniqueConstraint {
|
|
defs = append(defs, fmt.Sprintf(" CONSTRAINT %s UNIQUE (%s)", quote(c.Name), quoted(c.Columns)))
|
|
}
|
|
}
|
|
sql := fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (\n%s\n) ENGINE=InnoDB;\n\n", name, strings.Join(defs, ",\n"))
|
|
if _, err := io.WriteString(w.writer, sql); err != nil {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (w *Writer) execute(dbm *models.Database, conn string) error {
|
|
var b strings.Builder
|
|
old := w.writer
|
|
w.writer = &b
|
|
if err := w.writeDatabaseDDL(dbm); err != nil {
|
|
w.writer = old
|
|
return err
|
|
}
|
|
w.writer = old
|
|
db, err := sql.Open("mysql", conn)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to connect: %w", err)
|
|
}
|
|
defer db.Close()
|
|
for _, stmt := range strings.Split(b.String(), ";\n") {
|
|
stmt = stripComments(strings.TrimSpace(stmt))
|
|
if stmt == "" {
|
|
continue
|
|
}
|
|
if _, err := db.ExecContext(context.Background(), stmt); err != nil {
|
|
return fmt.Errorf("failed to execute SQL: %w", err)
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func stripComments(sqlText string) string {
|
|
lines := strings.Split(sqlText, "\n")
|
|
kept := lines[:0]
|
|
for _, line := range lines {
|
|
if !strings.HasPrefix(strings.TrimSpace(line), "--") {
|
|
kept = append(kept, line)
|
|
}
|
|
}
|
|
return strings.TrimSpace(strings.Join(kept, "\n"))
|
|
}
|
|
|
|
func quote(s string) string { return "`" + strings.ReplaceAll(s, "`", "``") + "`" }
|
|
func quoted(xs []string) string {
|
|
out := make([]string, len(xs))
|
|
for i, x := range xs {
|
|
out[i] = quote(x)
|
|
}
|
|
return strings.Join(out, ", ")
|
|
}
|