Files
ai_site/platform/internal/dbsync/manager.go
whm cb76824e94 fix: start system default sync channels by default
Create and restart paths enable IsSystemDefault channels; SyncPage auto-starts after save.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-05 17:50:27 +08:00

385 lines
9.9 KiB
Go
Raw Permalink 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 (
"context"
"database/sql"
"fmt"
"log"
"sync"
"time"
)
type Manager struct {
store *FileStore
mu sync.Mutex
runners map[string]context.CancelFunc
}
func NewManager(store *FileStore) *Manager {
return &Manager{store: store, runners: map[string]context.CancelFunc{}}
}
func (m *Manager) Store() *FileStore { return m.store }
func (m *Manager) StartAll(ctx context.Context) {
list, err := m.store.ListChannels()
if err != nil {
log.Printf("dbsync list: %v", err)
return
}
for _, ch := range list {
// 公司默认同步通道默认运行中(即使历史记录曾为 Enabled=false
if ch.Enabled || ch.IsSystemDefault {
_ = m.StartChannel(ch.ID)
}
}
go func() {
<-ctx.Done()
m.StopAll()
}()
}
func (m *Manager) StartChannel(id string) error {
m.mu.Lock()
defer m.mu.Unlock()
if _, ok := m.runners[id]; ok {
return nil
}
ch, err := m.store.GetChannel(id)
if err != nil {
return err
}
runCtx, cancel := context.WithCancel(context.Background())
m.runners[id] = cancel
go m.loop(runCtx, ch.ID)
_ = m.store.PatchStats(id, func(c *Channel) {
c.Enabled = true
c.LastError = ""
})
// 预热 remote避免首个 agent push 冷开 SQLite 触发超时
go func(driver Driver, dsn string) {
if err := WarmRemote(driver, dsn); err != nil {
log.Printf("dbsync warm remote %s: %v", id, err)
}
}(ch.Remote.Driver, ch.Remote.DSN)
return nil
}
func (m *Manager) StopChannel(id string) {
m.mu.Lock()
defer m.mu.Unlock()
if cancel, ok := m.runners[id]; ok {
cancel()
delete(m.runners, id)
}
if ch, err := m.store.GetChannel(id); err == nil {
InvalidateRemote(ch.Remote.Driver, ch.Remote.DSN)
}
_ = m.store.PatchStats(id, func(c *Channel) { c.Enabled = false })
}
func (m *Manager) StopAll() {
m.mu.Lock()
defer m.mu.Unlock()
for id, cancel := range m.runners {
cancel()
delete(m.runners, id)
}
}
func (m *Manager) loop(ctx context.Context, id string) {
ticks := 0
for {
ch, err := m.store.GetChannel(id)
if err != nil {
return
}
iv := time.Duration(ch.PollIntervalMS) * time.Millisecond
if iv < 100*time.Millisecond {
iv = 500 * time.Millisecond
}
if err := m.tick(ctx, ch); err != nil {
_ = m.store.PatchStats(id, func(c *Channel) { c.LastError = err.Error() })
log.Printf("dbsync channel %s: %v", id, err)
}
ticks++
// 双向通道自动对账:默认约每 15 分钟一次(限流)
autoEvery := 1800 // poll 500ms → ~15min
if ch.PollIntervalMS > 0 {
autoEvery = int((15 * time.Minute) / (time.Duration(ch.PollIntervalMS) * time.Millisecond))
if autoEvery < 60 {
autoEvery = 60
}
}
if ticks%autoEvery == 0 && ch.Direction == DirBidirectional {
if allow, _ := CanReconcile(ch, 15*time.Minute); allow {
if _, rerr := ReconcileChannel(ctx, ch); rerr != nil {
log.Printf("dbsync reconcile %s: %v", id, rerr)
} else {
_ = m.store.PatchStats(id, func(c *Channel) {
now := time.Now().UTC()
c.LastReconcileAt = &now
})
}
}
}
select {
case <-ctx.Done():
return
case <-time.After(iv):
}
}
}
func (m *Manager) tick(ctx context.Context, ch *Channel) error {
local, err := Open(ch.Local.Driver, ch.Local.DSN)
if err != nil {
return fmt.Errorf("open local: %w", err)
}
defer local.Close()
remote, err := Open(ch.Remote.Driver, ch.Remote.DSN)
if err != nil {
return fmt.Errorf("open remote: %w", err)
}
defer remote.Close()
if err := prepareEndpoint(ctx, local, ch.Local, ch.PKColumns); err != nil {
return fmt.Errorf("prepare local: %w", err)
}
if err := prepareEndpoint(ctx, remote, ch.Remote, ch.PKColumns); err != nil {
return fmt.Errorf("prepare remote: %w", err)
}
var n int
switch ch.Direction {
case DirRemoteToLocal:
n, err = m.drain(ctx, ch, "remote", remote, ch.Remote, local, ch.Local)
case DirBidirectional:
n1, e1 := m.drain(ctx, ch, "local", local, ch.Local, remote, ch.Remote)
n2, e2 := m.drain(ctx, ch, "remote", remote, ch.Remote, local, ch.Local)
n = n1 + n2
if e1 != nil {
err = e1
} else {
err = e2
}
default: // local_to_remote
n, err = m.drain(ctx, ch, "local", local, ch.Local, remote, ch.Remote)
}
now := time.Now().UTC()
_ = m.store.PatchStats(ch.ID, func(c *Channel) {
c.LastSyncAt = &now
c.Stats.LastBatch = n
if err != nil {
c.LastError = err.Error()
} else {
c.LastError = ""
}
})
return err
}
func prepareEndpoint(ctx context.Context, db *sql.DB, ep Endpoint, pks map[string]string) error {
if err := EnsureOutbox(ctx, db, ep.Driver); err != nil {
return err
}
if err := EnsureMeta(ctx, db, ep.Driver); err != nil {
return err
}
for _, t := range ep.Tables {
pk := "id"
if pks != nil && pks[t] != "" {
pk = pks[t]
}
if err := InstallTriggers(ctx, db, ep.Driver, t, pk); err != nil {
return err
}
}
return nil
}
func (m *Manager) drain(ctx context.Context, ch *Channel, sourceName string, src *sql.DB, srcEp Endpoint, dst *sql.DB, dstEp Endpoint) (int, error) {
rows, err := PollOutbox(ctx, src, srcEp.Driver, 100)
if err != nil {
return 0, err
}
if len(rows) == 0 {
return 0, nil
}
var done []int64
okCount := 0
for _, r := range rows {
pkCol := ch.PKColumns[r.TableName]
if pkCol == "" {
pkCol = "id"
}
payload := r.Payload
if r.Op != "delete" && (payload == "" || payload == "{}") {
js, _, ferr := FetchRowJSON(ctx, src, srcEp.Driver, r.TableName, pkCol, r.RowPK)
if ferr == sql.ErrNoRows {
// 行已删,改 delete
r.Op = "delete"
payload = fmt.Sprintf(`{%q:%q}`, pkCol, r.RowPK)
} else if ferr != nil {
_ = m.store.PatchStats(ch.ID, func(c *Channel) { c.Stats.Retries++ })
continue
} else {
payload = js
}
}
if r.Op == "delete" && payload == "" {
payload = fmt.Sprintf(`{%q:%q}`, pkCol, r.RowPK)
}
tgtVer, has, _ := GetMetaVersion(ctx, dst, dstEp.Driver, r.TableName, r.RowPK)
// 幂等:已同步过相同版本 → 跳过(防「多」)
if has && tgtVer == r.Version {
done = append(done, r.ID)
continue
}
if has && tgtVer > r.Version {
policy := ch.ConflictPolicy
if policy == "" {
policy = PolicyLWWSource
}
switch policy {
case PolicyLWWTarget:
loser := payload
winner := SnapshotTargetRow(ctx, dst, dstEp.Driver, r.TableName, pkCol, r.RowPK)
RecordLwwOverride(m.store, LwwOverride{
TenantID: ch.TenantID,
ChannelID: ch.ID,
Table: r.TableName,
RowPK: r.RowPK,
Op: r.Op,
Entry: EntryDrain,
Policy: string(PolicyLWWTarget),
Outcome: OutcomeKept,
LoserPayload: loser,
WinnerPayload: winner,
TargetVer: tgtVer,
SourceVer: r.Version,
})
done = append(done, r.ID)
continue
case PolicyLWWSource:
loser := SnapshotTargetRow(ctx, dst, dstEp.Driver, r.TableName, pkCol, r.RowPK)
if err := ApplyChange(ctx, dst, dstEp.Driver, r.TableName, pkCol, r.Op, payload, r.Version); err != nil {
_ = m.store.PatchStats(ch.ID, func(c *Channel) {
c.Stats.Retries++
c.LastError = err.Error()
})
continue
}
RecordLwwOverride(m.store, LwwOverride{
TenantID: ch.TenantID,
ChannelID: ch.ID,
Table: r.TableName,
RowPK: r.RowPK,
Op: r.Op,
Entry: EntryDrain,
Policy: string(PolicyLWWSource),
Outcome: OutcomeApplied,
LoserPayload: loser,
WinnerPayload: payload,
TargetVer: tgtVer,
SourceVer: r.Version,
})
done = append(done, r.ID)
okCount++
_ = m.store.PatchStats(ch.ID, func(c *Channel) {
if sourceName == "local" {
c.Stats.PushedOK++
} else {
c.Stats.PulledOK++
}
})
continue
default:
_ = m.store.AddConflict(Conflict{
TenantID: ch.TenantID,
ChannelID: ch.ID,
Table: r.TableName,
RowPK: r.RowPK,
Op: r.Op,
Source: sourceName,
Payload: payload,
TargetVer: tgtVer,
SourceVer: r.Version,
Message: "target version newer than source",
})
_ = m.store.PatchStats(ch.ID, func(c *Channel) { c.Stats.Conflicts++ })
done = append(done, r.ID)
continue
}
}
if err := ApplyChange(ctx, dst, dstEp.Driver, r.TableName, pkCol, r.Op, payload, r.Version); err != nil {
_ = m.store.PatchStats(ch.ID, func(c *Channel) {
c.Stats.Retries++
c.LastError = err.Error()
})
continue
}
done = append(done, r.ID)
okCount++
_ = m.store.PatchStats(ch.ID, func(c *Channel) {
if sourceName == "local" {
c.Stats.PushedOK++
} else {
c.Stats.PulledOK++
}
})
}
if err := MarkSynced(ctx, src, srcEp.Driver, done); err != nil {
return okCount, err
}
return okCount, nil
}
// PrepareChannel 连接两端、建 outbox/触发器,供「测试/启用」调用。
func PrepareChannel(ctx context.Context, ch *Channel) error {
if err := ValidateChannelAgainstDB(ctx, ch); err != nil {
return err
}
for _, ep := range []Endpoint{ch.Local, ch.Remote} {
db, err := Open(ep.Driver, ep.DSN)
if err != nil {
return err
}
if err := EnsureOutbox(ctx, db, ep.Driver); err != nil {
_ = db.Close()
return err
}
if err := EnsureMeta(ctx, db, ep.Driver); err != nil {
_ = db.Close()
return err
}
for _, t := range ep.Tables {
pk := "id"
if ch.PKColumns != nil && ch.PKColumns[t] != "" {
pk = ch.PKColumns[t]
}
if err := InstallTriggers(ctx, db, ep.Driver, t, pk); err != nil {
_ = db.Close()
return fmt.Errorf("%s.%s triggers: %w", ep.Driver, t, err)
}
}
_ = db.Close()
}
return nil
}
func TestEndpoint(ctx context.Context, ep Endpoint) TestResult {
db, err := Open(ep.Driver, ep.DSN)
if err != nil {
return TestResult{OK: false, Driver: string(ep.Driver), Message: err.Error()}
}
defer db.Close()
tables, err := ListTables(ctx, db, ep.Driver)
if err != nil {
return TestResult{OK: false, Driver: string(ep.Driver), Message: err.Error()}
}
return TestResult{OK: true, Driver: string(ep.Driver), Message: "connected", Tables: tables}
}