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