Files
ai_site/platform/internal/dbsync/checkpoint.go
whm 680d2b8cde feat: add sync checkpoint restore for last successful sync
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>
2026-08-06 11:49:00 +08:00

454 lines
12 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 (
"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()
}