Files
ai_site/platform/internal/dbsync/pull.go
whm b04b180d30 feat: harden loose-offline sync for user JWT, schema, and console ops
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>
2026-08-05 09:47:35 +08:00

255 lines
7.4 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"
"database/sql"
"fmt"
"strings"
"time"
)
const (
PullModeBootstrap = "bootstrap" // 全量分页灌库
PullModeRows = "rows" // 按 row_pks 取行
PullModePKs = "pks" // 只列主键(客户端本地 diff
)
const DefaultPullLimit = 200
const MaxPullLimit = 500
// PullRequest 形态 B平台只读线上 A把行/主键返回给本机 agent 写入 B。
type PullRequest struct {
Mode string `json:"mode"` // bootstrap | rows | pks空则 bootstrap
Table string `json:"table"`
RowPKs []string `json:"row_pks"`
AfterPK string `json:"after_pk"` // 分页游标(字典序)
Limit int `json:"limit"`
OnlineDBID string `json:"online_db_id"`
}
// PullItem 下行一条(客户端对本机 B upsert建议 WithApplying 防回声)。
type PullItem struct {
Table string `json:"table"`
Op string `json:"op"` // upsert
RowPK string `json:"row_pk"`
Row map[string]any `json:"row,omitempty"`
Version int64 `json:"version"`
}
// PullResult 下行响应。
type PullResult struct {
OK bool `json:"ok"`
Mode string `json:"mode"`
Table string `json:"table"`
PKColumn string `json:"pk_column"`
Columns []string `json:"columns,omitempty"` // 表结构;空表时也返回,便于本机建空表
Items []PullItem `json:"items,omitempty"`
PKs []string `json:"pks,omitempty"`
NextAfterPK string `json:"next_after_pk,omitempty"`
HasMore bool `json:"has_more"`
Message string `json:"message,omitempty"`
}
// PullFromRemote 从通道 remote线上 A读出数据供本机 agent 写入 B。
func PullFromRemote(ctx context.Context, ch *Channel, store *FileStore, req PullRequest) (*PullResult, error) {
if ch == nil {
return nil, fmt.Errorf("channel is nil")
}
mode := strings.ToLower(strings.TrimSpace(req.Mode))
if mode == "" {
mode = PullModeBootstrap
}
switch mode {
case PullModeBootstrap, PullModeRows, PullModePKs:
default:
return nil, fmt.Errorf("unsupported mode: %s (use bootstrap|rows|pks)", req.Mode)
}
table := strings.TrimSpace(req.Table)
if table == "" {
return nil, fmt.Errorf("table required")
}
if !tableInChannel(ch, table) {
return nil, fmt.Errorf("table %s 不在通道白名单", table)
}
pkCol := "id"
if ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[table]) != "" {
pkCol = ch.PKColumns[table]
}
limit := req.Limit
if limit <= 0 {
limit = DefaultPullLimit
}
if limit > MaxPullLimit {
limit = MaxPullLimit
}
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)
}
exists, err := tableExists(ctx, db, ch.Remote.Driver, table)
if err != nil {
return nil, Retryablef("table exists: %v", err)
}
if !exists {
return nil, fmt.Errorf("table %s not found on online A; ensure-schema or push first", table)
}
cols, err := listColumns(ctx, db, ch.Remote.Driver, table)
if err != nil {
return nil, Retryablef("columns: %v", err)
}
res := &PullResult{
OK: true,
Mode: mode,
Table: table,
PKColumn: pkCol,
Columns: cols,
}
switch mode {
case PullModePKs:
pks, next, more, err := listPKsPage(ctx, db, ch.Remote.Driver, table, pkCol, strings.TrimSpace(req.AfterPK), limit)
if err != nil {
return nil, Retryablef("list pks: %v", err)
}
res.PKs = pks
res.NextAfterPK = next
res.HasMore = more
res.Message = "pks from online A; client diffs then mode=rows"
case PullModeRows:
if len(req.RowPKs) == 0 {
return nil, fmt.Errorf("row_pks required for mode=rows")
}
if len(req.RowPKs) > MaxPullLimit {
return nil, fmt.Errorf("row_pks limit %d", MaxPullLimit)
}
items, err := fetchRowsByPKs(ctx, db, ch.Remote.Driver, table, pkCol, req.RowPKs)
if err != nil {
return nil, err
}
res.Items = items
res.Message = "rows from online A; apply to local B with applying flag"
default: // bootstrap
items, next, more, err := fetchRowsPage(ctx, db, ch.Remote.Driver, table, pkCol, strings.TrimSpace(req.AfterPK), limit)
if err != nil {
return nil, err
}
res.Items = items
res.NextAfterPK = next
res.HasMore = more
if len(items) == 0 && afterEmpty(req.AfterPK) {
res.Message = "empty table on online A; use columns to CREATE TABLE IF NOT EXISTS on local B"
} else {
res.Message = "bootstrap page from online A; loop until has_more=false; create local table from columns if missing"
}
}
if store != nil && (len(res.Items) > 0 || len(res.PKs) > 0) {
n := int64(len(res.Items))
if n == 0 {
n = int64(len(res.PKs))
}
_ = store.PatchStats(ch.ID, func(c *Channel) {
c.Stats.PulledOK += n
})
}
return res, nil
}
func listPKsPage(ctx context.Context, db *sql.DB, driver Driver, table, pkCol, afterPK string, limit int) (pks []string, next string, more bool, err error) {
q, args := buildPKPageQuery(driver, table, pkCol, afterPK, limit+1)
rows, err := db.QueryContext(ctx, q, args...)
if err != nil {
return nil, "", false, err
}
defer rows.Close()
for rows.Next() {
var v any
if err := rows.Scan(&v); err != nil {
return nil, "", false, err
}
pks = append(pks, fmt.Sprint(v))
}
if err := rows.Err(); err != nil {
return nil, "", false, err
}
if len(pks) > limit {
more = true
pks = pks[:limit]
}
if len(pks) > 0 {
next = pks[len(pks)-1]
}
return pks, next, more, nil
}
func fetchRowsPage(ctx context.Context, db *sql.DB, driver Driver, table, pkCol, afterPK string, limit int) (items []PullItem, next string, more bool, err error) {
pks, next, more, err := listPKsPage(ctx, db, driver, table, pkCol, afterPK, limit)
if err != nil {
return nil, "", false, Retryablef("list pks: %v", err)
}
items, err = fetchRowsByPKs(ctx, db, driver, table, pkCol, pks)
if err != nil {
return nil, "", false, err
}
return items, next, more, nil
}
func fetchRowsByPKs(ctx context.Context, db *sql.DB, driver Driver, table, pkCol string, pks []string) ([]PullItem, error) {
items := make([]PullItem, 0, len(pks))
for _, pk := range pks {
pk = strings.TrimSpace(pk)
if pk == "" {
continue
}
_, row, err := FetchRowJSON(ctx, db, driver, table, pkCol, pk)
if err == sql.ErrNoRows {
continue
}
if err != nil {
return nil, Retryablef("fetch row: %v", err)
}
ver, has, _ := GetMetaVersion(ctx, db, driver, table, pk)
if !has || ver <= 0 {
ver = time.Now().UnixNano()
}
items = append(items, PullItem{
Table: table,
Op: "upsert",
RowPK: pk,
Row: row,
Version: ver,
})
}
return items, nil
}
func afterEmpty(afterPK string) bool {
return strings.TrimSpace(afterPK) == ""
}
func buildPKPageQuery(driver Driver, table, pkCol, afterPK string, limit int) (string, []any) {
t := quoteIdent(driver, table)
p := quoteIdent(driver, pkCol)
switch driver {
case DriverPostgres:
if afterPK != "" {
return fmt.Sprintf(`SELECT %s FROM %s WHERE %s > $1 ORDER BY %s ASC LIMIT $2`, p, t, p, p), []any{afterPK, limit}
}
return fmt.Sprintf(`SELECT %s FROM %s ORDER BY %s ASC LIMIT $1`, p, t, p), []any{limit}
default:
if afterPK != "" {
return fmt.Sprintf(`SELECT %s FROM %s WHERE %s > ? ORDER BY %s ASC LIMIT ?`, p, t, p, p), []any{afterPK, limit}
}
return fmt.Sprintf(`SELECT %s FROM %s ORDER BY %s ASC LIMIT ?`, p, t, p), []any{limit}
}
}