255 lines
7.3 KiB
Go
255 lines
7.3 KiB
Go
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
|
|
}
|