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 }