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} } }