perf(recall): 累积延迟更新治理写放大 — 读路径零逐条写 + 单事务批量落盘 (t_8496e8b6)

根因: 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 分组规避)
This commit is contained in:
小怡 2026-09-12 17:57:31 +08:00
parent 38c31eeede
commit 2e04adf054
10 changed files with 836 additions and 25 deletions

View File

@ -11,6 +11,7 @@ import (
"time" "time"
"github.com/xiaoxue/memoryweave/internal/api" "github.com/xiaoxue/memoryweave/internal/api"
"github.com/xiaoxue/memoryweave/internal/storage"
) )
func main() { func main() {
@ -34,6 +35,8 @@ func main() {
<-sigCh <-sigCh
log.Println("[zhiyid] 收到关闭信号,正在退出...") log.Println("[zhiyid] 收到关闭信号,正在退出...")
// 累积延迟更新:把最后一个窗口的召回元数据落盘(防止统计增量丢失)
storage.StopRecallWriteBuffer()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel() defer cancel()
if err := srv.Shutdown(ctx); err != nil { if err := srv.Shutdown(ctx); err != nil {

View File

@ -52,6 +52,11 @@ func NewServer() http.Handler {
// 存储后端 // 存储后端
ldb := initStorageBackend(os.Getenv("STORAGE_BACKEND"), emb) ldb := initStorageBackend(os.Getenv("STORAGE_BACKEND"), emb)
// 召回元数据「累积延迟更新」缓冲2026-09-12 优化B
// 读路径不再逐条写 LanceDBMVCC 每写一行一个版本 → _versions 膨胀根因),
// 改为窗口内合并 + 单事务批量提交(每批 1 个版本)。
storage.InitRecallWriteBuffer(ldb)
// 启动时数据目录一致性检查(防止路径混乱导致读取废弃数据) // 启动时数据目录一致性检查(防止路径混乱导致读取废弃数据)
runStartupChecks(os.Getenv("STORAGE_BACKEND")) runStartupChecks(os.Getenv("STORAGE_BACKEND"))

View File

@ -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_updaterecall 读路径原行为)
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实际 %ddelta 未正确累加)", before+3, after)
}
fmt.Printf("BATCH_IPC_OK id=%s recall_count %d → %d\n", target, before, after)
}

View File

@ -46,10 +46,11 @@ func BenchmarkEmbedder_Batch50(b *testing.B) {
// ─── Recall Pipeline 基准 ───────────────────────────────── // ─── Recall Pipeline 基准 ─────────────────────────────────
func BenchmarkRecallPipeline_10Docs(b *testing.B) { func BenchmarkRecallPipeline_10Docs(b *testing.B) {
emb := &Embedder{endpoint: "http://localhost:8000/v1/embeddings"}
p := NewRecallPipeline( p := NewRecallPipeline(
&Embedder{endpoint: "http://localhost:8000/v1/embeddings"}, emb,
NewMemLanceClient(nil), NewMemLanceClient(nil),
NewReranker("http://localhost:8001/rerank"), NewReranker("http://localhost:8001/rerank", emb),
) )
b.ResetTimer() b.ResetTimer()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
@ -58,10 +59,11 @@ func BenchmarkRecallPipeline_10Docs(b *testing.B) {
} }
func BenchmarkRecallPipeline_50Docs(b *testing.B) { func BenchmarkRecallPipeline_50Docs(b *testing.B) {
emb := &Embedder{endpoint: "http://localhost:8000/v1/embeddings"}
p := NewRecallPipeline( p := NewRecallPipeline(
&Embedder{endpoint: "http://localhost:8000/v1/embeddings"}, emb,
NewMemLanceClient(nil), NewMemLanceClient(nil),
NewReranker("http://localhost:8001/rerank"), NewReranker("http://localhost:8001/rerank", emb),
) )
b.ResetTimer() b.ResetTimer()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {

View File

@ -44,6 +44,8 @@ type ipcReq struct {
Fields string `json:"fields,omitempty"` Fields string `json:"fields,omitempty"`
MinRecall int `json:"min_recall,omitempty"` MinRecall int `json:"min_recall,omitempty"`
QueryLimit int `json:"limit,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 { type ipcResp struct {
@ -450,6 +452,37 @@ func (rc *RustLanceDBClient) Update(table, id string, fields map[string]any) err
return nil 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) { func (rc *RustLanceDBClient) GetTopByQuality(agentID string, limit int) ([]models.MemoryRecord, error) {
_local.mu.RLock() _local.mu.RLock()
defer _local.mu.RUnlock() defer _local.mu.RUnlock()

View File

@ -73,7 +73,7 @@ func (p *RecallPipeline) Recall(query, namespace string, topK int, diversity flo
if len(filtered) == 0 && len(results) > 0 { if len(filtered) == 0 && len(results) > 0 {
filtered = results filtered = results
} }
go p.incrementRecallCount(results) p.recordRecallWrites(filtered)
return filtered, nil return filtered, nil
} }
} }
@ -258,26 +258,11 @@ func (p *RecallPipeline) Recall(query, namespace string, topK int, diversity flo
} }
} }
// 将召回结果写入本地缓存goroutine 依赖此 cache 做 increment // 召回元数据写2026-09-12 优化B累积延迟更新
// cache key = id与 Update() 中 lookup key 一致 // 旧实现:对每条结果同步 lancedb.Update → LanceDB MVCC 每写一行一个版本
for _, r := range results { // (实测 15 版本/分钟 → _versions 膨胀到 17.7G,根因所在)
if r.ID == "" { // 现在:只把命中 ID 放进 RecallWriteBuffer后台按窗口批量单事务提交每批 1 个版本)
continue p.recordRecallWrites(results)
}
_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)
// 共访追踪(即使无 prefetch pusher 也运行,用于持久化) // 共访追踪(即使无 prefetch pusher 也运行,用于持久化)
go func() { go func() {
@ -410,6 +395,28 @@ func cosineSimilarity(a, b []float32) float64 {
return sum 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) { func (p *RecallPipeline) incrementRecallCount(results []models.RecallResult) {
now := time.Now().Format(time.RFC3339) now := time.Now().Format(time.RFC3339)
for _, r := range results { for _, r := range results {

View File

@ -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(真实数据仅 22M75,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,
}
}

View File

@ -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 33 条不同记忆)", 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-1flush 后缓冲清空 —— 空窗口不产生任何写(版本 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)
}
}

View File

@ -350,6 +350,87 @@ impl LanceDBOps {
Ok(updated) 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: RFC3339last_recalled_at
pub fn update_recall_batch(&self, items_json: &str, ts: &str) -> Result<u64, Box<dyn std::error::Error>> {
#[derive(serde::Deserialize)]
struct Item {
id: String,
delta: i64,
}
let items: Vec<Item> = 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<i64, Vec<String>> = 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::<Vec<_>>().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 条 /// 全量扫描 memories 表(用于深整),上限 10000 条
// P2 2026-09-06 安全版遗忘候选全表扫描(替代 febc2c9 风暴版): // P2 2026-09-06 安全版遗忘候选全表扫描(替代 febc2c9 风暴版):
// 1. 强制 limit调用方传, main.rs 钳制硬上限 5000 // 1. 强制 limit调用方传, main.rs 钳制硬上限 5000

View File

@ -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::<Vec<serde_json::Value>>(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" => { "lancedb_query" => {
let min_recall = msg["min_recall"].as_i64().unwrap_or(5) as i64; let min_recall = msg["min_recall"].as_i64().unwrap_or(5) as i64;
let limit = msg["limit"].as_u64().unwrap_or(20) as usize; let limit = msg["limit"].as_u64().unwrap_or(20) as usize;