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 }