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>
255 lines
7.4 KiB
Go
255 lines
7.4 KiB
Go
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}
|
||
}
|
||
}
|