Files
ai_site/platform/internal/dbsync/manager.go
2026-07-31 10:19:22 +08:00

307 lines
7.4 KiB
Go

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 {
if ch.Enabled {
_ = 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 = ""
})
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)
}
_ = 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++
// 约每分钟主键对账一次,补漏(漏投递 / 触发器未装时的存量差)
if ticks%120 == 0 && ch.Direction == DirBidirectional {
if _, rerr := ReconcileChannel(ctx, ch); rerr != nil {
log.Printf("dbsync reconcile %s: %v", id, rerr)
}
}
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 {
switch ch.ConflictPolicy {
case PolicyLWWTarget:
done = append(done, r.ID)
continue
case PolicyLWWSource:
// fallthrough apply
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 {
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}
}