Files
agcore/src/memory/retriever.rs
T

303 lines
9.6 KiB
Rust

//! 记忆检索器 —— 基于 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 || 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()));
}
}