Skip to main content

p_memory/
preset.rs

1//! 预设检索:库预先配好的几种搜索方法。
2//!
3//! 每次返回固定的三个字段——记忆、图谱、笔记——各自独立排序,不混在一起。
4//! 图谱字段内部按流程的阶段分块:种子实体一块、命中的关系一块、铺开的关系与事件一块。
5//! 预设里的阈值(字符数、种子实体个数)都带默认值,调用方实例化时可以覆盖。
6
7use rusqlite::{params_from_iter, types::Value as SqlValue, Connection};
8use serde::{de::DeserializeOwned, Deserialize, Serialize};
9use serde_json::Value;
10use std::collections::{BTreeMap, BTreeSet};
11
12use crate::graph::{Entity, Event, Relation};
13use crate::search::{SearchHit, SearchRequest, SearchResult};
14use crate::storage::{self, KnowledgeBase};
15use crate::types::{MatchField, ReadFilter, RecordKind, SearchDiagnostics};
16use crate::{text, Error, Result};
17
18/// 库预先配好的搜索方法。
19#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
20#[serde(rename_all = "snake_case")]
21pub enum SearchPreset { Memory, Graph, Notes, Rag, Broad }
22
23impl SearchPreset {
24    pub const ALL: [Self; 5] = [Self::Memory, Self::Graph, Self::Notes, Self::Rag, Self::Broad];
25
26    pub fn as_str(self) -> &'static str {
27        match self {
28            Self::Memory => "memory", Self::Graph => "graph", Self::Notes => "notes",
29            Self::Rag => "rag", Self::Broad => "broad",
30        }
31    }
32
33    pub fn parse(value: &str) -> Result<Self> {
34        Self::ALL.into_iter().find(|preset| preset.as_str() == value)
35            .ok_or_else(|| Error::Validation(format!("preset must be memory, graph, notes, rag or broad, got {value}")))
36    }
37
38    /// 这次要不要出记忆那一路。
39    pub fn uses_memory(self) -> bool { matches!(self, Self::Memory | Self::Rag | Self::Broad) }
40    /// 这次要不要出图谱那一路。
41    pub fn uses_graph(self) -> bool { matches!(self, Self::Graph | Self::Rag | Self::Broad) }
42    /// 这次要不要出笔记那一路。
43    pub fn uses_notes(self) -> bool { matches!(self, Self::Notes | Self::Broad) }
44}
45
46/// 预设的阈值。输出量按字符数封顶(重排模型按字符数算,不看条数),
47/// 种子实体是中间量,仍按个数。
48#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
49#[serde(default)]
50pub struct PresetBudget {
51    /// 记忆那一路的字符数上限。
52    pub memory_chars: usize,
53    /// 笔记那一路的字符数上限。
54    pub notes_chars: usize,
55    /// 图谱那一路取前几个实体当种子。
56    pub seed_entities: usize,
57    /// 图谱第二步命中的关系的字符数上限。
58    pub graph_relations_chars: usize,
59    /// 图谱第三步铺开的关系与事件的字符数上限(关系与事件共用这一份)。
60    pub graph_context_chars: usize,
61    /// 笔记那一路「书名块」想要的条数;也是路径兜底的触发线——书名块不足这么多才用地名补。
62    pub note_titles: usize,
63}
64
65impl Default for PresetBudget {
66    fn default() -> Self {
67        Self { memory_chars: 2000, notes_chars: 3000, seed_entities: 4,
68            graph_relations_chars: 1000, graph_context_chars: 2000, note_titles: 5 }
69    }
70}
71
72#[derive(Debug, Clone, Serialize, Deserialize)]
73#[serde(default)]
74pub struct PresetRequest {
75    pub preset: SearchPreset,
76    pub query: String,
77    pub filter: ReadFilter,
78    /// 走向量路时用哪条向量空间;不给就只走全文。
79    pub embed_space: Option<String>,
80    /// 本次是否走全文路。
81    pub text: bool,
82    /// 本次是否走向量路(需要 `embed_space`)。
83    pub vector: bool,
84    /// 本次是否重排;未注册重排回调时被忽略。
85    pub rerank: bool,
86    /// 覆盖默认阈值。
87    pub budget: PresetBudget,
88    /// 每条路按字符数封顶之前,先取多少条候选。
89    pub candidate_limit: usize,
90}
91
92impl Default for PresetRequest {
93    fn default() -> Self {
94        Self { preset: SearchPreset::Rag, query: String::new(), filter: ReadFilter::default(),
95            embed_space: None, text: true, vector: true, rerank: true,
96            budget: PresetBudget::default(), candidate_limit: 64 }
97    }
98}
99
100/// 图谱那一路的产物,按流程阶段分块。
101#[derive(Debug, Clone, Default, Serialize, Deserialize)]
102pub struct GraphSection {
103    /// 第一步搜出来的种子实体(按分排,带别名)。
104    pub entities: Vec<Entity>,
105    /// 第二步:种子实体各自到自己的关系里敲查询词,命中的关系(按相关度排)。
106    pub relations: Vec<Relation>,
107    /// 第三步:这批实体两两之间的关系,不筛。
108    pub context_relations: Vec<Relation>,
109    /// 第三步:这批实体两两之间的事件(参与者至少两个落在集合内),不筛。
110    pub context_events: Vec<Event>,
111}
112
113/// 笔记那一路的产物,按「命中在哪」分块。
114#[derive(Debug, Clone, Default, Serialize, Deserialize)]
115pub struct NoteSection {
116    /// 书名块:文件名命中,排在最前。
117    pub titles: Vec<SearchHit>,
118    /// 内容块:切片正文命中,已去掉进了书名块的那些。
119    pub contents: Vec<SearchHit>,
120    /// 路径兜底:书名块条数不够时,用目录段补足的、排在最后的一批。
121    pub paths: Vec<SearchHit>,
122}
123
124/// 一次预设检索的结果。三个字段各自独立排序,没走的那一路是空的。
125#[derive(Debug, Clone, Serialize, Deserialize)]
126pub struct PresetResult {
127    pub preset: SearchPreset,
128    pub memories: Vec<SearchHit>,
129    pub graph: GraphSection,
130    pub notes: NoteSection,
131    pub revision: i64,
132    pub indexed_revision: i64,
133    pub diagnostics: SearchDiagnostics,
134}
135
136impl KnowledgeBase {
137    /// 按预设检索。三个字段各自独立排序,互不挤占。
138    pub fn search_preset(&self, request: &PresetRequest) -> Result<PresetResult> {
139        let query = request.query.trim();
140        if query.is_empty() { return Err(Error::Validation("a text query is required".into())); }
141        storage::validate_filter(&request.filter)?;
142        if request.budget.seed_entities == 0 { return Err(Error::Validation("seed_entities must be at least 1".into())); }
143        if request.candidate_limit == 0 { return Err(Error::Validation("candidate_limit must be at least 1".into())); }
144
145        let mut diagnostics = SearchDiagnostics::default();
146        let memories = if request.preset.uses_memory() {
147            let result = self.preset_hits(query, request, &[RecordKind::Memory], MatchField::All)?;
148            merge_diagnostics(&mut diagnostics, &result.diagnostics);
149            self.truncate_hits_by_chars(result.hits, request.budget.memory_chars)?
150        } else { Vec::new() };
151        let notes = if request.preset.uses_notes() {
152            self.preset_notes(query, request, &mut diagnostics)?
153        } else { NoteSection::default() };
154        let graph = if request.preset.uses_graph() {
155            self.preset_graph(query, request, &mut diagnostics)?
156        } else { GraphSection::default() };
157
158        let (revision, indexed_revision) = {
159            let state = self.read()?;
160            (storage::current_revision(state.conn())?, storage::meta(state.conn(), "indexed_revision")?)
161        };
162        Ok(PresetResult { preset: request.preset, memories, graph, notes, revision, indexed_revision, diagnostics })
163    }
164
165    /// 单一路的候选:复用现有的全文 + 向量融合,再按字符数截断。
166    fn preset_hits(&self, query: &str, request: &PresetRequest, kinds: &[RecordKind], match_field: MatchField) -> Result<SearchResult> {
167        let inner = SearchRequest {
168            query: query.to_string(), filter: request.filter.clone(), kinds: kinds.to_vec(),
169            limit: request.candidate_limit, embed_space: request.embed_space.clone(),
170            text: request.text, vector: request.vector, rerank: request.rerank,
171            match_field, ..Default::default()
172        };
173        self.search(&inner)
174    }
175
176    /// 笔记那一路:书名块在上、内容块在下,书名块条数不够时用地名兜底补足。
177    /// 三条探针分工不同——书名块与目录兜底是纯全文、只看名字/目录列、不走向量也不重排
178    /// (向量是语义相似,不属于「书名」);内容块按正文列取,沿用本次请求的全文 + 向量 + 重排。
179    fn preset_notes(&self, query: &str, request: &PresetRequest, diagnostics: &mut SearchDiagnostics) -> Result<NoteSection> {
180        // 书名块:文件名列,纯全文。
181        let title_result = self.preset_field_hits(query, request, MatchField::Name)?;
182        merge_diagnostics(diagnostics, &title_result.diagnostics);
183        let titles: Vec<SearchHit> = title_result.hits.into_iter().take(request.budget.note_titles).collect();
184        let title_ids: BTreeSet<i64> = titles.iter().map(|hit| hit.key.id).collect();
185
186        // 内容块:正文列,去掉已经进了书名块的 id,再按字符数封顶。
187        let content_result = self.preset_hits(query, request, &[RecordKind::Chunk], MatchField::Text)?;
188        merge_diagnostics(diagnostics, &content_result.diagnostics);
189        let content_hits: Vec<SearchHit> = content_result.hits.into_iter()
190            .filter(|hit| !title_ids.contains(&hit.key.id)).collect();
191        let contents = self.truncate_hits_by_chars(content_hits, request.budget.notes_chars)?;
192        let mut surfaced: BTreeSet<i64> = title_ids;
193        surfaced.extend(contents.iter().map(|hit| hit.key.id));
194
195        // 路径兜底:书名块条数不足才启用,只补差额,并排除已经出现在书名或内容块里的。
196        let mut paths = Vec::new();
197        if titles.len() < request.budget.note_titles {
198            let need = request.budget.note_titles - titles.len();
199            let path_result = self.preset_field_hits(query, request, MatchField::Path)?;
200            merge_diagnostics(diagnostics, &path_result.diagnostics);
201            paths = path_result.hits.into_iter()
202                .filter(|hit| !surfaced.contains(&hit.key.id)).take(need).collect();
203        }
204        Ok(NoteSection { titles, contents, paths })
205    }
206
207    /// 只看某一列的全文字段探针:不走向量、不重排,用于书名块与路径兜底。
208    fn preset_field_hits(&self, query: &str, request: &PresetRequest, match_field: MatchField) -> Result<SearchResult> {
209        let inner = SearchRequest {
210            query: query.to_string(), filter: request.filter.clone(), kinds: vec![RecordKind::Chunk],
211            limit: request.candidate_limit.max(request.budget.note_titles),
212            text: true, vector: false, rerank: false, match_field, ..Default::default()
213        };
214        self.search(&inner)
215    }
216
217    /// 按字符数封顶。长度取索引里的纯正文——切片正文只存在索引里;
218    /// 取不到正文的条目按 0 计。第一条无论多长都留下,避免「预算略小就整条路空掉」。
219    fn truncate_hits_by_chars(&self, hits: Vec<SearchHit>, chars: usize) -> Result<Vec<SearchHit>> {
220        if chars == 0 || hits.is_empty() { return Ok(Vec::new()); }
221        let ids: Vec<i64> = hits.iter().map(|hit| hit.key.id).collect();
222        // 索引取不到正文不该让整次检索失败:这一路退化成只按条数,不打断别的路。
223        let bodies = match self.index() { Ok(index) => index.bodies(&ids).unwrap_or_default(), Err(_) => BTreeMap::new() };
224        let mut used = 0usize;
225        let mut out = Vec::new();
226        for hit in hits {
227            let len = bodies.get(&hit.key.id).map(|body| body.chars().count()).unwrap_or(0);
228            if !out.is_empty() && used + len > chars { break; }
229            used += len;
230            out.push(hit);
231        }
232        Ok(out)
233    }
234
235    /// 图谱那一路:搜实体 → 种子各自敲自己的关系 → 两两之间铺开关系与事件。
236    fn preset_graph(&self, query: &str, request: &PresetRequest, diagnostics: &mut SearchDiagnostics) -> Result<GraphSection> {
237        // 第一步:查询词搜实体,取前几个当种子。
238        let entity_result = self.preset_hits(query, request, &[RecordKind::Entity], MatchField::All)?;
239        merge_diagnostics(diagnostics, &entity_result.diagnostics);
240        let seeds: Vec<i64> = entity_result.hits.iter().take(request.budget.seed_entities).map(|hit| hit.key.id).collect();
241        if seeds.is_empty() { return Ok(GraphSection::default()); }
242
243        let state = self.read()?;
244        let conn = state.conn();
245        // 种子实体连别名一起给出:实体名与别名都在它的正文里,取回的是记录本体。
246        let entity_filter = ReadFilter { tags: vec![], ..request.filter.clone() };
247        let loaded: BTreeMap<i64, Entity> = storage::load_many(conn, &seeds, &entity_filter)?;
248        let entities: Vec<Entity> = seeds.iter().filter_map(|id| loaded.get(id).cloned()).collect();
249
250        // 第二步:种子实体各自到「以它为端点」的关系里敲词,命中的留下并按相关度排。
251        // 敲的是「查询去掉实体名(连别名)之后的剩余词」——关系正文里本来就写着主语名,
252        // 拿实体名去敲几乎恒真、等于空转,还会把与谓词无关的关系带进来。剩余为空表示查询
253        // 就是实体名本身、没有谓词可敲,此时不筛,保留种子的全部端点关系。
254        let candidate_ids = incident_relations(conn, &seeds, &request.filter)?;
255        let candidates = storage::record_values(conn, &candidate_ids)?;
256        // 与全文路同一套扩散:剩余词里写「beta」时,也带上「alpha」这一类同义谓词。
257        let remainder = remaining_query_after_entities(query, &entities);
258        let graph_query = crate::graph::match_predicate_synonyms(conn, &request.filter.namespace, &remainder)?;
259        let ranked = rank_relations_by_query(&candidates, &graph_query);
260        let hit_ids = truncate_values_by_chars(&candidates, &ranked, RecordKind::Relation, request.budget.graph_relations_chars);
261        let relations: Vec<Relation> = decode_all(&candidates, &hit_ids)?;
262
263        // 第三步:实体集合 = 种子 + 命中关系的另一端;两两之间的关系与事件全部保留,不筛。
264        let mut members: BTreeSet<i64> = seeds.iter().copied().collect();
265        for relation in &relations { members.insert(relation.subject_id); members.insert(relation.object_id); }
266        let member_ids: Vec<i64> = members.into_iter().collect();
267        let hit_set: BTreeSet<i64> = hit_ids.iter().copied().collect();
268
269        let between_ids: Vec<i64> = relations_between(conn, &member_ids, &request.filter)?
270            .into_iter().filter(|id| !hit_set.contains(id)).collect();
271        let between_values = storage::record_values(conn, &between_ids)?;
272        let between_order: Vec<i64> = between_ids.iter().copied().filter(|id| between_values.contains_key(id)).collect();
273        let kept_relations = truncate_values_by_chars(&between_values, &between_order, RecordKind::Relation, request.budget.graph_context_chars);
274        let used = chars_of(&between_values, &kept_relations, RecordKind::Relation);
275        let context_relations: Vec<Relation> = decode_all(&between_values, &kept_relations)?;
276
277        let event_ids = events_between(conn, &member_ids, &request.filter)?;
278        let remaining = request.budget.graph_context_chars.saturating_sub(used);
279        // 事件这一批常常比预算大一到两个数量级(成员是高连接度实体时尤甚):先按正文长度
280        // 定下要留哪几条,再装配这几条。装配整批等于把它们全读出来解一遍再丢掉。
281        let lengths = storage::event_text_lengths(conn, &event_ids)?;
282        let event_order: Vec<i64> = event_ids.iter().copied().filter(|id| lengths.contains_key(id)).collect();
283        let kept_events = truncate_ids_by_lengths(&event_order, &lengths, remaining);
284        let event_values = storage::record_values(conn, &kept_events)?;
285        let context_events: Vec<Event> = decode_all(&event_values, &kept_events)?;
286
287        Ok(GraphSection { entities, relations, context_relations, context_events })
288    }
289}
290
291/// 把子路的诊断并进整次检索的诊断:开关取或,候选数累加,降级去重。
292fn merge_diagnostics(target: &mut SearchDiagnostics, source: &SearchDiagnostics) {
293    target.text_used |= source.text_used;
294    target.vector_used |= source.vector_used;
295    target.reranked |= source.reranked;
296    target.rerank_candidates += source.rerank_candidates;
297    target.rerank_truncated += source.rerank_truncated;
298    for degrade in &source.degraded {
299        if !target.degraded.contains(degrade) { target.degraded.push(*degrade); }
300    }
301}
302
303/// 查询词元在这段文本里的权重:命中的二字及以上词元算两份,单字算一份。
304fn token_weight(body: &str, tokens: &[String]) -> usize {
305    let present: BTreeSet<String> = text::tokenize(body).into_iter().collect();
306    tokens.iter().map(|token| {
307        if !present.contains(token) { 0 } else if token.chars().count() > 1 { 2 } else { 1 }
308    }).sum()
309}
310
311/// 第二步要敲的词:查询里去掉种子实体名(连别名)之后的剩余部分。
312/// 关系正文里本就写着主语名,拿实体名去敲几乎恒真、等于空转,去掉后剩下的才是真正的谓词线索。
313/// 名字或别名没出现在查询里时不改动。剩余为空(查询本身就是实体名)时返回空串,
314/// 交给 `rank_relations_by_query` 按「无词元即不筛」处理。
315fn remaining_query_after_entities(query: &str, entities: &[Entity]) -> String {
316    let mut remainder = query.to_string();
317    for entity in entities {
318        for term in std::iter::once(&entity.name).chain(entity.aliases.iter()) {
319            if !term.is_empty() { remainder = remainder.replace(term.as_str(), " "); }
320        }
321    }
322    remainder
323}
324
325/// 第二步的排序:命中权重高的在前,同权重按记录 id 升序,保证结果稳定。
326fn rank_relations_by_query(values: &BTreeMap<i64, Value>, query: &str) -> Vec<i64> {
327    let tokens = text::query_terms(query, true);
328    if tokens.is_empty() { return values.keys().copied().collect(); }
329    let mut scored: Vec<(usize, i64)> = values.iter()
330        .map(|(id, payload)| (token_weight(&storage::record_text(RecordKind::Relation, payload), &tokens), *id))
331        .filter(|(weight, _)| *weight > 0)
332        .collect();
333    scored.sort_by(|a, b| b.0.cmp(&a.0).then_with(|| a.1.cmp(&b.1)));
334    scored.into_iter().map(|(_, id)| id).collect()
335}
336
337/// 按字符数封顶,顺序不变;第一条无论多长都留下。
338fn truncate_values_by_chars(values: &BTreeMap<i64, Value>, order: &[i64], kind: RecordKind, chars: usize) -> Vec<i64> {
339    if chars == 0 { return Vec::new(); }
340    let mut used = 0usize;
341    let mut kept = Vec::new();
342    for id in order {
343        let Some(payload) = values.get(id) else { continue };
344        let len = storage::record_text(kind, payload).chars().count();
345        if !kept.is_empty() && used + len > chars { break; }
346        used += len;
347        kept.push(*id);
348    }
349    kept
350}
351
352/// 与 `truncate_values_by_chars` 同一条规则(顺序不变、第一条无论多长都留下),
353/// 只是长度由调用方先查好:要留的条数远小于取回的条数时,先定 id、再装配那几条。
354fn truncate_ids_by_lengths(order: &[i64], lengths: &BTreeMap<i64, usize>, chars: usize) -> Vec<i64> {
355    if chars == 0 { return Vec::new(); }
356    let mut used = 0usize;
357    let mut kept = Vec::new();
358    for id in order {
359        let Some(len) = lengths.get(id) else { continue };
360        if !kept.is_empty() && used + len > chars { break; }
361        used += len;
362        kept.push(*id);
363    }
364    kept
365}
366
367fn chars_of(values: &BTreeMap<i64, Value>, ids: &[i64], kind: RecordKind) -> usize {
368    ids.iter().filter_map(|id| values.get(id)).map(|payload| storage::record_text(kind, payload).chars().count()).sum()
369}
370
371fn decode_all<T: DeserializeOwned>(values: &BTreeMap<i64, Value>, ids: &[i64]) -> Result<Vec<T>> {
372    let mut out = Vec::new();
373    for id in ids {
374        if let Some(payload) = values.get(id) { out.push(serde_json::from_value(payload.clone())?); }
375    }
376    Ok(out)
377}
378
379/// 以这批实体为端点的关系 id(它是主语或宾语都算)。
380fn incident_relations(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
381    let (condition, values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
382    let placeholders = vec!["?"; entity_ids.len()].join(",");
383    let sql = format!("SELECT rl.record_id FROM relations rl JOIN records r ON r.id=rl.record_id \
384        WHERE (rl.subject_id IN ({placeholders}) OR rl.object_id IN ({placeholders})) AND {condition} ORDER BY rl.record_id");
385    let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id))
386        .chain(entity_ids.iter().map(|id| SqlValue::Integer(*id))).chain(values).collect();
387    let mut stmt = conn.prepare(&sql)?;
388    let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
389    Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
390}
391
392/// 两端都落在这批实体里的关系 id。
393fn relations_between(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
394    let (condition, values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
395    let placeholders = vec!["?"; entity_ids.len()].join(",");
396    let sql = format!("SELECT rl.record_id FROM relations rl JOIN records r ON r.id=rl.record_id \
397        WHERE rl.subject_id IN ({placeholders}) AND rl.object_id IN ({placeholders}) AND {condition} ORDER BY rl.record_id");
398    let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id))
399        .chain(entity_ids.iter().map(|id| SqlValue::Integer(*id))).chain(values).collect();
400    let mut stmt = conn.prepare(&sql)?;
401    let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
402    Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
403}
404
405/// 参与者里至少有两个落在这批实体里的事件 id。
406///
407/// 先在参与者表上按 `entity_id` 收窄并分组筛,再回 `records` 过滤领域:领域条件是挂在
408/// `event_id` 上的,同一个事件的所有参与行要么全过、要么全不过,所以先筛后筛结果相同。
409/// 反过来拿 `records` 当驱动表,会对领域内每一个事件逐条探测参与者表,代价高一个数量级。
410fn events_between(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
411    let (condition, values) = storage::filter_sql(filter, &[RecordKind::Event], false)?;
412    let placeholders = vec!["?"; entity_ids.len()].join(",");
413    let sql = format!("SELECT e.event_id FROM (SELECT event_id FROM event_participants \
414        WHERE entity_id IN ({placeholders}) GROUP BY event_id HAVING COUNT(DISTINCT entity_id) >= 2) e \
415        JOIN records r ON r.id=e.event_id WHERE {condition} ORDER BY e.event_id");
416    let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id)).chain(values).collect();
417    let mut stmt = conn.prepare(&sql)?;
418    let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
419    Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
420}