package memory import ( "context" "fmt" "strings" "git.yeij.top/AskaEth/Cyrene/ai-core/internal/model" ) // MemoryEntry 记忆条目别名(避免与model包冲突) type MemoryEntry = model.MemoryEntry // Retriever 记忆检索器 type Retriever struct { store *Store embedder Embedder // 文本转向量的接口 } // Embedder 文本嵌入接口 type Embedder interface { Embed(ctx context.Context, text string) ([]float64, error) IsAvailable() bool } // SimpleEmbedder 基于 n-gram 哈希的本地嵌入(无需外部API) // 用于嵌入API不可用时的降级方案,中文效果显著优于字符频率哈希。 type SimpleEmbedder struct{} const embedDim = 1536 // Embed 使用字符 bigram + trigram 哈希生成稀疏向量。 // 相似的中文短语会共享 n-gram → 哈希碰撞产生有意义的余弦相似度。 func (e *SimpleEmbedder) Embed(ctx context.Context, text string) ([]float64, error) { vec := make([]float64, embedDim) runes := []rune(strings.TrimSpace(text)) if len(runes) == 0 { return vec, nil } // 统计 n-gram 频率(bigram + trigram),用 TF 加权 grams := make(map[uint64]float64) addGram := func(start, n int) { if start+n > len(runes) { return } h := hashRunes(runes[start : start+n]) grams[h] += 1.0 } for i := 0; i < len(runes); i++ { addGram(i, 2) // bigram addGram(i, 3) // trigram } // 单字也加入,捕获关键词 for i := 0; i < len(runes); i++ { h := hashRunes(runes[i : i+1]) grams[h] += 0.3 } // 归一化后写入向量 var total float64 for _, v := range grams { total += v * v } if total == 0 { return vec, nil } norm := 1.0 / total // approximate L2 norm for h, v := range grams { idx := int(h % embedDim) vec[idx] += v * norm } return vec, nil } // hashRunes computes a simple FNV-like hash of rune slice. func hashRunes(r []rune) uint64 { var h uint64 = 14695981039346656037 for _, c := range r { h ^= uint64(c) h *= 1099511628211 } return h } // IsAvailable returns true (SimpleEmbedder is always available as fallback). func (e *SimpleEmbedder) IsAvailable() bool { return true } // NewRetriever 创建记忆检索器 func NewRetriever(store *Store, embedder Embedder) *Retriever { if embedder == nil { embedder = &SimpleEmbedder{} } return &Retriever{ store: store, embedder: embedder, } } // Retrieve 检索与查询相关的记忆 // 策略: 向量相似度 + 关键词匹配混合 → 按重要性降序返回 func (r *Retriever) Retrieve(ctx context.Context, userID string, query string) ([]MemoryEntry, error) { var allEntries []MemoryEntry seen := make(map[string]bool) // 1. 向量相似度检索 embedding, err := r.embedder.Embed(ctx, query) if err == nil { vecEntries, err := r.store.SearchByVector(ctx, userID, embedding, 8) if err == nil { for _, e := range vecEntries { if !seen[e.ID] { seen[e.ID] = true allEntries = append(allEntries, e) } } } } // 2. 关键词匹配检索(包含关键词标签匹配) keywordEntries, err := r.keywordSearch(ctx, userID, query) if err == nil { for _, e := range keywordEntries { if !seen[e.ID] { seen[e.ID] = true allEntries = append(allEntries, e) } } } // 3. 如果没有匹配,返回最近的重要记忆 if len(allEntries) == 0 { recentEntries, err := r.store.Query(ctx, model.MemoryQuery{ UserID: userID, Priority: model.MemoryImportant, Limit: 5, }) if err == nil { allEntries = recentEntries } } // 4. 去重合并:对高度相似的记忆只保留Importance更高的 allEntries = r.deduplicate(allEntries) // 5. 按重要性降序排列 sortByImportance(allEntries) // 限制返回数量 if len(allEntries) > 10 { allEntries = allEntries[:10] } return allEntries, nil } // RetrieveByCategory 按分类检索记忆 func (r *Retriever) RetrieveByCategory(ctx context.Context, userID string, category model.MemoryCategory, limit int) ([]MemoryEntry, error) { if limit <= 0 { limit = 20 } return r.store.Query(ctx, model.MemoryQuery{ UserID: userID, Category: category, Limit: limit, }) } // keywordSearch 关键词匹配检索(包含关键词标签和n-gram分词匹配) func (r *Retriever) keywordSearch(ctx context.Context, userID string, query string) ([]MemoryEntry, error) { // 将查询切分为中文分词 tokens(2-4字的滑动窗口) queryRunes := []rune(query) var tokens []string for size := 2; size <= 4; size++ { for i := 0; i+size <= len(queryRunes); i++ { tokens = append(tokens, string(queryRunes[i:i+size])) } } // 也保留完整查询 tokens = append(tokens, query) // 查询记忆 entries, err := r.store.Query(ctx, model.MemoryQuery{ UserID: userID, Limit: 100, }) if err != nil { return nil, err } type scoredEntry struct { entry MemoryEntry score int } var scored []scoredEntry seen := make(map[string]bool) for _, entry := range entries { if seen[entry.ID] { continue } s := 0 contentLower := strings.ToLower(entry.Content) summaryLower := strings.ToLower(entry.Summary) for _, tok := range tokens { tokLower := strings.ToLower(tok) if len([]rune(tok)) < 2 { continue // skip single-char matches (too noisy) } if strings.Contains(contentLower, tokLower) { s += 2 // content match is stronger } if strings.Contains(summaryLower, tokLower) { s += 3 // summary match is even stronger (distilled info) } } // 关键词标签直接匹配 for _, kw := range entry.Keywords { kwLower := strings.ToLower(kw) for _, tok := range tokens { if strings.Contains(strings.ToLower(tok), kwLower) || strings.Contains(kwLower, strings.ToLower(tok)) { s += 2 break } } } if s > 0 { seen[entry.ID] = true scored = append(scored, scoredEntry{entry, s}) } } // 按分数降序 for i := 0; i < len(scored); i++ { for j := i + 1; j < len(scored); j++ { if scored[j].score > scored[i].score { scored[i], scored[j] = scored[j], scored[i] } } } result := make([]MemoryEntry, 0, len(scored)) for _, s := range scored { result = append(result, s.entry) } return result, nil } // deduplicate 去重合并:对高度相似的记忆只保留 Importance 更高的 func (r *Retriever) deduplicate(entries []MemoryEntry) []MemoryEntry { if len(entries) < 2 { return entries } result := make([]MemoryEntry, 0, len(entries)) discarded := make(map[int]bool) for i := 0; i < len(entries); i++ { if discarded[i] { continue } for j := i + 1; j < len(entries); j++ { if discarded[j] { continue } score := entries[i].SimilarityScore(&entries[j]) if score >= deDupThreshold { // 保留更重要的那条 if entries[j].Importance > entries[i].Importance || (entries[j].Importance == entries[i].Importance && entries[j].Priority > entries[i].Priority) { discarded[i] = true break } else { discarded[j] = true } } } if !discarded[i] { result = append(result, entries[i]) } } return result } // sortByImportance 按 Importance 降序, Priority 降序排列 func sortByImportance(entries []MemoryEntry) { for i := 0; i < len(entries); i++ { for j := i + 1; j < len(entries); j++ { if entries[j].Importance > entries[i].Importance || (entries[j].Importance == entries[i].Importance && entries[j].Priority > entries[i].Priority) { entries[i], entries[j] = entries[j], entries[i] } } } } // Ensure fmt is used var _ = fmt.Sprintf