feat: add MySQL reader and writer
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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, ", ")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user