feat(memory): 添加记忆系统模块
This commit is contained in:
@@ -0,0 +1,297 @@
|
||||
//! 记忆检索器 —— 基于 TextOverlap (Dice 系数) 的单通道关键词检索。
|
||||
|
||||
use std::collections::HashSet;
|
||||
|
||||
use crate::memory::error::MemoryError;
|
||||
use crate::memory::knowledge::KnowledgeStore;
|
||||
use crate::memory::types::KnowledgePage;
|
||||
|
||||
/// 检索器配置。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RetrieverConfig {
|
||||
/// 最大返回条数(默认 20)。
|
||||
pub max_results: usize,
|
||||
/// 最低分数阈值 [0.0, 1.0](默认 0.1)。
|
||||
pub min_score: f32,
|
||||
}
|
||||
|
||||
impl Default for RetrieverConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_results: 20,
|
||||
min_score: 0.1,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 单条带评分的检索结果。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ScoredItem {
|
||||
pub page: KnowledgePage,
|
||||
/// TextOverlap 评分 [0.0, 1.0]
|
||||
pub score: f32,
|
||||
}
|
||||
|
||||
/// 检索结果。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RetrievalResult {
|
||||
pub items: Vec<ScoredItem>,
|
||||
pub query: String,
|
||||
}
|
||||
|
||||
/// 记忆检索器 —— 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分。
|
||||
pub struct MemoryRetriever {
|
||||
knowledge_store: KnowledgeStore,
|
||||
config: RetrieverConfig,
|
||||
/// 停用词表(用于关键词提取)。
|
||||
stop_words: HashSet<String>,
|
||||
}
|
||||
|
||||
impl MemoryRetriever {
|
||||
/// 创建一个新的 MemoryRetriever。
|
||||
pub fn new(knowledge_store: KnowledgeStore, config: RetrieverConfig) -> Self {
|
||||
Self {
|
||||
knowledge_store,
|
||||
config,
|
||||
stop_words: default_stop_words(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 替换停用词表。
|
||||
pub fn with_stop_words(mut self, stop_words: HashSet<String>) -> Self {
|
||||
self.stop_words = stop_words;
|
||||
self
|
||||
}
|
||||
|
||||
/// 检索相关知识页面。
|
||||
pub async fn retrieve(&self, query: &str) -> Result<RetrievalResult, MemoryError> {
|
||||
if query.is_empty() {
|
||||
return Ok(RetrievalResult {
|
||||
items: Vec::new(),
|
||||
query: query.to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
// 1. 关键词提取
|
||||
let keywords = extract_keywords(query, &self.stop_words);
|
||||
|
||||
// 2. 用关键词在 KnowledgeStore 中搜索
|
||||
let mut pages = Vec::new();
|
||||
for keyword in &keywords {
|
||||
let found = self.knowledge_store.search(keyword).await?;
|
||||
for page in found {
|
||||
if !pages.iter().any(|p: &KnowledgePage| p.id == page.id) {
|
||||
pages.push(page);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. TextOverlap 评分
|
||||
let mut items: Vec<ScoredItem> = pages
|
||||
.into_iter()
|
||||
.map(|page| {
|
||||
let score = text_overlap_score(query, &page);
|
||||
ScoredItem { page, score }
|
||||
})
|
||||
.collect();
|
||||
|
||||
// 4. 过滤 → 排序 → 截取
|
||||
items.retain(|i| i.score >= self.config.min_score);
|
||||
items.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
|
||||
items.truncate(self.config.max_results);
|
||||
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 从 query 中提取关键词:按非字母数字字符分割 → 转小写 → 过滤单字符和停用词。
|
||||
fn extract_keywords(query: &str, stop_words: &HashSet<String>) -> Vec<String> {
|
||||
query
|
||||
.split(|c: char| !c.is_alphanumeric())
|
||||
.filter_map(|s| {
|
||||
let lower = s.to_lowercase();
|
||||
if lower.is_empty() || lower.chars().count() < 2 {
|
||||
None
|
||||
} else if stop_words.contains(&lower) {
|
||||
None
|
||||
} else {
|
||||
Some(lower)
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// TextOverlap 评分(基于字符 bigram 的 Dice 系数 + 多字段加权)。
|
||||
///
|
||||
/// 字段权重:title 0.5 + summary 0.3 + content 0.2
|
||||
/// 中文场景按字符级 bigram 处理,不依赖分词器。
|
||||
pub fn text_overlap_score(query: &str, page: &KnowledgePage) -> f32 {
|
||||
let title = text_overlap_dice(query, &page.title);
|
||||
let summary = text_overlap_dice(query, &page.summary);
|
||||
let content = text_overlap_dice(query, &page.content);
|
||||
title * 0.5 + summary * 0.3 + content * 0.2
|
||||
}
|
||||
|
||||
/// Dice 系数(基于字符 bigram)。
|
||||
fn text_overlap_dice(query: &str, text: &str) -> f32 {
|
||||
let q_bigrams = char_bigrams(query);
|
||||
let t_bigrams = char_bigrams(text);
|
||||
if q_bigrams.is_empty() || t_bigrams.is_empty() {
|
||||
return 0.0;
|
||||
}
|
||||
let q_set: HashSet<&String> = q_bigrams.iter().collect();
|
||||
let t_set: HashSet<&String> = t_bigrams.iter().collect();
|
||||
let intersect = q_set.intersection(&t_set).count();
|
||||
let denom = q_set.len() + t_set.len();
|
||||
if denom == 0 {
|
||||
0.0
|
||||
} else {
|
||||
(2.0 * intersect as f32) / denom as f32
|
||||
}
|
||||
}
|
||||
|
||||
/// 提取字符 bigrams(用于字符级 Dice 系数)。
|
||||
fn char_bigrams(s: &str) -> Vec<String> {
|
||||
let chars: Vec<char> = s.chars().collect();
|
||||
chars.windows(2).map(|w| w.iter().collect()).collect()
|
||||
}
|
||||
|
||||
fn default_stop_words() -> HashSet<String> {
|
||||
[
|
||||
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has",
|
||||
"had", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall",
|
||||
"can", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what",
|
||||
"which", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor",
|
||||
"so", "if", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at",
|
||||
"by", "of", "as", "into", "through", "during", "before", "after", "above", "below",
|
||||
]
|
||||
.iter()
|
||||
.map(|s| s.to_string())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::memory::knowledge::KnowledgeStore;
|
||||
use crate::memory::{InMemoryStore, MemoryStore};
|
||||
use std::sync::Arc;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
fn make_page(id: &str, title: &str, summary: &str, content: &str) -> KnowledgePage {
|
||||
let now = OffsetDateTime::now_utc();
|
||||
KnowledgePage {
|
||||
id: id.to_string(),
|
||||
title: title.to_string(),
|
||||
summary: summary.to_string(),
|
||||
content: content.to_string(),
|
||||
tags: Vec::new(),
|
||||
references: Vec::new(),
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_empty_query() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default());
|
||||
let result = retriever.retrieve("").await.unwrap();
|
||||
assert!(result.items.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_finds_relevant_page() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
ks.add_page(make_page(
|
||||
"p1",
|
||||
"LangGraph StateGraph",
|
||||
"state management",
|
||||
"LangGraph uses StateGraph for state machines",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
ks.add_page(make_page("p2", "Other", "unrelated", "nothing matching"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default());
|
||||
let result = retriever.retrieve("LangGraph state").await.unwrap();
|
||||
assert!(!result.items.is_empty());
|
||||
assert_eq!(result.items[0].page.id, "p1");
|
||||
assert!(result.items[0].score > 0.0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_respects_min_score() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
ks.add_page(make_page("p1", "X", "Y", "Z")).await.unwrap();
|
||||
|
||||
let config = RetrieverConfig {
|
||||
max_results: 10,
|
||||
min_score: 0.99,
|
||||
};
|
||||
let retriever = MemoryRetriever::new(ks, config);
|
||||
let result = retriever.retrieve("totally unrelated content").await.unwrap();
|
||||
assert!(result.items.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_respects_max_results() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
for i in 0..5 {
|
||||
ks.add_page(make_page(
|
||||
&format!("p{i}"),
|
||||
"LangGraph",
|
||||
"framework",
|
||||
"agent runtime",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let config = RetrieverConfig {
|
||||
max_results: 2,
|
||||
min_score: 0.0,
|
||||
};
|
||||
let retriever = MemoryRetriever::new(ks, config);
|
||||
let result = retriever.retrieve("LangGraph").await.unwrap();
|
||||
assert_eq!(result.items.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_overlap_dice_zero_on_empty() {
|
||||
assert_eq!(text_overlap_dice("hello", ""), 0.0);
|
||||
assert_eq!(text_overlap_dice("", "hello"), 0.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn text_overlap_dice_identical() {
|
||||
let s = "hello world";
|
||||
assert!((text_overlap_dice(s, s) - 1.0).abs() < 0.001);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extract_keywords_filters_stop_words() {
|
||||
let stop = default_stop_words();
|
||||
let kws = extract_keywords("the quick brown fox is fast", &stop);
|
||||
assert!(!kws.contains(&"the".to_string()));
|
||||
assert!(!kws.contains(&"is".to_string()));
|
||||
assert!(kws.contains(&"quick".to_string()));
|
||||
assert!(kws.contains(&"brown".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extract_keywords_filters_single_chars() {
|
||||
let stop = default_stop_words();
|
||||
let kws = extract_keywords("a b c dog", &stop);
|
||||
assert!(!kws.contains(&"a".to_string()));
|
||||
assert!(!kws.contains(&"b".to_string()));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user