From 2e04adf054268a0f3da2572970de65c641e19611 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=B0=8F=E6=80=A1?= Date: Sat, 12 Sep 2026 17:57:31 +0800 Subject: [PATCH] =?UTF-8?q?perf(recall):=20=E7=B4=AF=E7=A7=AF=E5=BB=B6?= =?UTF-8?q?=E8=BF=9F=E6=9B=B4=E6=96=B0=E6=B2=BB=E7=90=86=E5=86=99=E6=94=BE?= =?UTF-8?q?=E5=A4=A7=20=E2=80=94=20=E8=AF=BB=E8=B7=AF=E5=BE=84=E9=9B=B6?= =?UTF-8?q?=E9=80=90=E6=9D=A1=E5=86=99=20+=20=E5=8D=95=E4=BA=8B=E5=8A=A1?= =?UTF-8?q?=E6=89=B9=E9=87=8F=E8=90=BD=E7=9B=98=20(t=5F8496e8b6)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 根因: recall 读路径对每条结果同步 lancedb.Update,LanceDB MVCC 每次提交=1 版本 实测 15 版本/分钟 / _versions 17.7G(真实数据 22M) 改动: - Go: 新增 RecallWriteBuffer(窗口合并+阈值触发+优雅退出落盘);recall.go 读路径改 Record();RustLanceDBClient.UpdateRecallBatch - Rust: lancedb_update_batch IPC — 按 delta 分组,每组一次 update(`recall_count + delta` + id IN (...)) 提交 - 单测 6 个 + IPC 端到端版本计数测试;bench_test.go 修 NewReranker 签名失配(阻塞包测试) 实测: 批量 3 条 → 1 个版本;逐条 3 次 → 3 个版本 (lance 不支持 CASE WHEN,已按 delta 分组规避) --- go/cmd/zhiyid/main.go | 3 + go/internal/api/server.go | 5 + .../storage/batch_ipc_integration_test.go | 165 +++++++++++ go/internal/storage/bench_test.go | 10 +- go/internal/storage/lancedb_ipc.go | 33 +++ go/internal/storage/recall.go | 49 +-- go/internal/storage/recall_write_buffer.go | 279 ++++++++++++++++++ .../storage/recall_write_buffer_test.go | 221 ++++++++++++++ rust/src/lancedb_ops.rs | 81 +++++ rust/src/main.rs | 15 + 10 files changed, 836 insertions(+), 25 deletions(-) create mode 100644 go/internal/storage/batch_ipc_integration_test.go create mode 100644 go/internal/storage/recall_write_buffer.go create mode 100644 go/internal/storage/recall_write_buffer_test.go diff --git a/go/cmd/zhiyid/main.go b/go/cmd/zhiyid/main.go index 8feffc5..58bfd47 100644 --- a/go/cmd/zhiyid/main.go +++ b/go/cmd/zhiyid/main.go @@ -11,6 +11,7 @@ import ( "time" "github.com/xiaoxue/memoryweave/internal/api" + "github.com/xiaoxue/memoryweave/internal/storage" ) func main() { @@ -34,6 +35,8 @@ func main() { <-sigCh log.Println("[zhiyid] 收到关闭信号,正在退出...") + // 累积延迟更新:把最后一个窗口的召回元数据落盘(防止统计增量丢失) + storage.StopRecallWriteBuffer() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if err := srv.Shutdown(ctx); err != nil { diff --git a/go/internal/api/server.go b/go/internal/api/server.go index bae8fcc..9d6a96c 100644 --- a/go/internal/api/server.go +++ b/go/internal/api/server.go @@ -52,6 +52,11 @@ func NewServer() http.Handler { // 存储后端 ldb := initStorageBackend(os.Getenv("STORAGE_BACKEND"), emb) + // 召回元数据「累积延迟更新」缓冲(2026-09-12 优化B) + // 读路径不再逐条写 LanceDB(MVCC 每写一行一个版本 → _versions 膨胀根因), + // 改为窗口内合并 + 单事务批量提交(每批 1 个版本)。 + storage.InitRecallWriteBuffer(ldb) + // 启动时数据目录一致性检查(防止路径混乱导致读取废弃数据) runStartupChecks(os.Getenv("STORAGE_BACKEND")) diff --git a/go/internal/storage/batch_ipc_integration_test.go b/go/internal/storage/batch_ipc_integration_test.go new file mode 100644 index 0000000..bcf1dbf --- /dev/null +++ b/go/internal/storage/batch_ipc_integration_test.go @@ -0,0 +1,165 @@ +package storage + +import ( + "encoding/binary" + "encoding/json" + "fmt" + "io" + "net" + "os" + "testing" + "time" +) + +// 端到端验证 lancedb_update_batch IPC(需要真实 sidecar socket)。 +// 用法: +// +// ZHIYI_TEST_IPC_SOCK=/tmp/ipc-batch-test.sock go test ./internal/storage/ -run TestBatchIPC -v +// +// 验证两点:① 批量调用按 delta 正确累加 recall_count;② 整批只产生 1 个 LanceDB 版本(由脚本比对 +// _versions 目录文件数确认,见 AC 记录)。 +func ipcCall(t *testing.T, sock string, req map[string]any) map[string]any { + t.Helper() + conn, err := net.DialTimeout("unix", sock, 5*time.Second) + if err != nil { + t.Fatalf("dial %s: %v", sock, err) + } + defer conn.Close() + body, _ := json.Marshal(req) + var lenBuf [4]byte + binary.BigEndian.PutUint32(lenBuf[:], uint32(len(body))) + if _, err := conn.Write(lenBuf[:]); err != nil { + t.Fatalf("write len: %v", err) + } + if _, err := conn.Write(body); err != nil { + t.Fatalf("write body: %v", err) + } + if _, err := io.ReadFull(conn, lenBuf[:]); err != nil { + t.Fatalf("read len: %v", err) + } + respBuf := make([]byte, binary.BigEndian.Uint32(lenBuf[:])) + if _, err := io.ReadFull(conn, respBuf); err != nil { + t.Fatalf("read body: %v", err) + } + var resp map[string]any + if err := json.Unmarshal(respBuf, &resp); err != nil { + t.Fatalf("unmarshal resp: %v", err) + } + if resp["status"] != "ok" { + t.Fatalf("IPC status=%v detail=%v", resp["status"], resp["error_detail"]) + } + return resp +} + +func scanRecallCount(t *testing.T, sock, wantID string) int64 { + t.Helper() + resp := ipcCall(t, sock, map[string]any{"type": "lancedb_scan", "limit": 2000}) + var records []struct { + ID string `json:"id"` + RecallCount int64 `json:"recall_count"` + } + if err := json.Unmarshal([]byte(resp["report_json"].(string)), &records); err != nil { + t.Fatalf("unmarshal scan: %v", err) + } + for _, r := range records { + if r.ID == wantID { + return r.RecallCount + } + } + t.Fatalf("id %s 不在 scan 结果中(%d 条)", wantID, len(records)) + return 0 +} + +func TestBatchIPCVersionAccounting(t *testing.T) { + sock := os.Getenv("ZHIYI_TEST_IPC_SOCK") + dir := os.Getenv("ZHIYI_TEST_DATA_DIR") + if sock == "" || dir == "" { + t.Skip("ZHIYI_TEST_IPC_SOCK / ZHIYI_TEST_DATA_DIR 未设置,跳过版本计数对比") + } + countVersions := func() int { + entries, err := os.ReadDir(dir + "/memories.lance/_versions") + if err != nil { + t.Fatalf("read versions dir: %v", err) + } + n := 0 + for _, e := range entries { + if !e.IsDir() && len(e.Name()) > 9 && e.Name()[len(e.Name())-9:] == ".manifest" { + n++ + } + } + return n + } + + resp := ipcCall(t, sock, map[string]any{"type": "lancedb_scan", "limit": 20}) + var recs []struct { + ID string `json:"id"` + } + if err := json.Unmarshal([]byte(resp["report_json"].(string)), &recs); err != nil || len(recs) < 6 { + t.Fatalf("样本不足: err=%v n=%d", err, len(recs)) + } + ts := time.Now().Format(time.RFC3339) + + // A) 新路径:一次批量(3 条不同记忆) + v0 := countVersions() + items, _ := json.Marshal([]RecallWriteItem{{ID: recs[0].ID, Delta: 1}, {ID: recs[1].ID, Delta: 1}, {ID: recs[2].ID, Delta: 1}}) + ipcCall(t, sock, map[string]any{"type": "lancedb_update_batch", "items": string(items), "ts": ts}) + v1 := countVersions() + batchVersions := v1 - v0 + + // B) 旧路径:3 次单条 lancedb_update(recall 读路径原行为) + fields, _ := json.Marshal([]map[string]string{ + {"column": "recall_count", "value": "1"}, + {"column": "last_recalled_at", "value": "'" + ts + "'"}, + {"column": "freshness", "value": "'verified'"}, + }) + for i := 3; i < 6; i++ { + ipcCall(t, sock, map[string]any{"type": "lancedb_update", "table": "memories", "id": recs[i].ID, "fields": string(fields)}) + } + v2 := countVersions() + perItemVersions := v2 - v1 + + fmt.Printf("VERSION_ACCOUNTING: 批量3条 → %d 个版本 | 逐条3次 → %d 个版本\n", batchVersions, perItemVersions) + if perItemVersions <= batchVersions { + t.Fatalf("逐条(%d) 应比批量(%d) 产生更多版本", perItemVersions, batchVersions) + } +} + +func TestBatchIPCUpdateRecallBatch(t *testing.T) { + sock := os.Getenv("ZHIYI_TEST_IPC_SOCK") + if sock == "" { + t.Skip("ZHIYI_TEST_IPC_SOCK 未设置,跳过 IPC 端到端验证") + } + + // 取一条真实记忆做样本 + resp := ipcCall(t, sock, map[string]any{"type": "lancedb_scan", "limit": 50}) + var recs []struct { + ID string `json:"id"` + RecallCount int64 `json:"recall_count"` + } + if err := json.Unmarshal([]byte(resp["report_json"].(string)), &recs); err != nil || len(recs) == 0 { + t.Fatalf("scan 无样本: err=%v n=%d", err, len(recs)) + } + target := recs[0].ID + before := scanRecallCount(t, sock, target) + + // 批量:同一条 delta=3 + 另一条 delta=1(验证合并语义与多行单事务) + second := target + if len(recs) > 1 { + second = recs[1].ID + } + items, _ := json.Marshal([]RecallWriteItem{{ID: target, Delta: 3}, {ID: second, Delta: 1}}) + ts := time.Now().Format(time.RFC3339) + out := ipcCall(t, sock, map[string]any{ + "type": "lancedb_update_batch", + "table": "memories", + "items": string(items), + "ts": ts, + }) + t.Logf("batch resp: %v", out) + + after := scanRecallCount(t, sock, target) + if after != before+3 { + t.Fatalf("recall_count 期望 %d,实际 %d(delta 未正确累加)", before+3, after) + } + fmt.Printf("BATCH_IPC_OK id=%s recall_count %d → %d\n", target, before, after) +} diff --git a/go/internal/storage/bench_test.go b/go/internal/storage/bench_test.go index 50ba1c1..13a42f7 100644 --- a/go/internal/storage/bench_test.go +++ b/go/internal/storage/bench_test.go @@ -46,10 +46,11 @@ func BenchmarkEmbedder_Batch50(b *testing.B) { // ─── Recall Pipeline 基准 ───────────────────────────────── func BenchmarkRecallPipeline_10Docs(b *testing.B) { + emb := &Embedder{endpoint: "http://localhost:8000/v1/embeddings"} p := NewRecallPipeline( - &Embedder{endpoint: "http://localhost:8000/v1/embeddings"}, + emb, NewMemLanceClient(nil), - NewReranker("http://localhost:8001/rerank"), + NewReranker("http://localhost:8001/rerank", emb), ) b.ResetTimer() for i := 0; i < b.N; i++ { @@ -58,10 +59,11 @@ func BenchmarkRecallPipeline_10Docs(b *testing.B) { } func BenchmarkRecallPipeline_50Docs(b *testing.B) { + emb := &Embedder{endpoint: "http://localhost:8000/v1/embeddings"} p := NewRecallPipeline( - &Embedder{endpoint: "http://localhost:8000/v1/embeddings"}, + emb, NewMemLanceClient(nil), - NewReranker("http://localhost:8001/rerank"), + NewReranker("http://localhost:8001/rerank", emb), ) b.ResetTimer() for i := 0; i < b.N; i++ { diff --git a/go/internal/storage/lancedb_ipc.go b/go/internal/storage/lancedb_ipc.go index ba2ff63..d8a5fa7 100644 --- a/go/internal/storage/lancedb_ipc.go +++ b/go/internal/storage/lancedb_ipc.go @@ -44,6 +44,8 @@ type ipcReq struct { Fields string `json:"fields,omitempty"` MinRecall int `json:"min_recall,omitempty"` QueryLimit int `json:"limit,omitempty"` + Items string `json:"items,omitempty"` // lancedb_update_batch: [{"id":..,"delta":N}] + TS string `json:"ts,omitempty"` // lancedb_update_batch: last_recalled_at } type ipcResp struct { @@ -450,6 +452,37 @@ func (rc *RustLanceDBClient) Update(table, id string, fields map[string]any) err return nil } +// UpdateRecallBatch 单事务批量更新召回元数据(2026-09-12 优化B)。 +// Rust 侧用一次 update 提交(CASE WHEN 表达式覆盖所有命中行)→ 整批只产生 1 个 LanceDB 版本。 +// recall_count 由 Rust 侧读「库现值 + delta」计算,避免 Go 缓存漂移。 +func (rc *RustLanceDBClient) UpdateRecallBatch(table string, items []RecallWriteItem, lastRecalledAt string) (int64, error) { + if len(items) == 0 { + return 0, nil + } + fj, err := json.Marshal(items) + if err != nil { + return 0, err + } + resp, err := rc.rpc(ipcReq{ + Type: "lancedb_update_batch", + Table: table, + Items: string(fj), + TS: lastRecalledAt, + }) + if err != nil { + return 0, err + } + var out struct { + Updated int64 `json:"updated"` + Items int `json:"items"` + } + if resp.ReportJSON != "" { + _ = json.Unmarshal([]byte(resp.ReportJSON), &out) + } + log.Printf("[ipc] UpdateRecallBatch: 合并 %d 条记忆 → 落盘 %d 行(单事务)", len(items), out.Updated) + return out.Updated, nil +} + func (rc *RustLanceDBClient) GetTopByQuality(agentID string, limit int) ([]models.MemoryRecord, error) { _local.mu.RLock() defer _local.mu.RUnlock() diff --git a/go/internal/storage/recall.go b/go/internal/storage/recall.go index e80b380..5c0e902 100644 --- a/go/internal/storage/recall.go +++ b/go/internal/storage/recall.go @@ -73,7 +73,7 @@ func (p *RecallPipeline) Recall(query, namespace string, topK int, diversity flo if len(filtered) == 0 && len(results) > 0 { filtered = results } - go p.incrementRecallCount(results) + p.recordRecallWrites(filtered) return filtered, nil } } @@ -258,26 +258,11 @@ func (p *RecallPipeline) Recall(query, namespace string, topK int, diversity flo } } - // 将召回结果写入本地缓存(goroutine 依赖此 cache 做 increment) - // cache key = id,与 Update() 中 lookup key 一致 - for _, r := range results { - if r.ID == "" { - continue - } - _local.mu.Lock() - if existing, ok := _local.memories[r.ID]; ok { - existing.RecallCount++ - } else { - _local.memories[r.ID] = &models.MemoryRecord{ - ID: r.ID, - RecallCount: 1, - } - } - _local.mu.Unlock() - } - - // Async increment recall_count + update last_recalled_at + importance - go p.incrementRecallCount(results) + // 召回元数据写(2026-09-12 优化B:累积延迟更新) + // 旧实现:对每条结果同步 lancedb.Update → LanceDB MVCC 每写一行一个版本 + // (实测 15 版本/分钟 → _versions 膨胀到 17.7G,根因所在) + // 现在:只把命中 ID 放进 RecallWriteBuffer,后台按窗口批量单事务提交(每批 1 个版本) + p.recordRecallWrites(results) // 共访追踪(即使无 prefetch pusher 也运行,用于持久化) go func() { @@ -410,6 +395,28 @@ func cosineSimilarity(a, b []float32) float64 { return sum } +// recordRecallWrites 记录召回元数据写。 +// 主路径:累积延迟更新缓冲(RecallWriteBufferInstance)——读路径零同步写、批量单事务落盘。 +// 兜底:缓冲未初始化(如单测/独立调用)时退回旧的逐条异步 Update。 +func (p *RecallPipeline) recordRecallWrites(results []models.RecallResult) { + ids := make([]string, 0, len(results)) + for _, r := range results { + if r.ID != "" { + ids = append(ids, r.ID) + } + } + if len(ids) == 0 { + return + } + if RecallWriteBufferInstance != nil { + RecallWriteBufferInstance.Record(ids) + return + } + go p.incrementRecallCount(results) +} + +// incrementRecallCount 旧版逐条同步写(每条结果 = 1 个 LanceDB 版本)。 +// 已不作为主路径,仅保留为缓冲不可用时的兜底。 func (p *RecallPipeline) incrementRecallCount(results []models.RecallResult) { now := time.Now().Format(time.RFC3339) for _, r := range results { diff --git a/go/internal/storage/recall_write_buffer.go b/go/internal/storage/recall_write_buffer.go new file mode 100644 index 0000000..da57303 --- /dev/null +++ b/go/internal/storage/recall_write_buffer.go @@ -0,0 +1,279 @@ +package storage + +import ( + "log" + "os" + "sort" + "strconv" + "sync" + "sync/atomic" + "time" +) + +// ─── 召回元数据「累积延迟更新」缓冲(2026-09-12 优化B)───────────────────────── +// +// 问题(根因):recall 读路径原先对**每条结果**同步调用 lancedb.Update +// (recall_count / last_recalled_at / freshness)。LanceDB 是 MVCC 存储—— +// **每次 update 提交产生一个版本**。实测生产:15 版本/分钟、_versions 目录 +// 膨胀到 17.72G(真实数据仅 22M,75,327 个 .manifest,放大 800 倍)。 +// +// 方案: +// +// Record() 读路径只把增量写进内存缓冲;**同一记忆在一个窗口内多次命中合并成 1 条 delta** +// Flush() 后台按窗口(默认 300s)+ 阈值(默认 256 条不同记忆)批量提交一次; +// Rust 侧 lancedb_update_batch 用**一次** update 提交(CASE WHEN 表达式) +// → 每批只产生 1 个 LanceDB 版本(旧实现:每行 1 个版本) +// +// 语义取舍(明确记录,便于日后审计): +// - last_recalled_at / freshness 最多延迟一个窗口(分钟级)。遗忘/衰减判定以「天」为单位,无影响。 +// - 进程崩溃会丢最后一个窗口的增量(召回统计,不是记忆数据本身),可接受。 +// - recall_count 由 Rust 侧「读库现值 + delta」计算(不是读 Go 缓存),顺带修掉旧实现里 +// 「本地缓存 +1 后再 +1」导致的计数漂移。 +type RecallWriteItem struct { + ID string `json:"id"` + Delta int `json:"delta"` +} + +// BatchRecallUpdater 单事务批量更新接口。 +// 由 Rust IPC 后端(RustLanceDBClient)实现;其他后端不支持时自动退化为逐条 Update。 +type BatchRecallUpdater interface { + UpdateRecallBatch(table string, items []RecallWriteItem, lastRecalledAt string) (int64, error) +} + +const ( + defaultRecallFlushInterval = 300 * time.Second + defaultRecallFlushMaxItems = 256 +) + +type pendingRecall struct { + delta int + lastAt time.Time +} + +// RecallWriteBuffer 累积召回元数据写,延迟批量落盘。 +type RecallWriteBuffer struct { + mu sync.Mutex + pending map[string]*pendingRecall + ldb LanceDB + batch BatchRecallUpdater + interval time.Duration + maxItems int + + stopCh chan struct{} + doneCh chan struct{} + stopped bool + startOnce sync.Once + + flushCount int64 + flushItems int64 + flushRows int64 + flushErrors int64 + fallbackRows int64 +} + +// RecallWriteBufferInstance 进程级单例(与 SearchCacheInstance / CoOccurTrackerInstance 同模式) +var RecallWriteBufferInstance *RecallWriteBuffer + +// NewRecallWriteBuffer 创建缓冲并启动后台 flush 协程;interval<=0 表示只按阈值触发。 +func NewRecallWriteBuffer(ldb LanceDB, interval time.Duration, maxItems int) *RecallWriteBuffer { + if maxItems <= 0 { + maxItems = defaultRecallFlushMaxItems + } + b := &RecallWriteBuffer{ + pending: make(map[string]*pendingRecall), + ldb: ldb, + interval: interval, + maxItems: maxItems, + stopCh: make(chan struct{}), + doneCh: make(chan struct{}), + } + // Rust IPC 后端支持单事务批量更新(CASE WHEN 一次提交) + if bu, ok := ldb.(BatchRecallUpdater); ok { + b.batch = bu + } + go b.loop() + return b +} + +// InitRecallWriteBuffer 初始化进程级单例(server 启动时调用一次)。 +// 窗口可用环境变量 RECALL_WRITE_FLUSH_SECONDS 覆盖(运维/验证用)。 +func InitRecallWriteBuffer(ldb LanceDB) *RecallWriteBuffer { + if RecallWriteBufferInstance != nil { + return RecallWriteBufferInstance + } + interval := defaultRecallFlushInterval + if v := os.Getenv("RECALL_WRITE_FLUSH_SECONDS"); v != "" { + if n, err := strconv.Atoi(v); err == nil && n >= 0 { + interval = time.Duration(n) * time.Second + } + } + b := NewRecallWriteBuffer(ldb, interval, defaultRecallFlushMaxItems) + RecallWriteBufferInstance = b + log.Printf("[recall-buffer] 已启用累积延迟更新: flush 窗口=%s, 阈值=%d 条, 单事务批量=%v", + interval, b.maxItems, b.batch != nil) + return b +} + +// StopRecallWriteBuffer 停止单例并落盘最后一个窗口(进程优雅退出时调用)。 +func StopRecallWriteBuffer() { + if b := RecallWriteBufferInstance; b != nil { + b.Stop() + } +} + +func (b *RecallWriteBuffer) loop() { + defer close(b.doneCh) + if b.interval <= 0 { + <-b.stopCh + return + } + t := time.NewTicker(b.interval) + defer t.Stop() + for { + select { + case <-b.stopCh: + return + case <-t.C: + b.Flush() + } + } +} + +// Record 累积一次召回命中的记忆 ID(同窗口内同 ID 合并 delta)。 +// 返回当前待写条数;达到阈值时异步触发一次 flush(不阻塞读路径)。 +func (b *RecallWriteBuffer) Record(ids []string) int { + now := time.Now() + b.mu.Lock() + for _, id := range ids { + if id == "" { + continue + } + p, ok := b.pending[id] + if !ok { + p = &pendingRecall{} + b.pending[id] = p + } + p.delta++ + p.lastAt = now + } + n := len(b.pending) + b.mu.Unlock() + + if n >= b.maxItems { + go b.Flush() + } + return n +} + +// Flush 取出当前窗口的全部增量并批量提交(每批最多 1 个 LanceDB 版本)。 +func (b *RecallWriteBuffer) Flush() (int, error) { + b.mu.Lock() + if len(b.pending) == 0 { + b.mu.Unlock() + return 0, nil + } + batch := make([]RecallWriteItem, 0, len(b.pending)) + var lastAt time.Time + for id, p := range b.pending { + if p.delta <= 0 { + continue + } + batch = append(batch, RecallWriteItem{ID: id, Delta: p.delta}) + if p.lastAt.After(lastAt) { + lastAt = p.lastAt + } + } + b.pending = make(map[string]*pendingRecall) + b.mu.Unlock() + + if len(batch) == 0 { + return 0, nil + } + sort.Slice(batch, func(i, j int) bool { return batch[i].ID < batch[j].ID }) + ts := lastAt.Format(time.RFC3339) + + var rows int64 + var err error + usedFallback := false + if b.batch != nil { + rows, err = b.batch.UpdateRecallBatch("memories", batch, ts) + } + if b.batch == nil || err != nil { + usedFallback = true + if err != nil { + log.Printf("[recall-buffer] 单事务批量提交失败(%v) → 退化逐条 Update(数据不丢,版本数不优化)", err) + atomic.AddInt64(&b.flushErrors, 1) + } + rows = 0 + for _, it := range batch { + if uerr := b.ldb.Update("memories", it.ID, map[string]any{ + "recall_count": map[string]string{"$inc": "1"}, + "last_recalled_at": ts, + "freshness": "verified", + }); uerr != nil { + log.Printf("[recall-buffer] 逐条 Update 失败 id=%s: %v", it.ID, uerr) + continue + } + rows++ + } + atomic.AddInt64(&b.fallbackRows, rows) + } + + // 落盘成功后同步进程内缓存(缓存只是热数据,权威值在 LanceDB) + if rows > 0 { + _local.mu.Lock() + for _, it := range batch { + if m, ok := _local.memories[it.ID]; ok { + m.RecallCount += it.Delta + if !lastAt.IsZero() { + m.LastRecalledAt = lastAt + } + } + } + _local.mu.Unlock() + } + + atomic.AddInt64(&b.flushCount, 1) + atomic.AddInt64(&b.flushItems, int64(len(batch))) + atomic.AddInt64(&b.flushRows, rows) + log.Printf("[recall-buffer] flush: 合并 %d 条记忆 → 落盘 %d 行, 单事务=%v, fallback=%v", + len(batch), rows, !usedFallback, usedFallback) + return int(rows), nil +} + +// Stop 停止后台 flush 协程,并把最后一个窗口落盘(进程优雅退出时调用)。 +func (b *RecallWriteBuffer) Stop() { + b.mu.Lock() + if b.stopped { + b.mu.Unlock() + return + } + b.stopped = true + b.mu.Unlock() + + close(b.stopCh) + select { + case <-b.doneCh: + case <-time.After(10 * time.Second): + log.Printf("[recall-buffer] 停止超时,直接落盘剩余窗口") + } + b.Flush() +} + +// Stats 观测用(flush 频率 / 合并率 / 版本写入行数)。 +func (b *RecallWriteBuffer) Stats() map[string]interface{} { + b.mu.Lock() + pendingN := len(b.pending) + b.mu.Unlock() + return map[string]interface{}{ + "pending": pendingN, + "flush_count": atomic.LoadInt64(&b.flushCount), + "flush_items": atomic.LoadInt64(&b.flushItems), + "flush_rows": atomic.LoadInt64(&b.flushRows), + "flush_errors": atomic.LoadInt64(&b.flushErrors), + "fallback_rows": atomic.LoadInt64(&b.fallbackRows), + "window": b.interval.String(), + "max_items": b.maxItems, + "single_tx": b.batch != nil, + } +} diff --git a/go/internal/storage/recall_write_buffer_test.go b/go/internal/storage/recall_write_buffer_test.go new file mode 100644 index 0000000..16f554b --- /dev/null +++ b/go/internal/storage/recall_write_buffer_test.go @@ -0,0 +1,221 @@ +package storage + +import ( + "fmt" + "sync" + "testing" + "time" + + "github.com/xiaoxue/memoryweave/internal/models" +) + +// fakeLDB 记录写入调用次数:验证「读路径零同步写 + 窗口内合并 + 单事务批量」。 +type fakeLDB struct { + mu sync.Mutex + batchCalls int + batchItems []RecallWriteItem + batchRowTotal int64 + updateCalls int + updateIDs []string + forceBatchErr bool +} + +func (f *fakeLDB) InsertEpisode(agentID, namespace, content, category string) (string, error) { + return "", nil +} +func (f *fakeLDB) GetTopByQuality(agentID string, limit int) ([]models.MemoryRecord, error) { + return nil, nil +} +func (f *fakeLDB) Stats() (map[string]interface{}, error) { return nil, nil } +func (f *fakeLDB) SoftDelete(id, reason string) error { return nil } +func (f *fakeLDB) GetVersionHistory(id string) ([]map[string]interface{}, error) { + return nil, nil +} +func (f *fakeLDB) GetCandidatesForForgetting() ([]map[string]interface{}, error) { + return nil, nil +} +func (f *fakeLDB) GetSkillCandidates(minRecalls, limit int) ([]models.MemoryRecord, error) { + return nil, nil +} +func (f *fakeLDB) Backup(path string) error { return nil } +func (f *fakeLDB) GetAuditLog(limit int) ([]map[string]interface{}, error) { return nil, nil } +func (f *fakeLDB) IncrementUseful(id string) {} +func (f *fakeLDB) IncrementNotUseful(id string) {} +func (f *fakeLDB) UpdateMemoryContent(id, newContent, source string) error { + return nil +} +func (f *fakeLDB) InsertMemory(m models.MemoryRecord) error { return nil } +func (f *fakeLDB) Search(table string, vector []float32, topK int, namespaceFilter string) ([]models.MemoryRecord, error) { + return nil, nil +} +func (f *fakeLDB) Insert(table string, record any) error { return nil } +func (f *fakeLDB) Update(table, id string, fields map[string]any) error { + f.mu.Lock() + f.updateCalls++ + f.updateIDs = append(f.updateIDs, id) + f.mu.Unlock() + return nil +} + +// BatchRecallUpdater 实现(单事务批量) +func (f *fakeLDB) UpdateRecallBatch(table string, items []RecallWriteItem, lastRecalledAt string) (int64, error) { + if f.forceBatchErr { + return 0, fmt.Errorf("forced batch failure") + } + f.mu.Lock() + f.batchCalls++ + f.batchItems = append(f.batchItems, items...) + f.batchRowTotal += int64(len(items)) + f.mu.Unlock() + return int64(len(items)), nil +} + +// AC-2:同一记忆在一个窗口内多次召回 → 只写 1 次(delta 合并),且是单事务批量调用 +func TestRecallWriteBufferCoalescesWithinWindow(t *testing.T) { + fake := &fakeLDB{} + b := NewRecallWriteBuffer(fake, 0, 1000) // 只手动 flush,由测试控制窗口 + defer b.Stop() + + // 模拟 5 次 recall,每次都命中同 3 条记忆(真实场景:prefetch 高频重复命中) + for i := 0; i < 5; i++ { + b.Record([]string{"mem_a", "mem_b", "mem_c", ""}) + } + + // 读路径零同步写 + fake.mu.Lock() + if fake.batchCalls != 0 || fake.updateCalls != 0 { + t.Fatalf("读路径不应产生任何写:batch=%d update=%d", fake.batchCalls, fake.updateCalls) + } + fake.mu.Unlock() + + rows, err := b.Flush() + if err != nil { + t.Fatalf("flush error: %v", err) + } + if rows != 3 { + t.Fatalf("flush rows = %d, want 3(3 条不同记忆)", rows) + } + + fake.mu.Lock() + defer fake.mu.Unlock() + if fake.batchCalls != 1 { + t.Fatalf("batch calls = %d, want 1(单事务批量)", fake.batchCalls) + } + if fake.updateCalls != 0 { + t.Fatalf("不应走逐条 Update 兜底(update calls=%d)", fake.updateCalls) + } + if len(fake.batchItems) != 3 { + t.Fatalf("batch items = %d, want 3", len(fake.batchItems)) + } + for _, it := range fake.batchItems { + if it.Delta != 5 { + t.Fatalf("id=%s delta = %d, want 5(窗口内 5 次命中合并)", it.ID, it.Delta) + } + } +} + +// AC-1:flush 后缓冲清空 —— 空窗口不产生任何写(版本 0 增长) +func TestRecallWriteBufferEmptyFlushNoWrite(t *testing.T) { + fake := &fakeLDB{} + b := NewRecallWriteBuffer(fake, 0, 1000) + defer b.Stop() + + if rows, err := b.Flush(); err != nil || rows != 0 { + t.Fatalf("空窗口 flush = (%d,%v), want (0,nil)", rows, err) + } + fake.mu.Lock() + defer fake.mu.Unlock() + if fake.batchCalls != 0 || fake.updateCalls != 0 { + t.Fatalf("空窗口不应产生写:batch=%d update=%d", fake.batchCalls, fake.updateCalls) + } +} + +// 阈值触发:待写条数达到 maxItems 时自动 flush(不阻塞 Record 调用) +func TestRecallWriteBufferThresholdFlush(t *testing.T) { + fake := &fakeLDB{} + b := NewRecallWriteBuffer(fake, 0, 3) + defer b.Stop() + + ids := []string{"m1", "m2", "m3", "m4"} + b.Record(ids) + + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + fake.mu.Lock() + done := fake.batchCalls > 0 + fake.mu.Unlock() + if done { + break + } + time.Sleep(20 * time.Millisecond) + } + fake.mu.Lock() + defer fake.mu.Unlock() + if fake.batchCalls == 0 { + t.Fatalf("达到阈值(%d)后应自动 flush,但 batchCalls=0", b.maxItems) + } +} + +// 窗口定时触发:interval 到点自动落盘 +func TestRecallWriteBufferTickerFlush(t *testing.T) { + fake := &fakeLDB{} + b := NewRecallWriteBuffer(fake, 100*time.Millisecond, 1000) + defer b.Stop() + + b.Record([]string{"mem_tick"}) + + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + b.mu.Lock() + idle := len(b.pending) == 0 + b.mu.Unlock() + if idle { + break + } + time.Sleep(20 * time.Millisecond) + } + fake.mu.Lock() + defer fake.mu.Unlock() + if fake.batchCalls != 1 { + t.Fatalf("窗口到点应 flush 1 次,实际 batchCalls=%d", fake.batchCalls) + } +} + +// 兜底:批量提交失败时退化为逐条 Update(数据不丢) +func TestRecallWriteBufferFallbackToPerItem(t *testing.T) { + fake := &fakeLDB{forceBatchErr: true} + b := NewRecallWriteBuffer(fake, 0, 1000) + defer b.Stop() + + b.Record([]string{"mem_x", "mem_y"}) + rows, err := b.Flush() + if err != nil { + t.Fatalf("flush 不应因批量失败而报错: %v", err) + } + if rows != 2 { + t.Fatalf("兜底 rows = %d, want 2", rows) + } + fake.mu.Lock() + defer fake.mu.Unlock() + if fake.updateCalls != 2 { + t.Fatalf("兜底逐条 Update 调用数 = %d, want 2", fake.updateCalls) + } + st := b.Stats() + if st["flush_errors"].(int64) != 1 { + t.Fatalf("flush_errors = %v, want 1", st["flush_errors"]) + } +} + +// Stop 时落盘最后一个窗口(优雅退出不丢统计) +func TestRecallWriteBufferStopFlushes(t *testing.T) { + fake := &fakeLDB{} + b := NewRecallWriteBuffer(fake, 0, 1000) + b.Record([]string{"mem_stop"}) + b.Stop() + + fake.mu.Lock() + defer fake.mu.Unlock() + if fake.batchCalls != 1 { + t.Fatalf("Stop 应落盘最后一个窗口,batchCalls=%d", fake.batchCalls) + } +} diff --git a/rust/src/lancedb_ops.rs b/rust/src/lancedb_ops.rs index 8bcbef0..c73a3bb 100644 --- a/rust/src/lancedb_ops.rs +++ b/rust/src/lancedb_ops.rs @@ -350,6 +350,87 @@ impl LanceDBOps { Ok(updated) } + /// 累积召回批量更新(2026-09-12 优化B)。 + /// + /// 背景:recall 读路径旧实现对每条结果调一次 update,而 LanceDB 是 MVCC—— + /// **每次 update 提交 = 一个版本**,实测 15 版本/分钟 → _versions 膨胀到 17.7G。 + /// + /// 做法(实测约束:lance 的 update 只支持基础 SQL 表达式,`CASE WHEN` 会被 + /// planner 拒绝 —— 见 2026-09-12 实测 `Expression 'CASE WHEN ...' is not supported SQL in lance`): + /// 把整批**按 delta 分组**,每组用一次 `recall_count + delta` 表达式 + `id IN (...)` 谓词提交。 + /// 真实召回流里 delta 绝大多数是 1(同一窗口内多次命中才 >1),所以 + /// N 条记忆的批量 ≈ **1 个版本**(旧实现:N 个版本)。 + /// + /// items_json: [{"id":"mem_xxx","delta":2}],ts: RFC3339(last_recalled_at) + pub fn update_recall_batch(&self, items_json: &str, ts: &str) -> Result> { + #[derive(serde::Deserialize)] + struct Item { + id: String, + delta: i64, + } + let items: Vec = serde_json::from_str(items_json)?; + if items.is_empty() { + return Ok(0); + } + let db = rt().block_on(lancedb::connect(self.data_dir.to_str().unwrap()).execute())?; + let tbl = rt().block_on(db.open_table("memories").execute())?; + + let esc = |s: &str| s.replace('\'', "''"); + let ts_lit = format!("'{}'", esc(ts)); + + // 按 delta 分组(BTreeMap 保证顺序稳定,便于日志/排障) + let mut groups: std::collections::BTreeMap> = std::collections::BTreeMap::new(); + for it in &items { + groups.entry(it.delta).or_default().push(it.id.clone()); + } + + let mut total: u64 = 0; + for (delta, ids) in &groups { + let predicate = format!( + "id IN ({})", + ids.iter().map(|id| format!("'{}'", esc(id))).collect::>().join(",") + ); + let op = tbl + .update() + .only_if(&predicate) + .column("recall_count", &format!("recall_count + {}", delta)) + .column("last_recalled_at", &ts_lit) + .column("freshness", "'verified'"); + match rt().block_on(op.execute()) { + Ok(n) => { + eprintln!( + "[lancedb] BATCH UPDATE recall DONE: delta={} ids={} rows={} (单事务)", + delta, + ids.len(), + n + ); + total += n; + } + Err(e) => { + // 兜底:谓词/表达式被拒时逐条更新(版本数退化为 N,但数据不丢、计数正确) + eprintln!( + "[lancedb] BATCH UPDATE(delta={}) failed ({}), fallback per-item ({} ids)", + delta, + e, + ids.len() + ); + for id in ids { + let fields = format!( + r#"[{{"column":"recall_count","value":"recall_count + {}"}},{{"column":"last_recalled_at","value":"{}"}},{{"column":"freshness","value":"'verified'"}}]"#, + delta, ts_lit + ); + match self.update("memories", id, &fields) { + Ok(n) => total += n, + Err(e2) => eprintln!("[lancedb] fallback update {} failed: {}", id, e2), + } + } + } + } + } + Ok(total) + } + + /// 全量扫描 memories 表(用于深整),上限 10000 条 // P2 2026-09-06 安全版遗忘候选全表扫描(替代 febc2c9 风暴版): // 1. 强制 limit(调用方传, main.rs 钳制硬上限 5000) diff --git a/rust/src/main.rs b/rust/src/main.rs index a1af04f..6f24db5 100644 --- a/rust/src/main.rs +++ b/rust/src/main.rs @@ -300,6 +300,21 @@ fn handle_client(mut stream: UnixStream, args: &Args, lancedb: LanceDBOps, state } } } + "lancedb_update_batch" => { + // 累积召回批量更新(2026-09-12 优化B):一次 update 提交 = 1 个 LanceDB 版本 + let _table = msg["table"].as_str().unwrap_or("memories"); + let items = msg["items"].as_str().unwrap_or("[]"); + let ts = msg["ts"].as_str().unwrap_or(""); + match lancedb.update_recall_batch(items, ts) { + Ok(n) => { + let cnt: usize = serde_json::from_str::>(items) + .map(|v| v.len()) + .unwrap_or(0); + send_ok(&mut stream, &format!(r#"{{"updated":{},"items":{}}}"#, n, cnt)); + } + Err(e) => send_error(&mut stream, "lancedb_update_batch", &e.to_string()), + } + } "lancedb_query" => { let min_recall = msg["min_recall"].as_i64().unwrap_or(5) as i64; let limit = msg["limit"].as_u64().unwrap_or(20) as usize;