287 lines
7.5 KiB
Go
287 lines
7.5 KiB
Go
// 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,
|
||
})
|
||
}
|
||
}
|
||
}
|