Capture rolling online DB snapshots after push/drain/reconcile and expose SyncPage 数据恢复 plus checkpoint/restore APIs; mark Z12h and restore done in coop docs. Co-authored-by: Cursor <cursoragent@cursor.com>
454 lines
12 KiB
Go
454 lines
12 KiB
Go
package dbsync
|
||
|
||
import (
|
||
"compress/gzip"
|
||
"context"
|
||
"database/sql"
|
||
"encoding/json"
|
||
"fmt"
|
||
"io"
|
||
"log"
|
||
"os"
|
||
"path/filepath"
|
||
"strings"
|
||
"sync"
|
||
"time"
|
||
)
|
||
|
||
const (
|
||
checkpointDebounce = 30 * time.Second
|
||
checkpointMaxRows = 500_000 // 单次快照行数上限,避免撑爆磁盘
|
||
checkpointFileLatest = "latest.json.gz"
|
||
checkpointFilePrev = "previous.json.gz"
|
||
checkpointFileMeta = "meta.json"
|
||
)
|
||
|
||
// CheckpointSlotMeta 一代快照摘要(不读全量)。
|
||
type CheckpointSlotMeta struct {
|
||
SyncedAt time.Time `json:"synced_at"`
|
||
Source string `json:"source"`
|
||
TableCount int `json:"table_count"`
|
||
RowCount int `json:"row_count"`
|
||
}
|
||
|
||
// CheckpointMeta 通道快照目录摘要。
|
||
type CheckpointMeta struct {
|
||
ChannelID string `json:"channel_id"`
|
||
Latest *CheckpointSlotMeta `json:"latest,omitempty"`
|
||
Previous *CheckpointSlotMeta `json:"previous,omitempty"`
|
||
}
|
||
|
||
// CheckpointTable 单表快照。
|
||
type CheckpointTable struct {
|
||
PKColumn string `json:"pk_column"`
|
||
Rows map[string]map[string]any `json:"rows"` // pk -> row
|
||
}
|
||
|
||
// CheckpointPayload 全量快照内容。
|
||
type CheckpointPayload struct {
|
||
ChannelID string `json:"channel_id"`
|
||
SyncedAt time.Time `json:"synced_at"`
|
||
Source string `json:"source"`
|
||
Tables map[string]CheckpointTable `json:"tables"`
|
||
}
|
||
|
||
// RestoreResult 恢复结果。
|
||
type RestoreResult struct {
|
||
OK bool `json:"ok"`
|
||
Which string `json:"which"`
|
||
SyncedAt time.Time `json:"synced_at"`
|
||
Source string `json:"source,omitempty"`
|
||
Tables int `json:"tables"`
|
||
Upserted int `json:"upserted"`
|
||
Deleted int `json:"deleted"`
|
||
RestoredAt time.Time `json:"restored_at"`
|
||
}
|
||
|
||
var (
|
||
cpSchedMu sync.Mutex
|
||
cpPending = map[string]*time.Timer{}
|
||
)
|
||
|
||
func (s *FileStore) Dir() string {
|
||
if s == nil {
|
||
return "./data/dbsync"
|
||
}
|
||
return s.dir
|
||
}
|
||
|
||
func checkpointDir(store *FileStore, channelID string) string {
|
||
return filepath.Join(store.Dir(), "checkpoints", strings.TrimSpace(channelID))
|
||
}
|
||
|
||
// ScheduleCheckpoint 防抖后写入线上库快照;失败仅打日志。
|
||
func ScheduleCheckpoint(store *FileStore, ch *Channel, source string) {
|
||
if store == nil || ch == nil || strings.TrimSpace(ch.ID) == "" {
|
||
return
|
||
}
|
||
id := strings.TrimSpace(ch.ID)
|
||
chCopy := *ch
|
||
src := strings.TrimSpace(source)
|
||
if src == "" {
|
||
src = "sync"
|
||
}
|
||
cpSchedMu.Lock()
|
||
defer cpSchedMu.Unlock()
|
||
if t, ok := cpPending[id]; ok {
|
||
t.Stop()
|
||
}
|
||
cpPending[id] = time.AfterFunc(checkpointDebounce, func() {
|
||
cpSchedMu.Lock()
|
||
delete(cpPending, id)
|
||
cpSchedMu.Unlock()
|
||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
|
||
defer cancel()
|
||
if err := SaveCheckpoint(ctx, store, &chCopy, src); err != nil {
|
||
log.Printf("dbsync checkpoint channel=%s source=%s: %v", id, src, err)
|
||
}
|
||
})
|
||
}
|
||
|
||
// SaveCheckpoint 导出通道 Remote 业务表,轮转 previous ← latest。
|
||
func SaveCheckpoint(ctx context.Context, store *FileStore, ch *Channel, source string) error {
|
||
if store == nil || ch == nil {
|
||
return fmt.Errorf("store/channel required")
|
||
}
|
||
dir := checkpointDir(store, ch.ID)
|
||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||
return err
|
||
}
|
||
payload, err := dumpRemoteCheckpoint(ctx, ch, source)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
latestPath := filepath.Join(dir, checkpointFileLatest)
|
||
prevPath := filepath.Join(dir, checkpointFilePrev)
|
||
// 轮转:现有 latest → previous
|
||
if _, err := os.Stat(latestPath); err == nil {
|
||
_ = os.Remove(prevPath)
|
||
if err := os.Rename(latestPath, prevPath); err != nil {
|
||
// Windows 上目标存在时 Rename 可能失败;已删 prev 再试
|
||
_ = os.Remove(prevPath)
|
||
if err2 := os.Rename(latestPath, prevPath); err2 != nil {
|
||
return fmt.Errorf("rotate checkpoint: %w", err2)
|
||
}
|
||
}
|
||
}
|
||
tmp := latestPath + ".tmp"
|
||
if err := writeCheckpointGzip(tmp, payload); err != nil {
|
||
_ = os.Remove(tmp)
|
||
return err
|
||
}
|
||
if err := os.Rename(tmp, latestPath); err != nil {
|
||
_ = os.Remove(latestPath)
|
||
if err2 := os.Rename(tmp, latestPath); err2 != nil {
|
||
_ = os.Remove(tmp)
|
||
return err2
|
||
}
|
||
}
|
||
meta := CheckpointMeta{ChannelID: ch.ID}
|
||
meta.Latest = slotMetaFromPayload(payload)
|
||
if prev, err := loadCheckpointPayload(prevPath); err == nil && prev != nil {
|
||
meta.Previous = slotMetaFromPayload(prev)
|
||
}
|
||
return writeCheckpointMeta(dir, meta)
|
||
}
|
||
|
||
func slotMetaFromPayload(p *CheckpointPayload) *CheckpointSlotMeta {
|
||
if p == nil {
|
||
return nil
|
||
}
|
||
rows := 0
|
||
for _, t := range p.Tables {
|
||
rows += len(t.Rows)
|
||
}
|
||
return &CheckpointSlotMeta{
|
||
SyncedAt: p.SyncedAt,
|
||
Source: p.Source,
|
||
TableCount: len(p.Tables),
|
||
RowCount: rows,
|
||
}
|
||
}
|
||
|
||
func dumpRemoteCheckpoint(ctx context.Context, ch *Channel, source string) (*CheckpointPayload, error) {
|
||
db, err := AcquireRemote(ch.Remote.Driver, ch.Remote.DSN)
|
||
if err != nil {
|
||
return nil, wrapOpenRemote(err)
|
||
}
|
||
names, err := listTablesForInspect(ctx, db, ch.Remote.Driver, false)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
out := &CheckpointPayload{
|
||
ChannelID: ch.ID,
|
||
SyncedAt: time.Now().UTC(),
|
||
Source: source,
|
||
Tables: make(map[string]CheckpointTable, len(names)),
|
||
}
|
||
totalRows := 0
|
||
for _, table := range names {
|
||
pkCol := "id"
|
||
if ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[table]) != "" {
|
||
pkCol = strings.TrimSpace(ch.PKColumns[table])
|
||
}
|
||
cols, err := listColumns(ctx, db, ch.Remote.Driver, table)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("columns %s: %w", table, err)
|
||
}
|
||
rows, err := fetchAllRows(ctx, db, ch.Remote.Driver, table, cols)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("dump %s: %w", table, err)
|
||
}
|
||
m := make(map[string]map[string]any, len(rows))
|
||
for _, row := range rows {
|
||
pk := fmt.Sprint(row[pkCol])
|
||
if pk == "" || pk == "<nil>" {
|
||
continue
|
||
}
|
||
m[pk] = row
|
||
}
|
||
totalRows += len(m)
|
||
if totalRows > checkpointMaxRows {
|
||
return nil, fmt.Errorf("checkpoint too large: >%d rows", checkpointMaxRows)
|
||
}
|
||
out.Tables[table] = CheckpointTable{PKColumn: pkCol, Rows: m}
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
func fetchAllRows(ctx context.Context, db *sql.DB, driver Driver, table string, cols []string) ([]map[string]any, error) {
|
||
if len(cols) == 0 {
|
||
return nil, nil
|
||
}
|
||
q := fmt.Sprintf(`SELECT * FROM %s`, quoteIdent(driver, table))
|
||
rows, err := db.QueryContext(ctx, q)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
colNames, err := rows.Columns()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
var out []map[string]any
|
||
for rows.Next() {
|
||
raw := make([]any, len(colNames))
|
||
ptrs := make([]any, len(colNames))
|
||
for i := range raw {
|
||
ptrs[i] = &raw[i]
|
||
}
|
||
if err := rows.Scan(ptrs...); err != nil {
|
||
return nil, err
|
||
}
|
||
m := make(map[string]any, len(colNames))
|
||
for i, c := range colNames {
|
||
m[c] = normalizeValue(raw[i])
|
||
}
|
||
out = append(out, m)
|
||
}
|
||
return out, rows.Err()
|
||
}
|
||
|
||
func writeCheckpointGzip(path string, payload *CheckpointPayload) error {
|
||
f, err := os.Create(path)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
defer f.Close()
|
||
zw := gzip.NewWriter(f)
|
||
enc := json.NewEncoder(zw)
|
||
if err := enc.Encode(payload); err != nil {
|
||
_ = zw.Close()
|
||
return err
|
||
}
|
||
if err := zw.Close(); err != nil {
|
||
return err
|
||
}
|
||
return f.Close()
|
||
}
|
||
|
||
func loadCheckpointPayload(path string) (*CheckpointPayload, error) {
|
||
f, err := os.Open(path)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer f.Close()
|
||
zr, err := gzip.NewReader(f)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer zr.Close()
|
||
var p CheckpointPayload
|
||
if err := json.NewDecoder(zr).Decode(&p); err != nil && err != io.EOF {
|
||
return nil, err
|
||
}
|
||
return &p, nil
|
||
}
|
||
|
||
func writeCheckpointMeta(dir string, meta CheckpointMeta) error {
|
||
b, err := json.MarshalIndent(meta, "", " ")
|
||
if err != nil {
|
||
return err
|
||
}
|
||
return os.WriteFile(filepath.Join(dir, checkpointFileMeta), b, 0o644)
|
||
}
|
||
|
||
// LoadCheckpointMeta 读取摘要;无快照时返回空 meta(不报错)。
|
||
func LoadCheckpointMeta(store *FileStore, channelID string) (*CheckpointMeta, error) {
|
||
if store == nil || strings.TrimSpace(channelID) == "" {
|
||
return &CheckpointMeta{}, nil
|
||
}
|
||
dir := checkpointDir(store, channelID)
|
||
b, err := os.ReadFile(filepath.Join(dir, checkpointFileMeta))
|
||
if err != nil {
|
||
if os.IsNotExist(err) {
|
||
// 尝试从文件推断
|
||
meta := &CheckpointMeta{ChannelID: channelID}
|
||
if p, e := loadCheckpointPayload(filepath.Join(dir, checkpointFileLatest)); e == nil {
|
||
meta.Latest = slotMetaFromPayload(p)
|
||
}
|
||
if p, e := loadCheckpointPayload(filepath.Join(dir, checkpointFilePrev)); e == nil {
|
||
meta.Previous = slotMetaFromPayload(p)
|
||
}
|
||
return meta, nil
|
||
}
|
||
return nil, err
|
||
}
|
||
var meta CheckpointMeta
|
||
if err := json.Unmarshal(b, &meta); err != nil {
|
||
return nil, err
|
||
}
|
||
meta.ChannelID = channelID
|
||
return &meta, nil
|
||
}
|
||
|
||
// LoadCheckpoint 加载 latest 或 previous。
|
||
func LoadCheckpoint(store *FileStore, channelID, which string) (*CheckpointPayload, error) {
|
||
which = strings.TrimSpace(which)
|
||
if which == "" {
|
||
which = "latest"
|
||
}
|
||
if which != "latest" && which != "previous" {
|
||
return nil, fmt.Errorf("which 须为 latest 或 previous")
|
||
}
|
||
name := checkpointFileLatest
|
||
if which == "previous" {
|
||
name = checkpointFilePrev
|
||
}
|
||
path := filepath.Join(checkpointDir(store, channelID), name)
|
||
p, err := loadCheckpointPayload(path)
|
||
if err != nil {
|
||
if os.IsNotExist(err) {
|
||
return nil, fmt.Errorf("无可用快照(%s)", which)
|
||
}
|
||
return nil, err
|
||
}
|
||
return p, nil
|
||
}
|
||
|
||
// RestoreCheckpoint 将线上库恢复为指定快照(upsert + 删除快照外 PK)。
|
||
func RestoreCheckpoint(ctx context.Context, store *FileStore, ch *Channel, which string) (*RestoreResult, error) {
|
||
if store == nil || ch == nil {
|
||
return nil, fmt.Errorf("store/channel required")
|
||
}
|
||
payload, err := LoadCheckpoint(store, ch.ID, which)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
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, err
|
||
}
|
||
verBase := time.Now().UnixNano()
|
||
upserted, deleted := 0, 0
|
||
for table, ct := range payload.Tables {
|
||
pkCol := strings.TrimSpace(ct.PKColumn)
|
||
if pkCol == "" {
|
||
pkCol = "id"
|
||
if ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[table]) != "" {
|
||
pkCol = strings.TrimSpace(ch.PKColumns[table])
|
||
}
|
||
}
|
||
// 先 upsert 快照行
|
||
i := 0
|
||
for pk, row := range ct.Rows {
|
||
if len(row) == 0 {
|
||
continue
|
||
}
|
||
if err := EnsureTableFromRow(ctx, db, ch.Remote.Driver, table, pkCol, row); err != nil {
|
||
return nil, fmt.Errorf("ensure %s: %w", table, err)
|
||
}
|
||
b, err := json.Marshal(row)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
ver := verBase + int64(i)
|
||
i++
|
||
if err := ApplyChange(ctx, db, ch.Remote.Driver, table, pkCol, "upsert", string(b), ver); err != nil {
|
||
return nil, fmt.Errorf("upsert %s pk=%s: %w", table, pk, err)
|
||
}
|
||
upserted++
|
||
}
|
||
// 删快照中不存在的行(含空表:清空线上多余行)
|
||
livePKs, err := listTablePKs(ctx, db, ch.Remote.Driver, table, pkCol)
|
||
if err != nil {
|
||
// 表可能尚不存在且快照也空
|
||
if len(ct.Rows) == 0 {
|
||
continue
|
||
}
|
||
return nil, fmt.Errorf("list pk %s: %w", table, err)
|
||
}
|
||
for _, pk := range livePKs {
|
||
if _, ok := ct.Rows[pk]; ok {
|
||
continue
|
||
}
|
||
delPayload, _ := json.Marshal(map[string]any{pkCol: pk})
|
||
ver := verBase + int64(i)
|
||
i++
|
||
if err := ApplyChange(ctx, db, ch.Remote.Driver, table, pkCol, "delete", string(delPayload), ver); err != nil {
|
||
return nil, fmt.Errorf("delete %s pk=%s: %w", table, pk, err)
|
||
}
|
||
deleted++
|
||
}
|
||
}
|
||
// 快照里没有、线上多出来的业务表:不自动 DROP(避免误伤);仅对快照内表做行级 prune
|
||
which = strings.TrimSpace(which)
|
||
if which == "" {
|
||
which = "latest"
|
||
}
|
||
return &RestoreResult{
|
||
OK: true,
|
||
Which: which,
|
||
SyncedAt: payload.SyncedAt,
|
||
Source: payload.Source,
|
||
Tables: len(payload.Tables),
|
||
Upserted: upserted,
|
||
Deleted: deleted,
|
||
RestoredAt: time.Now().UTC(),
|
||
}, nil
|
||
}
|
||
|
||
func listTablePKs(ctx context.Context, db *sql.DB, driver Driver, table, pkCol string) ([]string, error) {
|
||
q := fmt.Sprintf(`SELECT %s FROM %s`, quoteIdent(driver, pkCol), quoteIdent(driver, table))
|
||
rows, err := db.QueryContext(ctx, q)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
var out []string
|
||
for rows.Next() {
|
||
var v any
|
||
if err := rows.Scan(&v); err != nil {
|
||
return nil, err
|
||
}
|
||
s := fmt.Sprint(normalizeValue(v))
|
||
if s != "" && s != "<nil>" {
|
||
out = append(out, s)
|
||
}
|
||
}
|
||
return out, rows.Err()
|
||
}
|