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 == "" { 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 != "" { out = append(out, s) } } return out, rows.Err() }