Add UUID/FK channel checks, agent whitelist/push APIs, bindings, super-admin LWW audit with rollback, reconcile rate limits, and sync docs. Default customers stay opt-in; company conflict UI is removed. Co-authored-by: Cursor <cursoragent@cursor.com>
265 lines
7.9 KiB
Go
265 lines
7.9 KiB
Go
package dbsync
|
||
|
||
import (
|
||
"context"
|
||
"database/sql"
|
||
"fmt"
|
||
"strings"
|
||
)
|
||
|
||
// ValidateChannelConfig 静态校验(不连库):表名单、PK 列名约定。
|
||
func ValidateChannelConfig(ch *Channel) error {
|
||
if ch == nil {
|
||
return fmt.Errorf("channel is nil")
|
||
}
|
||
tables := uniqueTables(ch.Local.Tables, ch.Remote.Tables)
|
||
if len(tables) == 0 {
|
||
return fmt.Errorf("同步表白名单为空:请至少在 local 或 remote 填写表名")
|
||
}
|
||
for _, t := range tables {
|
||
t = strings.TrimSpace(t)
|
||
if t == "" {
|
||
return fmt.Errorf("表名不能为空")
|
||
}
|
||
if strings.HasPrefix(t, "_ajz_") {
|
||
return fmt.Errorf("禁止同步系统表: %s", t)
|
||
}
|
||
pk := pkColumn(ch, t)
|
||
if pk == "" {
|
||
return fmt.Errorf("表 %s 主键列名为空", t)
|
||
}
|
||
}
|
||
return nil
|
||
}
|
||
|
||
func pkColumn(ch *Channel, table string) string {
|
||
if ch.PKColumns != nil {
|
||
if v := strings.TrimSpace(ch.PKColumns[table]); v != "" {
|
||
return v
|
||
}
|
||
}
|
||
return "id"
|
||
}
|
||
|
||
// ValidateChannelAgainstDB 对可连接端做:TEXT/UUID 主键类型 + FK 闭包。
|
||
// 某端连不上时跳过该端(形态 B 下 local 常不可达),但至少一端须校验成功,否则拒绝。
|
||
func ValidateChannelAgainstDB(ctx context.Context, ch *Channel) error {
|
||
if err := ValidateChannelConfig(ch); err != nil {
|
||
return err
|
||
}
|
||
tables := uniqueTables(ch.Local.Tables, ch.Remote.Tables)
|
||
checked := 0
|
||
var lastSkip error
|
||
for _, ep := range []Endpoint{ch.Remote, ch.Local} {
|
||
if strings.TrimSpace(ep.DSN) == "" || ep.Driver == "" {
|
||
continue
|
||
}
|
||
epTables := ep.Tables
|
||
if len(epTables) == 0 {
|
||
epTables = tables
|
||
}
|
||
db, err := Open(ep.Driver, ep.DSN)
|
||
if err != nil {
|
||
lastSkip = fmt.Errorf("%s 无法连接(跳过库内校验): %w", ep.Driver, err)
|
||
continue
|
||
}
|
||
if err := validateEndpointSchema(ctx, db, ep.Driver, ch, epTables, tables); err != nil {
|
||
_ = db.Close()
|
||
return fmt.Errorf("[%s] %w", ep.Driver, err)
|
||
}
|
||
_ = db.Close()
|
||
checked++
|
||
}
|
||
if checked == 0 {
|
||
if lastSkip != nil {
|
||
return fmt.Errorf("无法对任何端做主键/外键校验:%v;请保证线上库(remote)DSN 可达后再保存", lastSkip)
|
||
}
|
||
return fmt.Errorf("无法对任何端做主键/外键校验:请配置可达的 remote DSN")
|
||
}
|
||
return nil
|
||
}
|
||
|
||
func validateEndpointSchema(ctx context.Context, db *sql.DB, driver Driver, ch *Channel, epTables, whitelist []string) error {
|
||
wl := map[string]struct{}{}
|
||
for _, t := range whitelist {
|
||
wl[strings.TrimSpace(t)] = struct{}{}
|
||
}
|
||
for _, t := range epTables {
|
||
t = strings.TrimSpace(t)
|
||
if t == "" {
|
||
continue
|
||
}
|
||
pk := pkColumn(ch, t)
|
||
typ, err := DescribeColumnType(ctx, db, driver, t, pk)
|
||
if err != nil {
|
||
return fmt.Errorf("表 %s 主键列 %s: %w", t, pk, err)
|
||
}
|
||
if !isTextLikePK(typ) {
|
||
return fmt.Errorf("表 %s 主键 %s 类型为 %q,同步表须为 TEXT/VARCHAR/UUID 类(禁止自增整数作同步键)", t, pk, typ)
|
||
}
|
||
fks, err := ListForeignKeys(ctx, db, driver, t)
|
||
if err != nil {
|
||
return fmt.Errorf("表 %s 外键: %w", t, err)
|
||
}
|
||
for _, fk := range fks {
|
||
_, childIn := wl[fk.ChildTable]
|
||
_, parentIn := wl[fk.ParentTable]
|
||
if childIn != parentIn {
|
||
missing := fk.ParentTable
|
||
if !childIn {
|
||
missing = fk.ChildTable
|
||
}
|
||
return fmt.Errorf("外键闭包不完整:%s.%s → %s.%s,请将 %s 一并加入同步白名单",
|
||
fk.ChildTable, fk.ChildColumn, fk.ParentTable, fk.ParentColumn, missing)
|
||
}
|
||
}
|
||
}
|
||
return nil
|
||
}
|
||
|
||
func isTextLikePK(typ string) bool {
|
||
raw := strings.ToUpper(strings.TrimSpace(typ))
|
||
if raw == "" {
|
||
return false
|
||
}
|
||
token := strings.Fields(raw)[0]
|
||
token = strings.Split(token, "(")[0]
|
||
switch token {
|
||
case "TEXT", "VARCHAR", "CHAR", "CHARACTER", "UUID", "NVARCHAR", "NCHAR", "STRING", "CITEXT", "CLOB":
|
||
return true
|
||
case "INT", "INTEGER", "BIGINT", "SMALLINT", "TINYINT", "MEDIUMINT",
|
||
"SERIAL", "BIGSERIAL", "SMALLSERIAL",
|
||
"NUMERIC", "DECIMAL", "NUMBER", "FLOAT", "DOUBLE", "REAL", "BOOLEAN", "BOOL":
|
||
return false
|
||
default:
|
||
if strings.Contains(raw, "CHAR") || strings.Contains(raw, "TEXT") || strings.Contains(raw, "UUID") || strings.Contains(raw, "CLOB") {
|
||
return true
|
||
}
|
||
return false
|
||
}
|
||
}
|
||
|
||
// ForeignKey 子表 → 父表
|
||
type ForeignKey struct {
|
||
ChildTable string
|
||
ChildColumn string
|
||
ParentTable string
|
||
ParentColumn string
|
||
}
|
||
|
||
// DescribeColumnType 返回列类型字符串(驱动相关原文)。
|
||
func DescribeColumnType(ctx context.Context, db *sql.DB, driver Driver, table, column string) (string, error) {
|
||
table = strings.TrimSpace(table)
|
||
column = strings.TrimSpace(column)
|
||
switch driver {
|
||
case DriverSQLite:
|
||
rows, err := db.QueryContext(ctx, fmt.Sprintf(`PRAGMA table_info(%s)`, quoteIdent(driver, table)))
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
defer rows.Close()
|
||
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 "", err
|
||
}
|
||
if strings.EqualFold(name, column) {
|
||
if typ == "" {
|
||
typ = "TEXT" // sqlite 松类型兜底:无声明时按 TEXT 处理需调用方结合;此处空则拒
|
||
return "", fmt.Errorf("列 %s 无类型声明(sqlite);请显式声明为 TEXT", column)
|
||
}
|
||
return typ, nil
|
||
}
|
||
}
|
||
return "", fmt.Errorf("列不存在")
|
||
case DriverMySQL:
|
||
var typ string
|
||
err := db.QueryRowContext(ctx, `
|
||
SELECT DATA_TYPE FROM information_schema.columns
|
||
WHERE table_schema = DATABASE() AND table_name = ? AND column_name = ?`, table, column).Scan(&typ)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
return typ, nil
|
||
case DriverPostgres:
|
||
var typ string
|
||
err := db.QueryRowContext(ctx, `
|
||
SELECT data_type FROM information_schema.columns
|
||
WHERE table_schema = 'public' AND table_name = $1 AND column_name = $2`, table, column).Scan(&typ)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
return typ, nil
|
||
default:
|
||
return "", fmt.Errorf("unsupported driver")
|
||
}
|
||
}
|
||
|
||
// ListForeignKeys 列出以 table 为子表的外键。
|
||
func ListForeignKeys(ctx context.Context, db *sql.DB, driver Driver, table string) ([]ForeignKey, error) {
|
||
table = strings.TrimSpace(table)
|
||
switch driver {
|
||
case DriverSQLite:
|
||
rows, err := db.QueryContext(ctx, fmt.Sprintf(`PRAGMA foreign_key_list(%s)`, quoteIdent(driver, table)))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
var out []ForeignKey
|
||
for rows.Next() {
|
||
var id, seq int
|
||
var parent, from, to, onUpdate, onDelete, match string
|
||
if err := rows.Scan(&id, &seq, &parent, &from, &to, &onUpdate, &onDelete, &match); err != nil {
|
||
return nil, err
|
||
}
|
||
out = append(out, ForeignKey{
|
||
ChildTable: table, ChildColumn: from,
|
||
ParentTable: parent, ParentColumn: to,
|
||
})
|
||
}
|
||
return out, rows.Err()
|
||
case DriverMySQL:
|
||
rows, err := db.QueryContext(ctx, `
|
||
SELECT TABLE_NAME, COLUMN_NAME, REFERENCED_TABLE_NAME, REFERENCED_COLUMN_NAME
|
||
FROM information_schema.KEY_COLUMN_USAGE
|
||
WHERE table_schema = DATABASE() AND TABLE_NAME = ?
|
||
AND REFERENCED_TABLE_NAME IS NOT NULL`, table)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
return scanFKRows(rows)
|
||
case DriverPostgres:
|
||
rows, err := db.QueryContext(ctx, `
|
||
SELECT tc.table_name, kcu.column_name, ccu.table_name, ccu.column_name
|
||
FROM information_schema.table_constraints AS tc
|
||
JOIN information_schema.key_column_usage AS kcu
|
||
ON tc.constraint_name = kcu.constraint_name AND tc.table_schema = kcu.table_schema
|
||
JOIN information_schema.constraint_column_usage AS ccu
|
||
ON ccu.constraint_name = tc.constraint_name AND ccu.table_schema = tc.table_schema
|
||
WHERE tc.constraint_type = 'FOREIGN KEY' AND tc.table_schema = 'public' AND tc.table_name = $1`, table)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
return scanFKRows(rows)
|
||
default:
|
||
return nil, fmt.Errorf("unsupported driver")
|
||
}
|
||
}
|
||
|
||
func scanFKRows(rows *sql.Rows) ([]ForeignKey, error) {
|
||
var out []ForeignKey
|
||
for rows.Next() {
|
||
var ctab, ccol, ptab, pcol string
|
||
if err := rows.Scan(&ctab, &ccol, &ptab, &pcol); err != nil {
|
||
return nil, err
|
||
}
|
||
out = append(out, ForeignKey{ChildTable: ctab, ChildColumn: ccol, ParentTable: ptab, ParentColumn: pcol})
|
||
}
|
||
return out, rows.Err()
|
||
}
|