Files
ai_site/platform/internal/dbsync/fingerprint.go
whm 4fae7ae718 feat: add schema fingerprint and Z10d alter; refresh sync coop docs
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>
2026-08-05 18:16:55 +08:00

200 lines
5.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}