diff --git a/backend/ai-core/internal/memory/retriever.go b/backend/ai-core/internal/memory/retriever.go index 4424dea..bbde1a5 100644 --- a/backend/ai-core/internal/memory/retriever.go +++ b/backend/ai-core/internal/memory/retriever.go @@ -23,27 +23,67 @@ type Embedder interface { IsAvailable() bool } -// SimpleEmbedder 基于关键词的简单嵌入(MVP阶段可用,无需外部API) +// SimpleEmbedder 基于 n-gram 哈希的本地嵌入(无需外部API) +// 用于嵌入API不可用时的降级方案,中文效果显著优于字符频率哈希。 type SimpleEmbedder struct{} -// Embed 简单的关键词哈希嵌入(用于MVP快速验证) -func (e *SimpleEmbedder) Embed(ctx context.Context, text string) ([]float64, error) { - // 生成一个简单的1536维特征向量 - // 基于字符频率的简单表示,用于MVP阶段 - vec := make([]float64, 1536) +const embedDim = 1536 - runes := []rune(strings.ToLower(text)) - for i, r := range runes { - idx := int(r) % 1536 - vec[idx] += 1.0 / float64(len(runes)) - // 考虑位置信息 - posIdx := (int(r) + i) % 1536 - vec[posIdx] += 0.5 / float64(len(runes)) +// 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 } @@ -127,67 +167,87 @@ func (r *Retriever) RetrieveByCategory(ctx context.Context, userID string, categ }) } -// keywordSearch 关键词匹配检索(包含关键词标签匹配) +// 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, - Priority: model.MemoryImportant, - Limit: 50, + UserID: userID, + Limit: 100, }) if err != nil { return nil, err } - // 关键词匹配过滤 - var matched []MemoryEntry - queryLower := strings.ToLower(query) + 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) - // 内容/摘要匹配 - if strings.Contains(contentLower, queryLower) || strings.Contains(summaryLower, queryLower) { - matched = append(matched, entry) - continue + 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 { - if strings.Contains(queryLower, strings.ToLower(kw)) || - strings.Contains(strings.ToLower(kw), queryLower) { - matched = append(matched, entry) - break - } - } - } - - // 也匹配普通记忆 - normalEntries, err := r.store.Query(ctx, model.MemoryQuery{ - UserID: userID, - Priority: model.MemoryNormal, - Limit: 100, - }) - if err == nil { - for _, entry := range normalEntries { - contentLower := strings.ToLower(entry.Content) - summaryLower := strings.ToLower(entry.Summary) - if strings.Contains(contentLower, queryLower) || strings.Contains(summaryLower, queryLower) { - matched = append(matched, entry) - continue - } - for _, kw := range entry.Keywords { - if strings.Contains(queryLower, strings.ToLower(kw)) || - strings.Contains(strings.ToLower(kw), queryLower) { - matched = append(matched, entry) + 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}) + } } - return matched, nil + // 按分数降序 + 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 更高的