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, ", ") }