Enable Binding-scoped agent push/pull, empty-table schema ensure, SyncPage inspect/drop-table, default module import, and agent-bound publish docs from the 宇恒联调意见. Co-authored-by: Cursor <cursoragent@cursor.com>
278 lines
7.3 KiB
Go
278 lines
7.3 KiB
Go
package dbsync
|
||
|
||
import (
|
||
"context"
|
||
"database/sql"
|
||
"fmt"
|
||
"strings"
|
||
)
|
||
|
||
// TableInspect 线上/本机库一张表的摘要(控制台验同步用)。
|
||
type TableInspect struct {
|
||
Name string `json:"name"`
|
||
RowCount int64 `json:"row_count"`
|
||
Columns []string `json:"columns"`
|
||
ColumnCount int `json:"column_count"`
|
||
}
|
||
|
||
// InspectResult 库体检结果。
|
||
type InspectResult struct {
|
||
OK bool `json:"ok"`
|
||
Side string `json:"side"` // remote | local
|
||
Driver string `json:"driver"`
|
||
DSNHint string `json:"dsn_hint,omitempty"`
|
||
Tables []TableInspect `json:"tables"`
|
||
Message string `json:"message,omitempty"`
|
||
}
|
||
|
||
// PreviewResult 单表行预览。
|
||
type PreviewResult struct {
|
||
OK bool `json:"ok"`
|
||
Table string `json:"table"`
|
||
Columns []string `json:"columns"`
|
||
Rows []map[string]any `json:"rows"`
|
||
Total int64 `json:"total"`
|
||
Limit int `json:"limit"`
|
||
Message string `json:"message,omitempty"`
|
||
}
|
||
|
||
// InspectEndpoint 列出业务表 + 行数 + 列名(含 _ajz_ 同步系统表,便于对照)。
|
||
func InspectEndpoint(ctx context.Context, ep Endpoint, side string, includeSyncMeta bool) (*InspectResult, error) {
|
||
db, owned, err := openInspectDB(ep)
|
||
if err != nil {
|
||
return &InspectResult{OK: false, Side: side, Driver: string(ep.Driver), Message: err.Error()}, err
|
||
}
|
||
if owned {
|
||
defer db.Close()
|
||
}
|
||
names, err := listTablesForInspect(ctx, db, ep.Driver, includeSyncMeta)
|
||
if err != nil {
|
||
return &InspectResult{OK: false, Side: side, Driver: string(ep.Driver), Message: err.Error()}, err
|
||
}
|
||
out := &InspectResult{
|
||
OK: true,
|
||
Side: side,
|
||
Driver: string(ep.Driver),
|
||
DSNHint: dsnFileHint(ep.DSN),
|
||
Tables: make([]TableInspect, 0, len(names)),
|
||
}
|
||
for _, name := range names {
|
||
ti := TableInspect{Name: name}
|
||
cols, _ := listColumns(ctx, db, ep.Driver, name)
|
||
ti.Columns = cols
|
||
ti.ColumnCount = len(cols)
|
||
n, _ := countRows(ctx, db, ep.Driver, name)
|
||
ti.RowCount = n
|
||
out.Tables = append(out.Tables, ti)
|
||
}
|
||
if len(out.Tables) == 0 {
|
||
out.Message = "库可连接,但尚无业务表(可能还未 push)"
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// PreviewTable 预览表前 limit 行(默认 50,最大 200)。
|
||
func PreviewTable(ctx context.Context, ep Endpoint, table string, limit int) (*PreviewResult, error) {
|
||
table = strings.TrimSpace(table)
|
||
if table == "" {
|
||
return nil, fmt.Errorf("table required")
|
||
}
|
||
if limit <= 0 {
|
||
limit = 50
|
||
}
|
||
if limit > 200 {
|
||
limit = 200
|
||
}
|
||
db, owned, err := openInspectDB(ep)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if owned {
|
||
defer db.Close()
|
||
}
|
||
cols, err := listColumns(ctx, db, ep.Driver, table)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
total, _ := countRows(ctx, db, ep.Driver, table)
|
||
rows, err := fetchPreviewRows(ctx, db, ep.Driver, table, cols, limit)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return &PreviewResult{
|
||
OK: true,
|
||
Table: table,
|
||
Columns: cols,
|
||
Rows: rows,
|
||
Total: total,
|
||
Limit: limit,
|
||
}, nil
|
||
}
|
||
|
||
func openInspectDB(ep Endpoint) (*sql.DB, bool, error) {
|
||
// 控制台验库用短连接,不进 remote 池,避免 SQLite 文件被长期占用。
|
||
db, err := Open(ep.Driver, ep.DSN)
|
||
if err != nil {
|
||
return nil, false, err
|
||
}
|
||
return db, true, nil
|
||
}
|
||
|
||
func listTablesForInspect(ctx context.Context, db *sql.DB, driver Driver, includeSyncMeta bool) ([]string, error) {
|
||
var q string
|
||
switch driver {
|
||
case DriverSQLite:
|
||
if includeSyncMeta {
|
||
q = `SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name`
|
||
} else {
|
||
q = `SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' AND name NOT LIKE '_ajz_%' ORDER BY name`
|
||
}
|
||
case DriverMySQL:
|
||
if includeSyncMeta {
|
||
q = `SELECT table_name FROM information_schema.tables WHERE table_schema = DATABASE() ORDER BY table_name`
|
||
} else {
|
||
q = `SELECT table_name FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name NOT LIKE '\_ajz\_%' ORDER BY table_name`
|
||
}
|
||
case DriverPostgres:
|
||
if includeSyncMeta {
|
||
q = `SELECT tablename FROM pg_tables WHERE schemaname='public' ORDER BY tablename`
|
||
} else {
|
||
q = `SELECT tablename FROM pg_tables WHERE schemaname='public' AND tablename NOT LIKE '\_ajz\_%' ORDER BY tablename`
|
||
}
|
||
default:
|
||
return nil, fmt.Errorf("unsupported driver")
|
||
}
|
||
rows, err := db.QueryContext(ctx, q)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
var out []string
|
||
for rows.Next() {
|
||
var n string
|
||
if err := rows.Scan(&n); err != nil {
|
||
return nil, err
|
||
}
|
||
out = append(out, n)
|
||
}
|
||
return out, rows.Err()
|
||
}
|
||
|
||
func listColumns(ctx context.Context, db *sql.DB, driver Driver, table string) ([]string, error) {
|
||
var (
|
||
rows *sql.Rows
|
||
err error
|
||
)
|
||
switch driver {
|
||
case DriverSQLite:
|
||
rows, err = db.QueryContext(ctx, fmt.Sprintf(`PRAGMA table_info(%s)`, quoteIdent(driver, table)))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
var cols []string
|
||
for rows.Next() {
|
||
var cid int
|
||
var name, typ string
|
||
var notnull, pk int
|
||
var dflt sql.NullString
|
||
if err := rows.Scan(&cid, &name, &typ, ¬null, &dflt, &pk); err != nil {
|
||
return nil, err
|
||
}
|
||
cols = append(cols, name)
|
||
}
|
||
return cols, rows.Err()
|
||
case DriverMySQL:
|
||
rows, err = db.QueryContext(ctx, `
|
||
SELECT COLUMN_NAME FROM information_schema.columns
|
||
WHERE table_schema = DATABASE() AND table_name = ? ORDER BY ORDINAL_POSITION`, table)
|
||
case DriverPostgres:
|
||
rows, err = db.QueryContext(ctx, `
|
||
SELECT column_name FROM information_schema.columns
|
||
WHERE table_schema = 'public' AND table_name = $1 ORDER BY ordinal_position`, table)
|
||
default:
|
||
return nil, fmt.Errorf("unsupported")
|
||
}
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
var cols []string
|
||
for rows.Next() {
|
||
var n string
|
||
if err := rows.Scan(&n); err != nil {
|
||
return nil, err
|
||
}
|
||
cols = append(cols, n)
|
||
}
|
||
return cols, rows.Err()
|
||
}
|
||
|
||
func countRows(ctx context.Context, db *sql.DB, driver Driver, table string) (int64, error) {
|
||
q := fmt.Sprintf(`SELECT COUNT(*) FROM %s`, quoteIdent(driver, table))
|
||
var n int64
|
||
err := db.QueryRowContext(ctx, q).Scan(&n)
|
||
return n, err
|
||
}
|
||
|
||
func fetchPreviewRows(ctx context.Context, db *sql.DB, driver Driver, table string, cols []string, limit int) ([]map[string]any, error) {
|
||
if len(cols) == 0 {
|
||
return []map[string]any{}, nil
|
||
}
|
||
q := fmt.Sprintf(`SELECT * FROM %s LIMIT %d`, quoteIdent(driver, table), limit)
|
||
rows, err := db.QueryContext(ctx, q)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
colNames, err := rows.Columns()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
var out []map[string]any
|
||
for rows.Next() {
|
||
raw := make([]any, len(colNames))
|
||
ptrs := make([]any, len(colNames))
|
||
for i := range raw {
|
||
ptrs[i] = &raw[i]
|
||
}
|
||
if err := rows.Scan(ptrs...); err != nil {
|
||
return nil, err
|
||
}
|
||
m := make(map[string]any, len(colNames))
|
||
for i, c := range colNames {
|
||
m[c] = normalizeSQLValue(raw[i])
|
||
}
|
||
out = append(out, m)
|
||
}
|
||
return out, rows.Err()
|
||
}
|
||
|
||
func normalizeSQLValue(v any) any {
|
||
switch x := v.(type) {
|
||
case nil:
|
||
return nil
|
||
case []byte:
|
||
return string(x)
|
||
default:
|
||
return x
|
||
}
|
||
}
|
||
|
||
func dsnFileHint(dsn string) string {
|
||
dsn = strings.TrimSpace(dsn)
|
||
if i := strings.LastIndexAny(dsn, `/\`); i >= 0 {
|
||
rest := dsn[i+1:]
|
||
if j := strings.IndexAny(rest, "?#"); j >= 0 {
|
||
rest = rest[:j]
|
||
}
|
||
if strings.HasSuffix(strings.ToLower(rest), ".db") {
|
||
return rest
|
||
}
|
||
}
|
||
if len(dsn) > 64 {
|
||
return dsn[:64] + "…"
|
||
}
|
||
return dsn
|
||
}
|