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:
parent
38c31eeede
commit
2e04adf054
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
// 读路径不再逐条写 LanceDB(MVCC 每写一行一个版本 → _versions 膨胀根因),
|
||||||
|
// 改为窗口内合并 + 单事务批量提交(每批 1 个版本)。
|
||||||
|
storage.InitRecallWriteBuffer(ldb)
|
||||||
|
|
||||||
// 启动时数据目录一致性检查(防止路径混乱导致读取废弃数据)
|
// 启动时数据目录一致性检查(防止路径混乱导致读取废弃数据)
|
||||||
runStartupChecks(os.Getenv("STORAGE_BACKEND"))
|
runStartupChecks(os.Getenv("STORAGE_BACKEND"))
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
@ -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++ {
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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: RFC3339(last_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)
|
||||||
|
|
|
||||||
|
|
@ -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;
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue