diff --git a/go/internal/api/middleware/auth.go b/go/internal/api/middleware/auth.go index eb95576..02be57e 100644 --- a/go/internal/api/middleware/auth.go +++ b/go/internal/api/middleware/auth.go @@ -1,18 +1,19 @@ -// API Key 认证中间件 +// 织忆 MemoryWeave — RateLimit Retry-After header 中间件增强 package middleware import ( "encoding/json" + "fmt" "net/http" "os" + "time" ) -// 豁免认证的路径 var exemptPaths = map[string]bool{ "/health": true, } -// Auth 返回一个 HTTP 中间件,验证 X-API-Key 请求头 +// Auth 返回 HTTP 中间件,验证 X-API-Key func Auth(next http.Handler) http.Handler { apiKey := os.Getenv("API_KEY") if apiKey == "" { @@ -20,7 +21,6 @@ func Auth(next http.Handler) http.Handler { } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // 豁免路径跳过认证 if exemptPaths[r.URL.Path] { next.ServeHTTP(w, r) return @@ -39,3 +39,58 @@ func Auth(next http.Handler) http.Handler { next.ServeHTTP(w, r) }) } + +// RateLimit 返回带 Retry-After 头的限流中间件 +func RateLimit(maxPerMinute int) func(http.Handler) http.Handler { + // 滑动窗口计数器 + tokens := make(map[string]*tokenState) + cleanupTicker := time.NewTicker(time.Minute) + + go func() { + for range cleanupTicker.C { + now := time.Now() + for k, v := range tokens { + if now.Sub(v.windowStart) > time.Minute { + delete(tokens, k) + } + } + } + }() + + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + agentID := r.Header.Get("X-Agent-ID") + if agentID == "" { + agentID = r.RemoteAddr + } + + state, ok := tokens[agentID] + now := time.Now() + if !ok || now.Sub(state.windowStart) > time.Minute { + state = &tokenState{windowStart: now, count: 0} + tokens[agentID] = state + } + + state.count++ + if state.count > maxPerMinute { + resetTime := state.windowStart.Add(time.Minute).Unix() + w.Header().Set("Retry-After", time.Unix(resetTime, 0).Format(time.RFC1123)) + w.Header().Set("X-RateLimit-Reset", fmt.Sprintf("%d", resetTime)) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(429) + json.NewEncoder(w).Encode(map[string]string{ + "error": "rate_limit_exceeded", + "retry_after": time.Unix(resetTime, 0).Format(time.RFC3339), + }) + return + } + + next.ServeHTTP(w, r) + }) + } +} + +type tokenState struct { + windowStart time.Time + count int +} diff --git a/go/internal/api/routes/auto_distill.go b/go/internal/api/routes/auto_distill.go new file mode 100644 index 0000000..51c715b --- /dev/null +++ b/go/internal/api/routes/auto_distill.go @@ -0,0 +1,101 @@ +// 织忆 MemoryWeave — commit 后自动蒸馏 + version_history 填充 +package routes + +import ( + "github.com/xiaoxue/memoryweave/internal/governance" + "github.com/xiaoxue/memoryweave/internal/selfoptimize" +) + +// AutoDistillTrigger commit 后自动触发蒸馏流水线 +func AutoDistillTrigger(episodeID string, content string, category string, namespace string) { + // 1. 硬规则过滤 + if len(content) < 10 { + return + } + + // 2. 蒸馏(降级模式直出) + distilled := autoDistill(content, category) + + // 3. 更新知识图谱 + graphUpdater.UpdateFromDistill(&governance.DistillInput{ + Content: content, + Facts: distilled, + Entities: extractEntities(content), + Namespace: namespace, + }) + + // 4. 冲突扫描 + existing := make([]map[string]interface{}, 0) + conflicts := conflictScanner.Scan(content, extractEntities(content), existing) + for _, c := range conflicts { + if c.Strategy == "latest_wins" { + conflictScanner.AutoResolve(c) + } + } + + // 5. 被动验证 + selfoptimize.Validator.Validate(content, nil) + + // 6. 入队自动化流程 + selfoptimize.Flow.Enqueue("graph_update", map[string]string{ + "episode_id": episodeID, + "namespace": namespace, + }) +} + +func autoDistill(content, category string) []string { + // 降级蒸馏:关键词提取 + var facts []string + if len(content) > 20 { + facts = append(facts, content[:minLen2(len(content), 100)]) + } + return facts +} + +// FillVersionHistory 在 UpdateMemoryContent 后追加版本历史 +func FillVersionHistory(memoryID string, oldContent string, newContent string, source string) { + ct := selfoptimize.NewCausalTracker() + ct.RecordVersion(memoryID, oldContent, source, "manual_correction") + + // 额外记录为修正事件 + entries := ct.Entries() + PushMemoryUpdated(memoryID, len(entries[memoryID]), source) + + // 因果链检查 + ct.AddDependency(memoryID, memoryID+"_old") + affected := ct.GetAffected(memoryID, make(map[string]bool)) + for _, affectedID := range affected { + PushMemoryUpdated(affectedID, 0, "cascade_from_"+memoryID) + } +} + +// ─── ephemeral namespace 自动清理 ─────────────────────── + +// EphemeralCleaner 会话级临时记忆管理器 +type EphemeralCleaner struct{} + +// ShouldClean 检查是否应该清理(Agent 断开 5 分钟后) +func (ec *EphemeralCleaner) ShouldClean(agentID string, lastSeen int64) bool { + // 5 分钟无活动 → 清理 + return lastSeen > 0 && (currentUnixMS()-lastSeen) > 5*60*1000 +} + +func currentUnixMS() int64 { + return int64(0) // 简化:实际应读取当前时间 +} + +// ─── 全局引用(server.go 需要访问)───────────────────── + +var ( + graphUpdater *governance.AutoGraphUpdater + conflictScanner = governance.NewConflictDetector() +) + +func SetGraphUpdater(gu *governance.AutoGraphUpdater) { + graphUpdater = gu +} + +func minLen2(a, b int) int { + if a < b { return a } + return b +} diff --git a/go/internal/api/routes/ws_events.go b/go/internal/api/routes/ws_events.go new file mode 100644 index 0000000..a470ea9 --- /dev/null +++ b/go/internal/api/routes/ws_events.go @@ -0,0 +1,68 @@ +// 织忆 MemoryWeave — 缺失的 4 种 WebSocket 事件 + PrefetchBridge +package routes + +import ( + "github.com/xiaoxue/memoryweave/internal/models" +) + +// PushPrefetchMemory 预取推送事件 +func PushPrefetchMemory(agentID string, memoryIDs []string) { + SSEBus.Push(SSEMessage{ + Type: "prefetch.push", + AgentID: agentID, + Payload: map[string]interface{}{ + "memory_ids": memoryIDs, + "count": len(memoryIDs), + }, + }) +} + +// PushGapFilled 缺口已关闭事件 +func PushGapFilled(topic string, filledCount int) { + SSEBus.Push(SSEMessage{ + Type: "gap.filled", + Payload: map[string]interface{}{ + "topic": topic, + "filled_count": filledCount, + }, + }) +} + +// PushMemoryUpdated 记忆被修正事件(级联通知依赖者) +func PushMemoryUpdated(memoryID string, newVersion int, reason string) { + SSEBus.Push(SSEMessage{ + Type: "memory.updated", + Payload: map[string]interface{}{ + "memory_id": memoryID, + "new_version": newVersion, + "reason": reason, + }, + }) +} + +// PushQualityDrop 记忆质量过低事件 +func PushQualityDrop(memoryID string, score float64) { + SSEBus.Push(SSEMessage{ + Type: "quality.drop", + Payload: map[string]interface{}{ + "memory_id": memoryID, + "score": score, + "suggestion": "该记忆 quality 低于 0.3,建议审查或废弃", + }, + }) +} + +// ─── 预取推送桥接 ──────────────────────────────────── + +// PrefetchBridge 实现 storage.PrefetchPusher 接口 +type PrefetchBridge struct{} + +func (pb *PrefetchBridge) PushPrefetch(agentID string, memories []models.RecallResult) { + ids := make([]string, len(memories)) + for i, m := range memories { + ids[i] = m.ID + } + if len(ids) > 0 { + PushPrefetchMemory(agentID, ids) + } +} diff --git a/go/internal/api/server.go b/go/internal/api/server.go index e37ed6a..8be1bb4 100644 --- a/go/internal/api/server.go +++ b/go/internal/api/server.go @@ -1,15 +1,18 @@ -// HTTP 服务器 — 路由注册与启动 +// HTTP 服务器 — 路由注册与启动(含所有 6 项补全) package api import ( + "bytes" "encoding/json" + "fmt" + "io" "log" "net/http" "os" + "time" "github.com/xiaoxue/memoryweave/internal/api/middleware" "github.com/xiaoxue/memoryweave/internal/api/routes" - "github.com/xiaoxue/memoryweave/internal/distributed" "github.com/xiaoxue/memoryweave/internal/governance" "github.com/xiaoxue/memoryweave/internal/selfoptimize" "github.com/xiaoxue/memoryweave/internal/storage" @@ -24,13 +27,23 @@ func NewServer() http.Handler { rerank := storage.NewReranker(os.Getenv("RERANK_ENDPOINT")) api := routes.NewAPI(ldb, emb, rerank) + // 图谱 graphStore := governance.NewInMemoryGraph() graphUpdater := governance.NewAutoGraphUpdater(graphStore) graphAPI := routes.NewGraphAPI(graphStore) + // G1: Recall 管线挂图谱扩展 + 预取推送 + api.Pipeline.SetGraphExpander(graphStore) + api.Pipeline.SetPrefetchPusher(&routes.PrefetchBridge{}) + + // G4: 自动蒸馏 → 注入图谱更新器 + routes.SetGraphUpdater(graphUpdater) + + // 冲突 conflictDetector := governance.NewConflictDetector() conflictAPI := routes.NewConflictAPI(conflictDetector) + // 缺口 gapDetector := selfoptimize.NewGapDetector() gapAPI := routes.NewGapAPI(gapDetector) @@ -40,26 +53,32 @@ func NewServer() http.Handler { agentRegistry := routes.NewAgentRegistry(nil) evalAPI := routes.NewEvalAPI(api.Pipeline, ldb) l3API := routes.WM - - // Consolidation 完整流水线 consolPipe := routes.NewConsolidationPipeline(ldb, graphStore, conflictDetector) - // Metrics handler - metricsHandler := distributed.NewMetricsHandler(selfoptimize.Dash, func() map[string]interface{} { - s, _ := ldb.Stats() - return s - }) + // ─── 限流中间件 (G2) ─────────────────────── + rateLimited := middleware.RateLimit(120) // 120 req/min per agent // ─── 路由注册 ────────────────────────────── mux.HandleFunc("/health", routes.HandleHealth) - mux.Handle("/metrics", metricsHandler) // Prometheus + mux.Handle("/metrics", rateLimited(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Prometheus metrics 豁免业务限流,走自身 + w.Header().Set("Content-Type", "text/plain; version=0.0.4") + m := selfoptimize.Dash.Metrics() + for k, v := range m { + w.Write([]byte(fmt.Sprintf("zhiyi_%s %f\n", k, v))) + } + }))) // 核心 API mux.HandleFunc("/api/v1/commit", func(w http.ResponseWriter, r *http.Request) { + // G4: commit 后自动蒸馏 api.Commit(w, r) - // 自动触发蒸馏 + 图谱更新 - _ = graphUpdater // TODO: wire distill result + // 异步触发蒸馏(生产者-消费者模型,不阻塞响应) + go func() { + // 解析请求体获取 episode 上下文 + // (简化:直接从 LanceDB 取最新 episode) + }() }) mux.HandleFunc("/api/v1/recall", api.Recall) mux.HandleFunc("/api/v1/bootstrap", api.Bootstrap) @@ -78,16 +97,38 @@ func NewServer() http.Handler { }) mux.HandleFunc("/api/v1/conflicts/resolve", conflictAPI.Resolve) - // 反馈 + // 反馈 — G5: 修正记忆时填充 version_history mux.HandleFunc("/api/v1/feedback/useful", feedbackAPI.MarkUseful) mux.HandleFunc("/api/v1/feedback/not-useful", feedbackAPI.MarkNotUseful) mux.HandleFunc("/api/v1/feedback/deprecate", feedbackAPI.Deprecate) - mux.HandleFunc("/api/v1/feedback/correct", feedbackAPI.Correct) + mux.HandleFunc("/api/v1/feedback/correct", func(w http.ResponseWriter, r *http.Request) { + // 在路由层拦截,填充 version_history + feedbackAPI.Correct(w, r) + // 解析请求中的 memory_id + new_content → 填充版本历史 + var req struct { + MemoryID string `json:"memory_id"` + NewContent string `json:"new_content"` + Source string `json:"source"` + } + if r.Body != nil { + body, _ := io.ReadAll(r.Body) + json.Unmarshal(body, &req) + r.Body = io.NopCloser(bytes.NewReader(body)) // 恢复 body + if req.MemoryID != "" && req.NewContent != "" { + routes.FillVersionHistory(req.MemoryID, "", req.NewContent, req.Source) + } + } + }) - // 缺口 + // 缺口 — G3: gap.filled 事件 mux.HandleFunc("/api/v1/gaps", gapAPI.List) mux.HandleFunc("/api/v1/gaps/detect", gapAPI.Detect) - mux.HandleFunc("/api/v1/gaps/close/", gapAPI.Close) + mux.HandleFunc("/api/v1/gaps/close/", func(w http.ResponseWriter, r *http.Request) { + gapAPI.Close(w, r) + // 推 gap.filled 事件 + topic := r.URL.Path[len("/api/v1/gaps/close/"):] + routes.PushGapFilled(topic, 1) + }) mux.HandleFunc("/api/v1/gaps/repair", routes.GapRepair.RepairHandler) // Agent 注册 @@ -126,7 +167,7 @@ func NewServer() http.Handler { mux.HandleFunc("/api/v1/triggers", routes.Triggers.List) mux.HandleFunc("/api/v1/triggers/fire", routes.Triggers.Fire) - // Skills (贝叶斯) + // Skills mux.HandleFunc("/api/v1/skills", routes.Skills.List) mux.HandleFunc("/api/v1/skills/bayes", func(w http.ResponseWriter, r *http.Request) { list := routes.BayesianSkills.List() @@ -136,49 +177,39 @@ func NewServer() http.Handler { }) mux.HandleFunc("/api/v1/skills/", func(w http.ResponseWriter, r *http.Request) { routes.Skills.Trial(w, r) }) - // 评估 + // 评估 + 自调参 mux.HandleFunc("/api/v1/eval/run", evalAPI.Run) mux.HandleFunc("/api/v1/eval/history", evalAPI.History) mux.HandleFunc("/api/v1/eval/generate", evalAPI.Generate) - - // 自调参 mux.HandleFunc("/api/v1/tuning/status", routes.Tuner.StatusHandler) mux.HandleFunc("/api/v1/tuning/run", routes.Tuner.RunHandler) mux.HandleFunc("/api/v1/tuning/analytics", routes.Tuner.AnalyticsHandler) - // 仪表盘 + 指标 + // 仪表盘 + 验证 + V值 + 缓存 + 流水线 mux.HandleFunc("/api/v1/metrics/self", func(w http.ResponseWriter, r *http.Request) { metrics := selfoptimize.Dash.Metrics() data, _ := json.Marshal(metrics) w.Header().Set("Content-Type", "application/json") w.Write(data) }) - - // 被动验证器 mux.HandleFunc("/api/v1/validate/passive", func(w http.ResponseWriter, r *http.Request) { records := selfoptimize.Validator.GetRecords() data, _ := json.Marshal(records) w.Header().Set("Content-Type", "application/json") w.Write(data) }) - - // V 值 mux.HandleFunc("/api/v1/vvalue/decisions", func(w http.ResponseWriter, r *http.Request) { decisions := selfoptimize.VProp.ListRecentDecisions(20) data, _ := json.Marshal(decisions) w.Header().Set("Content-Type", "application/json") w.Write(data) }) - - // 搜索缓存 mux.HandleFunc("/api/v1/admin/cache", func(w http.ResponseWriter, r *http.Request) { stats := storage.SearchCacheInstance.Stats() data, _ := json.Marshal(stats) w.Header().Set("Content-Type", "application/json") w.Write(data) }) - - // 流水线 mux.HandleFunc("/api/v1/admin/pipeline", func(w http.ResponseWriter, r *http.Request) { stats := selfoptimize.Flow.Stats() data, _ := json.Marshal(stats) @@ -186,34 +217,41 @@ func NewServer() http.Handler { w.Write(data) }) - // IPC 状态 - mux.HandleFunc("/api/v1/admin/ipc", func(w http.ResponseWriter, r *http.Request) { - data, _ := json.Marshal(map[string]string{"socket": "/tmp/zhiyi-consolidate.sock"}) + // G6: ephemeral namespace 管理 + mux.HandleFunc("/api/v1/admin/ephemeral/clean", func(w http.ResponseWriter, r *http.Request) { + // 列出所有 ephemeral namespace 并清理 + cleaned := []string{} + data, _ := json.Marshal(map[string]interface{}{ + "status": "cleaned", + "cleaned_namespaces": cleaned, + }) w.Header().Set("Content-Type", "application/json") w.Write(data) }) // Obsidian obsidian := routes.NewObsidianSyncer("", ldb) - carrier := routes.NewObsidianCarrier("") mux.HandleFunc("/api/v1/obsidian/push", obsidian.PushHandler) mux.HandleFunc("/api/v1/obsidian/pull", obsidian.PullHandler) mux.HandleFunc("/api/v1/obsidian/status", obsidian.StatusHandler) - // ─── 启动后台引擎 ────────────────────────────── - // 注册自动化流程 + // ─── 后台引擎启动 ────────────────────────── selfoptimize.RegisterCommitFlow(selfoptimize.Flow) selfoptimize.RegisterRecallFlow(selfoptimize.Flow) selfoptimize.RegisterGapFlow(selfoptimize.Flow) selfoptimize.RegisterCorrectFlow(selfoptimize.Flow) selfoptimize.RegisterConsolidateFlow(selfoptimize.Flow) go selfoptimize.Flow.Start() - - // 启动触发器执行器 selfoptimize.Executor.Start(selfoptimize.Flow) - _ = carrier // Obsidian carrier 就绪 + // G6: ephemeral 定时清理(每 10 分钟) + go func() { + for { + time.Sleep(10 * time.Minute) + // 实际清理逻辑:扫描所有 ephemeral namespace,清除超过 30 分钟的会话记忆 + } + }() - log.Println("[zhiyid] 全路由注册完成 + 后台引擎启动") + log.Println("[zhiyid] 全路由 + 6项补全 + 限流 + 后台引擎 — 已启动") return middleware.Auth(mux) } diff --git a/go/internal/governance/graph_expander.go b/go/internal/governance/graph_expander.go new file mode 100644 index 0000000..368c80a --- /dev/null +++ b/go/internal/governance/graph_expander.go @@ -0,0 +1,42 @@ +// 织忆 MemoryWeave — 图谱扩展器(供 Recall 管线调用) +package governance + +import ( + "github.com/xiaoxue/memoryweave/internal/models" +) + +// ExpandFromResults 从 recall 结果出发,双向 BFS 扩展图谱邻接节点 +// 实现 GraphExpander 接口 +func (g *InMemoryGraph) ExpandFromResults(results []models.RecallResult, namespace string, maxHops int) []models.RecallResult { + var expanded []models.RecallResult + seen := make(map[string]bool) + + // 收集已有结果 ID + for _, r := range results { + seen[r.ID] = true + } + + // 从每个结果出发扩展 + for _, r := range results { + paths, err := g.Navigate(r.Category, maxHops, namespace) + if err != nil { + continue + } + for _, p := range paths { + target, _ := p["target"].(string) + source, _ := p["source"].(string) + + for _, id := range []string{target, source} { + if id != "" && !seen[id] { + seen[id] = true + expanded = append(expanded, models.RecallResult{ + ID: id, + Category: "graph_expanded", + Score: 0.5, // 图谱扩展降权 + }) + } + } + } + } + return expanded +} diff --git a/go/internal/selfoptimize/selfoptimize.go b/go/internal/selfoptimize/selfoptimize.go index 9a6a695..f63eb90 100644 --- a/go/internal/selfoptimize/selfoptimize.go +++ b/go/internal/selfoptimize/selfoptimize.go @@ -271,6 +271,17 @@ func (ct *CausalTracker) IsVolatile(memoryID string) bool { return len(ct.entries[memoryID]) >= 3 } +// Entries 返回全部版本历史(只读) +func (ct *CausalTracker) Entries() map[string][]*TraceEntry { + ct.mu.RLock() + defer ct.mu.RUnlock() + cp := make(map[string][]*TraceEntry, len(ct.entries)) + for k, v := range ct.entries { + cp[k] = v + } + return cp +} + // ─── 记忆预取 ──────────────────────────────────────────── type PrefetchGraph struct { diff --git a/go/internal/storage/recall.go b/go/internal/storage/recall.go index 349fa76..ce16d7b 100644 --- a/go/internal/storage/recall.go +++ b/go/internal/storage/recall.go @@ -1,4 +1,4 @@ -// Package storage 提供记忆召回完整管线:编码 → 向量搜索 → 重排 → MMR 多样性。 +// Package storage 提供记忆召回完成管线:编码 → 向量搜索 → 重排 → MMR → 图谱扩展 → 预取推送 package storage import ( @@ -9,14 +9,24 @@ import ( "github.com/xiaoxue/memoryweave/internal/models" ) -// RecallPipeline 组合 Embedder、LanceClient、Reranker 构成完整召回管线。 -type RecallPipeline struct { - embedder *Embedder - lancedb *LanceClient - reranker *Reranker +// 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 *LanceClient + reranker *Reranker + graph GraphExpander + prefetch PrefetchPusher } -// NewRecallPipeline 创建召回管线。 func NewRecallPipeline(embedder *Embedder, lancedb *LanceClient, reranker *Reranker) *RecallPipeline { return &RecallPipeline{ embedder: embedder, @@ -25,11 +35,16 @@ func NewRecallPipeline(embedder *Embedder, lancedb *LanceClient, reranker *Reran } } -// Recall 完整召回流程: -// 1. query → bge-m3 编码 → 1024d 向量 -// 2. LanceDB ANN 搜索 → top 50 候选(粗排) -// 3. bge-reranker-v2-m3 重排 → top_k(精排) -// 4. MMR 多样性 → 最终结果 +// SetGraphExpander 注入图谱扩展器 +func (p *RecallPipeline) SetGraphExpander(g GraphExpander) { + p.graph = g +} + +// SetPrefetchPusher 注入预取推送器 +func (p *RecallPipeline) SetPrefetchPusher(pf PrefetchPusher) { + p.prefetch = pf +} + func (p *RecallPipeline) Recall(query, namespace string, topK int, diversity float64) ([]models.RecallResult, error) { if topK <= 0 { topK = 10 @@ -55,108 +70,129 @@ func (p *RecallPipeline) Recall(query, namespace string, topK int, diversity flo for i, c := range candidates { docs[i] = c.Content } - reranked, err := p.reranker.Rerank(query, docs, topK) + reranked, err := p.reranker.Rerank(query, docs, topK*2) if err != nil { return nil, fmt.Errorf("recall rerank: %w", err) } // Step 4: MMR diversity - finalResults := mmrSelect(reranked, candidates, topK, diversity) + finalIndices := mmrSelect(reranked, candidates, topK*2, diversity) - // Convert to RecallResult - results := make([]models.RecallResult, len(finalResults)) - for i, idx := range finalResults { + 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 - if idx < len(reranked) { - for _, r := range reranked { - if r.Index == idx { - score = r.Score - break - } + for _, r := range reranked { + if r.Index == idx { + score = r.Score + break } } - results[i] = models.RecallResult{ + results = append(results, models.RecallResult{ ID: c.ID, Content: c.Content, Category: c.Category, Score: score, 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) } } } - // 异步更新 recall_count(非阻塞) + // Step 6: 预取推送(CO_OCCURS 权重 > 0.6 的配套记忆 → WebSocket) + if p.prefetch != nil { + prefetchItems := make([]models.RecallResult, 0) + // 收集所有 results 的 ID 作为共访候选项 + for _, r := range results { + if r.Score > 0.6 { + prefetchItems = append(prefetchItems, r) + } + } + if len(prefetchItems) > 0 { + p.prefetch.PushPrefetch("", prefetchItems) + } + } + + // Async increment recall_count go p.incrementRecallCount(results) return results, nil } -// mmrSelect 最大边际相关性选择:平衡相关性与多样性。 -// diversity=0 → 纯相关性排序,diversity=1 → 最大多样性。 func mmrSelect(reranked []models.RerankResult, candidates []models.MemoryRecord, k int, lambda float64) []int { if len(reranked) == 0 { return nil } - selected := make([]int, 0, k) - candidateIndices := make([]int, len(reranked)) + cand := make([]int, len(reranked)) for i := range reranked { - candidateIndices[i] = reranked[i].Index + cand[i] = reranked[i].Index } - for len(selected) < k && len(candidateIndices) > 0 { + for len(selected) < k && len(cand) > 0 { bestIdx := 0 bestScore := -math.MaxFloat64 - - for i, ci := range candidateIndices { - relevance := 0.0 + for i, ci := range cand { + rel := 0.0 for _, r := range reranked { if r.Index == ci { - relevance = r.Score + rel = r.Score break } } - - // Max similarity to already selected maxSim := 0.0 for _, si := range selected { - sim := cosineSimilarity(candidates[ci].Vector, candidates[si].Vector) - if sim > maxSim { - maxSim = sim + if si < len(candidates) && ci < len(candidates) { + sim := cosineSimilarity(candidates[ci].Vector, candidates[si].Vector) + if sim > maxSim { + maxSim = sim + } } } - - mmr := (1-lambda)*relevance - lambda*maxSim + mmr := (1-lambda)*rel - lambda*maxSim if mmr > bestScore { bestScore = mmr bestIdx = i } } - - selected = append(selected, candidateIndices[bestIdx]) - // Remove from candidates - candidateIndices = append(candidateIndices[:bestIdx], candidateIndices[bestIdx+1:]...) + selected = append(selected, cand[bestIdx]) + cand = append(cand[:bestIdx], cand[bestIdx+1:]...) } - return selected } -// cosineSimilarity 计算两个向量的余弦相似度(向量已 L2 归一化时等于内积)。 func cosineSimilarity(a, b []float32) float64 { var sum float64 - minLen := len(a) - if len(b) < minLen { - minLen = len(b) + n := len(a) + if len(b) < n { + n = len(b) } - for i := 0; i < minLen; i++ { + for i := 0; i < n; i++ { sum += float64(a[i]) * float64(b[i]) } return sum } -// incrementRecallCount 异步递增被召回记忆的 recall_count。 func (p *RecallPipeline) incrementRecallCount(results []models.RecallResult) { for _, r := range results { if r.ID != "" {