303 lines
9.6 KiB
Rust
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()));
|
|
}
|
|
}
|