feat: add MySQL reader and writer

This commit is contained in:
SG Command
2026-10-03 12:36:39 +02:00
parent b38f53c603
commit 7a9219b6e3
53 changed files with 12311 additions and 0 deletions
+1
View File
@@ -20,6 +20,7 @@ const (
PostgresqlDatabaseType DatabaseType = "pgsql" // PostgreSQL database
MSSQLDatabaseType DatabaseType = "mssql" // Microsoft SQL Server database
SqlLiteDatabaseType DatabaseType = "sqlite" // SQLite database
MySQLDatabaseType DatabaseType = "mysql" // MySQL/MariaDB database
)
// Database represents the complete database schema
+245
View File
@@ -0,0 +1,245 @@
package mysql
import (
"context"
"database/sql"
"fmt"
"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/readers"
)
type Reader struct {
options *readers.ReaderOptions
db *sql.DB
ctx context.Context
}
func NewReader(options *readers.ReaderOptions) *Reader {
return &Reader{options: options, ctx: context.Background()}
}
func (r *Reader) ReadDatabase() (*models.Database, error) {
if r.options == nil || r.options.ConnectionString == "" {
return nil, fmt.Errorf("connection string is required")
}
if err := r.connect(); err != nil {
return nil, fmt.Errorf("failed to connect: %w", err)
}
defer r.close()
var name, version string
if err := r.db.QueryRowContext(r.ctx, "SELECT DATABASE()").Scan(&name); err != nil {
return nil, fmt.Errorf("failed to get database name: %w", err)
}
_ = r.db.QueryRowContext(r.ctx, "SELECT VERSION()").Scan(&version)
db := models.InitDatabase(name)
db.DatabaseType, db.SourceFormat, db.DatabaseVersion = models.MySQLDatabaseType, "mysql", version
schemas, err := r.querySchemas(name)
if err != nil {
return nil, fmt.Errorf("failed to query schemas: %w", err)
}
for _, schema := range schemas {
tables, err := r.queryTables(schema.Name)
if err != nil {
return nil, err
}
schema.Tables = tables
for _, table := range tables {
table.Columns, err = r.queryColumns(schema.Name, table.Name)
if err != nil {
return nil, err
}
table.Constraints, err = r.queryConstraints(schema.Name, table.Name)
if err != nil {
return nil, err
}
table.Indexes, err = r.queryIndexes(schema.Name, table.Name)
if err != nil {
return nil, err
}
table.RefSchema = schema
for _, c := range table.Constraints {
if c.Type == models.ForeignKeyConstraint {
r.deriveRelationship(table, c)
}
}
}
schema.RefDatabase = db
db.Schemas = append(db.Schemas, schema)
}
return db, nil
}
func (r *Reader) ReadSchema() (*models.Schema, error) {
db, err := r.ReadDatabase()
if err != nil {
return nil, err
}
if len(db.Schemas) == 0 {
return nil, fmt.Errorf("no schemas found in database")
}
return db.Schemas[0], nil
}
func (r *Reader) ReadTable() (*models.Table, error) {
s, err := r.ReadSchema()
if err != nil {
return nil, err
}
if len(s.Tables) == 0 {
return nil, fmt.Errorf("no tables found in schema")
}
return s.Tables[0], nil
}
func (r *Reader) connect() error {
db, err := sql.Open("mysql", r.options.ConnectionString)
if err != nil {
return err
}
if err = db.PingContext(r.ctx); err != nil {
db.Close()
return err
}
r.db = db
return nil
}
func (r *Reader) close() {
if r.db != nil {
_ = r.db.Close()
}
}
func (r *Reader) mapDataType(t string) string { return mariadb.ConvertMariaDBToCanonical(t) }
func (r *Reader) querySchemas(current string) ([]*models.Schema, error) {
rows, err := r.db.QueryContext(r.ctx, "SELECT SCHEMA_NAME FROM information_schema.SCHEMATA WHERE SCHEMA_NAME = ?", current)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*models.Schema
for rows.Next() {
var n string
if err := rows.Scan(&n); err != nil {
return nil, err
}
out = append(out, models.InitSchema(n))
}
return out, rows.Err()
}
func (r *Reader) queryTables(schema string) ([]*models.Table, error) {
rows, err := r.db.QueryContext(r.ctx, "SELECT TABLE_NAME FROM information_schema.TABLES WHERE TABLE_SCHEMA = ? AND TABLE_TYPE = 'BASE TABLE' ORDER BY TABLE_NAME", schema)
if err != nil {
return nil, err
}
defer rows.Close()
var out []*models.Table
for rows.Next() {
var n string
if err := rows.Scan(&n); err != nil {
return nil, err
}
out = append(out, models.InitTable(n, schema))
}
return out, rows.Err()
}
func (r *Reader) queryColumns(schema, table string) (map[string]*models.Column, error) {
rows, err := r.db.QueryContext(r.ctx, `SELECT COLUMN_NAME, COLUMN_TYPE, IS_NULLABLE, COLUMN_DEFAULT, ORDINAL_POSITION, EXTRA, COLUMN_COMMENT FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? ORDER BY ORDINAL_POSITION`, schema, table)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[string]*models.Column{}
for rows.Next() {
var name, typ, nullable, extra, comment string
var def sql.NullString
var pos int
if err := rows.Scan(&name, &typ, &nullable, &def, &pos, &extra, &comment); err != nil {
return nil, err
}
c := models.InitColumn(name, table, schema)
c.Type = r.mapDataType(typ)
c.NotNull = strings.EqualFold(nullable, "NO")
c.Sequence = uint(pos)
c.Comment = comment
if def.Valid {
c.Default = def.String
}
c.AutoIncrement = strings.Contains(strings.ToLower(extra), "auto_increment")
out[name] = c
}
return out, rows.Err()
}
func (r *Reader) queryConstraints(schema, table string) (map[string]*models.Constraint, error) {
rows, err := r.db.QueryContext(r.ctx, `SELECT CONSTRAINT_NAME, CONSTRAINT_TYPE, COLUMN_NAME, REFERENCED_TABLE_SCHEMA, REFERENCED_TABLE_NAME, REFERENCED_COLUMN_NAME, ORDINAL_POSITION FROM information_schema.KEY_COLUMN_USAGE k JOIN information_schema.TABLE_CONSTRAINTS t USING (CONSTRAINT_SCHEMA, TABLE_NAME, CONSTRAINT_NAME) WHERE k.TABLE_SCHEMA=? AND k.TABLE_NAME=? ORDER BY CONSTRAINT_NAME, ORDINAL_POSITION`, schema, table)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[string]*models.Constraint{}
for rows.Next() {
var name, typ, col, rs, rt, rc string
var pos int
if err := rows.Scan(&name, &typ, &col, &rs, &rt, &rc, &pos); err != nil {
return nil, err
}
c := out[name]
if c == nil {
ct := models.UniqueConstraint
if typ == "PRIMARY KEY" {
ct = models.PrimaryKeyConstraint
}
if typ == "FOREIGN KEY" {
ct = models.ForeignKeyConstraint
}
c = models.InitConstraint(name, ct)
c.Schema = schema
c.Table = table
c.ReferencedSchema = rs
c.ReferencedTable = rt
out[name] = c
}
c.Columns = append(c.Columns, col)
if rc != "" {
c.ReferencedColumns = append(c.ReferencedColumns, rc)
}
}
return out, rows.Err()
}
func (r *Reader) queryIndexes(schema, table string) (map[string]*models.Index, error) {
rows, err := r.db.QueryContext(r.ctx, `SELECT INDEX_NAME, NON_UNIQUE, COLUMN_NAME, SEQ_IN_INDEX, INDEX_TYPE FROM information_schema.STATISTICS WHERE TABLE_SCHEMA=? AND TABLE_NAME=? AND INDEX_NAME <> 'PRIMARY' ORDER BY INDEX_NAME, SEQ_IN_INDEX`, schema, table)
if err != nil {
return nil, err
}
defer rows.Close()
out := map[string]*models.Index{}
for rows.Next() {
var name, col, typ string
var non, seq int
if err := rows.Scan(&name, &non, &col, &seq, &typ); err != nil {
return nil, err
}
i := out[name]
if i == nil {
i = models.InitIndex(name, table, schema)
i.Unique = non == 0
i.Type = strings.ToLower(typ)
out[name] = i
}
i.Columns = append(i.Columns, col)
}
return out, rows.Err()
}
func (r *Reader) deriveRelationship(t *models.Table, c *models.Constraint) {
n := fmt.Sprintf("%s_to_%s", t.Name, c.ReferencedTable)
rel := models.InitRelationship(n, models.OneToMany)
rel.FromTable = t.Name
rel.FromSchema = t.Schema
rel.FromColumns = append([]string(nil), c.Columns...)
rel.ToTable = c.ReferencedTable
rel.ToSchema = c.ReferencedSchema
rel.ToColumns = append([]string(nil), c.ReferencedColumns...)
rel.ForeignKey = c.Name
t.Relationships[n] = rel
}
+21
View File
@@ -0,0 +1,21 @@
package mysql
import (
"git.warky.dev/wdevs/relspecgo/pkg/readers"
"testing"
)
func TestReaderMapDataType(t *testing.T) {
r := NewReader(&readers.ReaderOptions{})
for _, tc := range []struct{ input, want string }{{"varchar(64)", "string"}, {"bigint unsigned", "int64"}, {"datetime", "timestamp"}, {"json", "json"}} {
if got := r.mapDataType(tc.input); got != tc.want {
t.Errorf("mapDataType(%q) = %q, want %q", tc.input, got, tc.want)
}
}
}
func TestReaderRequiresConnectionString(t *testing.T) {
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil {
t.Fatal("expected missing connection string error")
}
}
+152
View File
@@ -0,0 +1,152 @@
package mysql
import (
"context"
"database/sql"
"fmt"
"io"
"os"
"sort"
"strings"
"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)
}
if w.writer == nil {
if w.options.OutputPath != "" {
f, err := os.Create(w.options.OutputPath)
if err != nil {
return err
}
defer f.Close()
w.writer = f
} else {
w.writer = os.Stdout
}
}
fmt.Fprintf(w.writer, "-- MySQL Database Schema\n-- Database: %s\n-- Generated by RelSpec\n\n", db.Name)
for _, s := range db.Schemas {
if err := w.WriteSchema(s); err != nil {
return err
}
}
return nil
}
func (w *Writer) WriteSchema(s *models.Schema) error {
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 {
if w.writer == nil {
w.writer = os.Stdout
}
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 { return cols[i].Sequence < cols[j].Sequence })
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))
}
}
for _, c := range t.Constraints {
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 _, c := range t.Constraints {
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.WriteDatabase(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 = strings.TrimSpace(stmt)
if stmt == "" || strings.HasPrefix(stmt, "--") {
continue
}
if _, err := db.ExecContext(context.Background(), stmt); err != nil {
return fmt.Errorf("failed to execute SQL: %w", err)
}
}
return nil
}
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, ", ")
}
+37
View File
@@ -0,0 +1,37 @@
package mysql
import (
"bytes"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/writers"
)
func TestWriterGeneratesMySQLDDL(t *testing.T) {
db := models.InitDatabase("app")
s := models.InitSchema("app")
table := models.InitTable("users", "app")
table.Columns["id"] = models.InitColumn("id", "users", "app")
table.Columns["id"].Type = "int"
table.Columns["id"].IsPrimaryKey = true
table.Columns["id"].NotNull = true
table.Columns["name"] = models.InitColumn("name", "users", "app")
table.Columns["name"].Type = "string"
table.Columns["name"].Length = 80
s.Tables = append(s.Tables, table)
db.Schemas = append(db.Schemas, s)
var out bytes.Buffer
w := NewWriter(&writers.WriterOptions{Metadata: map[string]interface{}{}})
w.writer = &out
if err := w.WriteDatabase(db); err != nil {
t.Fatal(err)
}
got := out.String()
for _, want := range []string{"CREATE TABLE IF NOT EXISTS `app`.`users`", "`id` INT NOT NULL", "`name` VARCHAR(80)", "PRIMARY KEY (`id`)"} {
if !strings.Contains(got, want) {
t.Errorf("DDL missing %q:\n%s", want, got)
}
}
}