Ship ensure ADD COLUMN and fingerprint API for Yuheng reconcile, and rewrite §0.2–0.3 cooperation checklist in the联调意见. Co-authored-by: Cursor <cursoragent@cursor.com>
200 lines
5.0 KiB
Go
200 lines
5.0 KiB
Go
package dbsync
|
||
|
||
import (
|
||
"context"
|
||
"crypto/sha256"
|
||
"database/sql"
|
||
"encoding/hex"
|
||
"encoding/json"
|
||
"fmt"
|
||
"sort"
|
||
"strings"
|
||
)
|
||
|
||
// FingerprintTable 单表指纹(与宇恒 sync_fingerprint 算法对齐)。
|
||
type FingerprintTable struct {
|
||
Name string `json:"name"`
|
||
PKColumn string `json:"pk_column"`
|
||
RowCount int64 `json:"row_count"`
|
||
ContentHash string `json:"content_hash"`
|
||
Columns []string `json:"columns,omitempty"`
|
||
Error string `json:"error,omitempty"`
|
||
}
|
||
|
||
// FingerprintRequest agent 拉线上表指纹。
|
||
type FingerprintRequest struct {
|
||
OnlineDBID string `json:"online_db_id"`
|
||
Tables []string `json:"tables,omitempty"`
|
||
}
|
||
|
||
// FingerprintResult 批量指纹。
|
||
type FingerprintResult struct {
|
||
OK bool `json:"ok"`
|
||
Tables []FingerprintTable `json:"tables"`
|
||
}
|
||
|
||
// FingerprintRemoteTables 计算线上 A 业务表 row_count + content_hash。
|
||
// 哈希:sha256( table+"|"+pk+"|" + Σ json.dumps(row, sort_keys=True)+"\n" ) 取前 32 hex。
|
||
func FingerprintRemoteTables(ctx context.Context, ch *Channel, req FingerprintRequest) (*FingerprintResult, error) {
|
||
if ch == nil {
|
||
return nil, fmt.Errorf("channel is nil")
|
||
}
|
||
db, err := AcquireRemote(ch.Remote.Driver, ch.Remote.DSN)
|
||
if err != nil {
|
||
return nil, wrapOpenRemote(err)
|
||
}
|
||
if err := EnsureMeta(ctx, db, ch.Remote.Driver); err != nil {
|
||
return nil, Retryablef("ensure meta: %v", err)
|
||
}
|
||
names, err := ListTables(ctx, db, ch.Remote.Driver)
|
||
if err != nil {
|
||
return nil, Retryablef("list tables: %v", err)
|
||
}
|
||
want := map[string]struct{}{}
|
||
for _, t := range req.Tables {
|
||
t = strings.TrimSpace(t)
|
||
if t != "" {
|
||
want[t] = struct{}{}
|
||
}
|
||
}
|
||
out := &FingerprintResult{OK: true, Tables: make([]FingerprintTable, 0, len(names))}
|
||
for _, name := range names {
|
||
if strings.HasPrefix(name, "_ajz_") {
|
||
continue
|
||
}
|
||
if len(want) > 0 {
|
||
if _, ok := want[name]; !ok {
|
||
continue
|
||
}
|
||
}
|
||
if !tableInChannel(ch, name) {
|
||
continue
|
||
}
|
||
fp, ferr := fingerprintOneTable(ctx, db, ch, name)
|
||
if ferr != nil {
|
||
fp.Error = ferr.Error()
|
||
out.OK = false
|
||
}
|
||
out.Tables = append(out.Tables, fp)
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
func fingerprintOneTable(ctx context.Context, db *sql.DB, ch *Channel, table string) (FingerprintTable, error) {
|
||
cols, err := listColumns(ctx, db, ch.Remote.Driver, table)
|
||
if err != nil {
|
||
return FingerprintTable{Name: table}, err
|
||
}
|
||
pkCol := "id"
|
||
if ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[table]) != "" {
|
||
pkCol = ch.PKColumns[table]
|
||
}
|
||
if !containsStr(cols, pkCol) {
|
||
for _, c := range []string{"id", "__row_id", "_row_id"} {
|
||
if containsStr(cols, c) {
|
||
pkCol = c
|
||
break
|
||
}
|
||
}
|
||
if !containsStr(cols, pkCol) && len(cols) > 0 {
|
||
pkCol = cols[0]
|
||
}
|
||
}
|
||
n, err := countRows(ctx, db, ch.Remote.Driver, table)
|
||
if err != nil {
|
||
return FingerprintTable{Name: table, PKColumn: pkCol, Columns: cols}, err
|
||
}
|
||
hash, err := contentHashOrdered(ctx, db, ch.Remote.Driver, table, pkCol)
|
||
if err != nil {
|
||
return FingerprintTable{Name: table, PKColumn: pkCol, RowCount: n, Columns: cols}, err
|
||
}
|
||
return FingerprintTable{
|
||
Name: table,
|
||
PKColumn: pkCol,
|
||
RowCount: n,
|
||
ContentHash: hash,
|
||
Columns: cols,
|
||
}, nil
|
||
}
|
||
|
||
func containsStr(ss []string, x string) bool {
|
||
for _, s := range ss {
|
||
if s == x {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
|
||
func contentHashOrdered(ctx context.Context, db *sql.DB, driver Driver, table, pkCol string) (string, error) {
|
||
q := fmt.Sprintf(`SELECT * FROM %s ORDER BY %s`, quoteIdent(driver, table), quoteIdent(driver, pkCol))
|
||
rows, err := db.QueryContext(ctx, q)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
defer rows.Close()
|
||
colNames, err := rows.Columns()
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
h := sha256.New()
|
||
_, _ = h.Write([]byte(table + "|" + pkCol + "|"))
|
||
ptrs := make([]any, len(colNames))
|
||
vals := make([]any, len(colNames))
|
||
for i := range vals {
|
||
ptrs[i] = &vals[i]
|
||
}
|
||
for rows.Next() {
|
||
for i := range vals {
|
||
vals[i] = nil
|
||
}
|
||
if err := rows.Scan(ptrs...); err != nil {
|
||
return "", err
|
||
}
|
||
m := make(map[string]any, len(colNames))
|
||
for i, c := range colNames {
|
||
m[c] = normalizeSQLValue(vals[i])
|
||
}
|
||
blob, err := stableJSON(m)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
_, _ = h.Write(blob)
|
||
_, _ = h.Write([]byte("\n"))
|
||
}
|
||
if err := rows.Err(); err != nil {
|
||
return "", err
|
||
}
|
||
sum := hex.EncodeToString(h.Sum(nil))
|
||
if len(sum) > 32 {
|
||
sum = sum[:32]
|
||
}
|
||
return sum, nil
|
||
}
|
||
|
||
// stableJSON 按 key 排序,贴近 Python json.dumps(..., sort_keys=True, default=str)。
|
||
func stableJSON(m map[string]any) ([]byte, error) {
|
||
keys := make([]string, 0, len(m))
|
||
for k := range m {
|
||
keys = append(keys, k)
|
||
}
|
||
sort.Strings(keys)
|
||
parts := make([]string, 0, len(keys))
|
||
for _, k := range keys {
|
||
kb, err := json.Marshal(k)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
vb, err := json.Marshal(m[k])
|
||
if err != nil {
|
||
// fallback string
|
||
vb, err = json.Marshal(fmt.Sprint(m[k]))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
}
|
||
parts = append(parts, string(kb)+":"+string(vb))
|
||
}
|
||
return []byte("{" + strings.Join(parts, ",") + "}"), nil
|
||
}
|