memoryweave/go/internal/storage/recall.go

287 lines
7.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// Package storage 提供记忆召回完成管线:编码 → 向量搜索 → 重排 → MMR → 图谱扩展 → 预取推送
package storage
import (
"encoding/json"
"fmt"
"math"
"time"
"github.com/xiaoxue/memoryweave/internal/models"
)
// MissRecorder 缺口记录接口(由 routes 层注入 gapDetector.RecordMiss
type MissRecorder func(query, namespace string)
// PrefetchPusher 预取推送接口(由 routes 层注入)
type PrefetchPusher interface {
PushPrefetch(agentID string, memories []models.RecallResult)
}
// GraphExpander 图谱扩展接口(由 governance 层注入)
type GraphExpander interface {
ExpandFromResults(results []models.RecallResult, namespace string, maxHops int) []models.RecallResult
}
type RecallPipeline struct {
embedder *Embedder
lancedb LanceDB
reranker *Reranker
graph GraphExpander
prefetch PrefetchPusher
recordMiss MissRecorder
}
func NewRecallPipeline(embedder *Embedder, lancedb LanceDB, reranker *Reranker) *RecallPipeline {
return &RecallPipeline{
embedder: embedder,
lancedb: lancedb,
reranker: reranker,
}
}
// SetGraphExpander 注入图谱扩展器
func (p *RecallPipeline) SetGraphExpander(g GraphExpander) {
p.graph = g
}
// SetPrefetchPusher 注入预取推送器
func (p *RecallPipeline) SetPrefetchPusher(pf PrefetchPusher) {
p.prefetch = pf
}
// SetMissRecorder 注入缺口记录器gapDetector.RecordMiss
func (p *RecallPipeline) SetMissRecorder(mr MissRecorder) {
p.recordMiss = mr
}
func (p *RecallPipeline) Recall(query, namespace string, topK int, diversity float64) ([]models.RecallResult, error) {
if topK <= 0 {
topK = 10
}
// Step 0: 搜索缓存§2.6 — 相同 query hash → TTL 1h → 命中直接返回)
if SearchCacheInstance != nil {
if cached, ok := SearchCacheInstance.Get(query, namespace); ok {
var results []models.RecallResult
if err := json.Unmarshal(cached, &results); err == nil {
return results, nil
}
}
}
// Step 1: Encode query
queryVec, err := p.embedder.EncodeSingle(query)
if err != nil {
return nil, fmt.Errorf("recall encode: %w", err)
}
// Step 2: Coarse search — own namespace + shared§1.3 召回规则)
// 默认搜自己的 namespace + shared显式 include_namespaces 时才跨域
var candidates []models.MemoryRecord
namespaces := []string{namespace}
if namespace != "shared" {
namespaces = append(namespaces, "shared")
}
for _, ns := range namespaces {
nsResults, err := p.lancedb.Search("memories", queryVec, 50, ns)
if err != nil {
continue
}
candidates = append(candidates, nsResults...)
}
// 去重(按 ID
seen := make(map[string]bool)
unique := make([]models.MemoryRecord, 0, len(candidates))
for _, c := range candidates {
if !seen[c.ID] {
seen[c.ID] = true
unique = append(unique, c)
}
}
candidates = unique
// 截断到 top 50
if len(candidates) > 50 {
candidates = candidates[:50]
}
if len(candidates) == 0 {
// 自动记录召回缺口gap_scan 触发器会在 30min 冷却后分类)
if p.recordMiss != nil {
p.recordMiss(query, namespace)
}
return []models.RecallResult{}, nil
}
// Step 3: Rerank to top_k (graceful fallback)
docs := make([]string, len(candidates))
for i, c := range candidates {
docs[i] = c.Content
}
reranked, err := p.reranker.Rerank(query, docs, topK*2)
if err != nil {
// 重排不可用时降级:直接使用向量搜索的原始排序
reranked = make([]models.RerankResult, len(candidates))
for i := range candidates {
reranked[i] = models.RerankResult{Index: i, Score: 0.5, Text: candidates[i].Content}
}
}
// Step 4: MMR diversity
finalIndices := MMRSelect(reranked, candidates, topK*2, diversity)
results := make([]models.RecallResult, 0, topK)
for i, idx := range finalIndices {
if i >= topK {
break
}
if idx < len(candidates) {
c := candidates[idx]
score := 0.0
for _, r := range reranked {
if r.Index == idx {
score = r.Score
break
}
}
results = append(results, models.RecallResult{
ID: c.ID,
Content: c.Content,
Category: c.Category,
Score: score*0.7 + c.QualityScore*0.3,
Timestamp: c.CreatedAt.Format(time.RFC3339),
})
}
}
// Step 5: 图谱多跳扩展(结果 < 5 条时双向 BFS 扩展 1 跳)
if len(results) < 5 && p.graph != nil {
expanded := p.graph.ExpandFromResults(results, namespace, 1)
for _, e := range expanded {
already := false
for _, r := range results {
if r.ID == e.ID {
already = true
break
}
}
if !already {
results = append(results, e)
}
}
}
// Step 6: 去重(相同 content 只保留最高分)
seenContent := make(map[string]bool)
deduped := make([]models.RecallResult, 0, len(results))
for _, r := range results {
if !seenContent[r.Content] {
seenContent[r.Content] = true
deduped = append(deduped, r)
}
}
results = deduped
// Step 6.5: 写搜索缓存§2.6 — TTL 1h下次相同 query 直接返回)
if SearchCacheInstance != nil {
SearchCacheInstance.Set(query, namespace, results)
}
// Step 7: 预取推送CO_OCCURS 权重 > 0.6 的配套记忆 → WebSocket
if p.prefetch != nil {
// 使用 CO_OCCURS 追踪器收集真实预取候选项
var prefetchIDs []string
for _, r := range results {
if r.ID != "" {
prefetchIDs = append(prefetchIDs, r.ID)
}
}
if len(prefetchIDs) > 0 {
// 记录共访统计到 CO_OCCURS tracker
if CoOccurTrackerInstance != nil {
CoOccurTrackerInstance.RecordRecall(prefetchIDs, namespace)
candidates := CoOccurTrackerInstance.CollectCandidates(prefetchIDs)
if len(candidates) > 0 {
prefetchItems := make([]models.RecallResult, 0, len(candidates))
for _, cid := range candidates {
prefetchItems = append(prefetchItems, models.RecallResult{ID: cid})
}
p.prefetch.PushPrefetch("", prefetchItems)
}
}
}
}
// Async increment recall_count + update last_recalled_at + importance
go p.incrementRecallCount(results)
return results, nil
}
// MMRSelect 多样性去重 — 开放给 recall/debug 诊断端点
func MMRSelect(reranked []models.RerankResult, candidates []models.MemoryRecord, k int, lambda float64) []int {
if len(reranked) == 0 {
return nil
}
selected := make([]int, 0, k)
cand := make([]int, len(reranked))
for i := range reranked {
cand[i] = reranked[i].Index
}
for len(selected) < k && len(cand) > 0 {
bestIdx := 0
bestScore := -math.MaxFloat64
for i, ci := range cand {
rel := 0.0
for _, r := range reranked {
if r.Index == ci {
rel = r.Score
break
}
}
maxSim := 0.0
for _, si := range selected {
if si < len(candidates) && ci < len(candidates) {
sim := cosineSimilarity(candidates[ci].Vector, candidates[si].Vector)
if sim > maxSim {
maxSim = sim
}
}
}
mmr := (1-lambda)*rel - lambda*maxSim
if mmr > bestScore {
bestScore = mmr
bestIdx = i
}
}
selected = append(selected, cand[bestIdx])
cand = append(cand[:bestIdx], cand[bestIdx+1:]...)
}
return selected
}
func cosineSimilarity(a, b []float32) float64 {
var sum float64
n := len(a)
if len(b) < n {
n = len(b)
}
for i := 0; i < n; i++ {
sum += float64(a[i]) * float64(b[i])
}
return sum
}
func (p *RecallPipeline) incrementRecallCount(results []models.RecallResult) {
now := time.Now().Format(time.RFC3339)
for _, r := range results {
if r.ID != "" {
_ = p.lancedb.Update("memories", r.ID, map[string]any{
"recall_count": map[string]string{"$inc": "1"},
"last_recalled_at": now,
})
}
}
}