Skip to main content

p_memory/
search.rs

1use crate::{embeddings, graph::{Entity, Neighborhood, Relation}, storage::{self, KnowledgeBase}, text, types::*, Error, Result};
2use parking_lot::Mutex;
3use rusqlite::{params_from_iter, types::Value as SqlValue, OptionalExtension};
4use serde::{Deserialize, Serialize};
5use serde_json::Value;
6use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
7use std::sync::Arc;
8
9fn default_limit() -> usize { 10 }
10/// 阶段打点:只在注册了事件 sink 时才计时,没注册就是一次空判断。
11fn mark(stages: &mut Option<crate::events::StageTimer>, name: &str) {
12    if let Some(stages) = stages.as_mut() { stages.mark(name); }
13}
14fn yes() -> bool { true }
15pub fn default_kinds() -> Vec<RecordKind> { vec![RecordKind::Memory, RecordKind::Entity, RecordKind::Relation, RecordKind::Event, RecordKind::Chunk] }
16
17/// 一个标记文本在 strings 表里的 id;库里没有这个标记就是 `None`。
18fn string_id(conn: &rusqlite::Connection, value: &str) -> Result<Option<i64>> {
19    Ok(conn.query_row("SELECT id FROM strings WHERE text=?1", [text::normalized_tag(value)], |r| r.get(0)).optional()?)
20}
21
22/// 把 `filter` 里的标记文本折算成 id:索引里只存 id,折算在进索引之前做完。
23/// 换不到 id 说明库里没有这个标记,本次不可能有命中,直接给空结果,不必进索引碰。
24fn index_filter(conn: &rusqlite::Connection, filter: &ReadFilter, kinds: &[RecordKind]) -> Result<Option<crate::index::IndexFilter>> {
25    let Some(namespace) = string_id(conn, &filter.namespace)? else { return Ok(None) };
26    let mut scopes = Vec::with_capacity(filter.scopes.len());
27    for scope in &filter.scopes {
28        match string_id(conn, scope)? { Some(id) => scopes.push(id), None => return Ok(None) }
29    }
30    let mut tags = Vec::with_capacity(filter.tags.len());
31    for tag in &filter.tags {
32        match string_id(conn, tag)? { Some(id) => tags.push(id), None => return Ok(None) }
33    }
34    Ok(Some(crate::index::IndexFilter { namespace, scopes, kinds: kinds.iter().map(|kind| kind.code()).collect(), tags, note_ids: filter.note_ids.clone() }))
35}
36
37/// Graph-First 剪枝:先在图邻域里取到候选实体 id,向量检索只在这些 id 内打分。
38/// `SearchRequest::prune` 为 `None` 时保持全量检索,不改默认行为。
39#[derive(Debug, Clone, Serialize, Deserialize)]
40pub struct GraphPrune {
41    /// 图展开起点(实体 record_id)。
42    pub root: i64,
43    /// 展开跳数,至少 1。
44    pub depth: usize,
45    /// 邻域节点上限,至少 1。
46    pub limit: usize,
47}
48
49#[derive(Debug, Clone, Serialize, Deserialize)]
50#[serde(default)]
51pub struct SearchRequest {
52    pub query: String, pub filter: ReadFilter, pub kinds: Vec<RecordKind>,
53    pub limit: usize, pub candidate_limit: Option<usize>,
54    /// 走向量路时用哪条向量空间;库用该空间注册的回调嵌入查询词,宿主不接触向量。
55    pub embed_space: Option<String>,
56    pub text_weight: f64, pub prune: Option<GraphPrune>,
57    /// 本次是否走全文路。
58    #[serde(default = "yes")] pub text: bool,
59    /// 本次是否走向量路(需要 `embed_space`)。
60    #[serde(default = "yes")] pub vector: bool,
61    /// 本次是否重排;未注册重排回调时被忽略。
62    #[serde(default = "yes")] pub rerank: bool,
63    /// 是否额外返回过滤后的匹配总量。
64    #[serde(default)] pub with_total: bool,
65    /// 本次全文路限定在哪一列命中:`all`(默认,正文+名字)、`text`、`name`、`path`。
66    #[serde(default)] pub match_field: MatchField,
67    /// 切片折叠后,每条命中最多聚合这一篇里排名最高的几片(含本条)到 `top_chunks`;
68    /// 0 表示不聚合,只留排名最高的那一片。只看名字/目录列时不做聚合(那两列是笔记级的)。
69    #[serde(default = "default_top_chunks_per_note")] pub top_chunks_per_note: usize,
70}
71fn default_top_chunks_per_note() -> usize { 3 }
72impl Default for SearchRequest {
73    fn default() -> Self {
74        Self { query: String::new(), filter: ReadFilter::default(), kinds: default_kinds(), limit: default_limit(),
75            candidate_limit: None, embed_space: None, text_weight: 1.0, prune: None,
76            text: true, vector: true, rerank: true, with_total: false, match_field: MatchField::All,
77            top_chunks_per_note: default_top_chunks_per_note() }
78    }
79}
80/// 折叠后挂在代表命中上的一个片段:同一篇笔记里也命中本次查询的那一片。
81/// 只带定位,不带正文——正文按 `id` 从索引取回(`notes.get_chunk`)。
82#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
83pub struct ChunkRef {
84    /// 切片记录 id。
85    pub id: i64,
86    /// 该片在原文里的起始行(1 起)。
87    pub offset: usize,
88}
89#[derive(Debug, Clone, Serialize, Deserialize)]
90pub struct SearchHit {
91    pub key: RecordKey,
92    /// Reciprocal rank fusion score, k=60. This is not a probability.
93    pub score: f64, pub text_score: Option<f64>, pub vector_scores: BTreeMap<String, f64>,
94    /// 本条目在重排回调那里的分数;未重排时为 `None`。
95    #[serde(default)] pub rerank_score: Option<f64>,
96    /// 本条目是切片时:其所属笔记里命中本次查询的切片总数(按请求的匹配列计,含本条)。
97    /// 非切片命中、或请求只看名字/目录列时为 `None`。
98    /// 它只描述「这篇文档有多少相关片段」,与结果窗口、翻页无关。
99    #[serde(default)] pub note_chunks: Option<usize>,
100    /// 本条代表的那一篇里,命中本次查询、排名最高的若干片段(第 0 条就是本条自身),
101    /// 按名次排列、至多 `top_chunks_per_note` 条。同一篇内的次序取索引里的相关度
102    /// (严格词元命中排在宽松命中之前),与命中本身的 `score` 不同源。
103    /// 被折叠掉的片段在这里找回来,调用方不必为「同一篇还有别的相关片段」再搜一次。
104    /// 非切片命中、或请求只看名字/目录列时为空。
105    #[serde(default)] pub top_chunks: Vec<ChunkRef>,
106    pub record: Value,
107}
108#[derive(Debug, Clone, Serialize, Deserialize)]
109pub struct SearchResult {
110    pub hits: Vec<SearchHit>, pub revision: i64, pub indexed_revision: i64,
111    /// 过滤之后、截断之前的匹配数;`with_total=false` 时为 `None`。
112    #[serde(default)] pub total: Option<usize>,
113    #[serde(default)] pub diagnostics: SearchDiagnostics,
114}
115
116/// 一条命中,附上它在图上挂载的实体邻域。
117#[derive(Debug, Clone, Serialize, Deserialize)]
118pub struct ContextualHit { pub hit: SearchHit, pub context: Neighborhood }
119
120// ── 宿主重排回调 ─────────────────────────────────────────────────────
121
122/// 重排回调:`(查询词, 文档) -> 分数`,与输入文档等长、按序给出。
123/// 重排不绑定向量空间,一个进程内一个回调;多宿主各在自己的进程里注册。
124pub trait Reranker: Send {
125    fn rerank(&mut self, query: &str, documents: &[String]) -> std::result::Result<Vec<f32>, String>;
126}
127
128impl<F> Reranker for F
129where F: FnMut(&str, &[String]) -> std::result::Result<Vec<f32>, String> + Send {
130    fn rerank(&mut self, query: &str, documents: &[String]) -> std::result::Result<Vec<f32>, String> { self(query, documents) }
131}
132
133fn default_max_tokens_total() -> usize { 8192 }
134fn default_max_tokens_per_doc() -> usize { 1024 }
135fn default_max_candidates() -> usize { 50 }
136
137/// 两路候选按名次交替合并、去重:全文第 0 名、向量第 0 名、全文第 1 名、向量第 1 名……
138/// 用于把两路结果直接交给重排——两路各占约一半名额,不偏向任何一路。
139fn merge_candidates(text: &[(RecordKey, f64)], vector: &[(RecordKey, f64)]) -> Vec<(RecordKey, f64)> {
140    let mut seen = HashSet::new();
141    let mut merged = Vec::with_capacity(text.len() + vector.len());
142    let mut index = 0;
143    while index < text.len() || index < vector.len() {
144        if let Some(item) = text.get(index) { if seen.insert(item.0) { merged.push(*item); } }
145        if let Some(item) = vector.get(index) { if seen.insert(item.0) { merged.push(*item); } }
146        index += 1;
147    }
148    merged
149}
150
151/// 宿主注册重排回调时一并声明的预算。条数上限与 token 上限是**与门**,任一到顶就停:
152/// 几十字符的短候选,token 预算能装几百条,而重排的成本由批次条数主导,所以条数必须单独封顶;
153/// 长正文则由 token 封顶(模型吃不下更多)。两者超出的候选不是变慢就是直接报错。
154#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
155pub struct RerankerOptions {
156    /// 送进重排的总预算:查询词 + 所有候选文档,从前往后累加到超额为止。
157    #[serde(default = "default_max_tokens_total")] pub max_tokens_total: usize,
158    /// 送进重排的候选条数上限。
159    #[serde(default = "default_max_candidates")] pub max_candidates: usize,
160    #[serde(default = "default_max_tokens_per_doc")] pub max_tokens_per_doc: usize,
161    /// 查询词的 token 预算,`None` 表示不截断。
162    #[serde(default)] pub max_tokens_query: Option<usize>,
163}
164impl Default for RerankerOptions {
165    fn default() -> Self { Self { max_tokens_total: default_max_tokens_total(), max_candidates: default_max_candidates(), max_tokens_per_doc: default_max_tokens_per_doc(), max_tokens_query: None } }
166}
167
168pub(crate) struct RerankerEntry { pub options: RerankerOptions, pub reranker: Box<dyn Reranker> }
169
170/// 进程内的重排回调位。`Arc` 让检索线程拿出去后立刻放掉注册表锁。
171#[derive(Default)]
172pub(crate) struct RerankerRegistry { entry: Mutex<Option<Arc<Mutex<RerankerEntry>>>> }
173
174impl RerankerRegistry {
175    pub fn new() -> Self { Self::default() }
176    pub fn is_registered(&self) -> bool { self.entry.lock().is_some() }
177    pub fn get(&self) -> Option<Arc<Mutex<RerankerEntry>>> { self.entry.lock().clone() }
178    pub fn register(&self, entry: RerankerEntry) { *self.entry.lock() = Some(Arc::new(Mutex::new(entry))); }
179    pub fn remove(&self) -> bool { self.entry.lock().take().is_some() }
180}
181
182/// 注册校验用的样本:真实调用一次,条数与有限性不符即拒绝绑定。
183const RERANK_SAMPLE_DOCS: [&str; 2] = ["重排校验样本一", "rerank probe two"];
184
185impl KnowledgeBase {
186    /// 注册重排回调。调用前按声明的定长约束强制截断,宿主不自己重排。
187    pub fn register_reranker<F: Reranker + 'static>(&self, reranker: F) -> Result<()> {
188        self.register_reranker_with(reranker, RerankerOptions::default())
189    }
190
191    pub fn register_reranker_with<F: Reranker + 'static>(&self, reranker: F, options: RerankerOptions) -> Result<()> {
192        if options.max_tokens_total == 0 { return Err(Error::Validation("max_tokens_total must be at least 1".into())); }
193        if options.max_candidates == 0 { return Err(Error::Validation("max_candidates must be at least 1".into())); }
194        if options.max_tokens_per_doc == 0 { return Err(Error::Validation("max_tokens_per_doc must be at least 1".into())); }
195        if options.max_tokens_query == Some(0) { return Err(Error::Validation("max_tokens_query must be positive".into())); }
196        let mut entry = RerankerEntry { options, reranker: Box::new(reranker) };
197        let documents: Vec<String> = RERANK_SAMPLE_DOCS.iter().map(|sample| (*sample).to_string()).collect();
198        // 样本返回 Ok 就校验条数与有限性;回调当场不可用时无从校验形状,允许绑定,
199        // 可用性留到检索时降级并写进诊断——这与「重排挂了不报错」的降级口径一致。
200        if let Ok(produced) = entry.reranker.rerank("校验样本", &documents) {
201            if produced.len() != documents.len() {
202                return Err(Error::Validation(format!("reranker returned {} scores for {} documents", produced.len(), documents.len())));
203            }
204            if produced.iter().any(|score| !score.is_finite()) { return Err(Error::Validation("reranker scores must be finite".into())); }
205        }
206        self.engine.rerankers.register(entry);
207        Ok(())
208    }
209
210    pub fn unregister_reranker(&self) -> bool { self.engine.rerankers.remove() }
211
212    pub fn reranker_registered(&self) -> bool { self.engine.rerankers.is_registered() }
213
214    /// 全文路。派生索引查询失败时返回 `Index`,由调用方触发重建后重试。
215    fn search_text(&self, conn: &rusqlite::Connection, query: &str, filter: &ReadFilter, kinds: &[RecordKind], limit: usize, field: MatchField) -> Result<Vec<(RecordKey, f64)>> {
216        self.sync_index_if_behind(conn)?;
217        let Some(index_filter) = index_filter(conn, filter, kinds)? else { return Ok(Vec::new()) };
218        // 领域登记了谓词等价词时先扩散:把同义写法一并纳入召回(如「beta」补「alpha」)。
219        let expanded = crate::graph::match_predicate_synonyms(conn, &filter.namespace, query)?;
220        self.index()?.search_in(&expanded, &index_filter, limit, field)
221    }
222
223    pub fn search(&self, request: &SearchRequest) -> Result<SearchResult> {
224        storage::validate_filter(&request.filter)?;
225        storage::validate_limit(request.limit)?;
226        let query = request.query.trim();
227        if query.is_empty() { return Err(Error::Validation("a text query is required".into())); }
228        // 向量路要同时满足「开关打开」与「给出目标空间」;缺任一条件就不走向量路,不是错误。
229        let vector_path = request.vector && request.embed_space.is_some();
230        if !request.text && !vector_path { return Err(Error::Validation("enable at least one of text or vector".into())); }
231        if !request.text_weight.is_finite() || request.text_weight <= 0.0 { return Err(Error::Validation("text_weight must be finite and positive".into())); }
232        let limit = request.candidate_limit.unwrap_or((request.limit * 5).max(1_000).min(10_000));
233        storage::validate_limit(limit)?;
234        if limit < request.limit { return Err(Error::Validation("candidate_limit must be at least limit".into())); }
235        // 事件只在宿主注册了 sink 时才产出:没注册就全程不构造、不格式化。
236        let sink = self.engine.events.get();
237        let started = sink.is_some().then(|| std::time::Instant::now());
238        let mut stages = sink.is_some().then(crate::events::StageTimer::start);
239        // Graph-First 剪枝:图邻域的候选 id 在取库锁之前算好(build_graph 自己会先取一次锁)。
240        let allowed = match &request.prune {
241            Some(prune) => {
242                if prune.depth == 0 { return Err(Error::Validation("graph prune depth must be at least 1".into())); }
243                storage::validate_limit(prune.limit)?;
244                let mut ids: HashSet<i64> = self.graph().build_graph(&request.filter)?.ego_ids(prune.root, prune.depth, prune.limit).into_iter().collect();
245                ids.insert(prune.root);
246                Some(ids)
247            }
248            None => None,
249        };
250        mark(&mut stages, "prepare");
251        let mut diagnostics = SearchDiagnostics::default();
252        // 向量路的第一步是「库自己把查询词嵌入」——宿主只给词,不给向量。
253        // 这一步是模型往返,必须在取库锁之前做完。
254        let mut embedded_query: Option<(embeddings::EmbeddingSpace, Vec<f32>)> = None;
255        // 向量路的候选限在「该领域实际启用了向量化的类型」上:请求的 kinds 与它取交集。
256        // 交集为空就整条向量路不走——空的 kinds 在打分侧表示「不过滤」,
257        // 直接传下去会把已关闭类型的存量向量也捞回来。
258        let mut vector_kinds: Vec<RecordKind> = Vec::new();
259        if let Some(space_id) = request.embed_space.as_deref().filter(|_| request.vector) {
260            let gated = { let state = self.read()?; embeddings::namespace_vectorization(state.conn(), &request.filter.namespace)? };
261            if !gated {
262                diagnostics.degraded.push(Degrade::NamespaceDisabled);
263            } else {
264                // 只让「已启用、且这一档自己已经补齐」的类型走向量:
265                // 某一档没补完,就把它从向量路里剔除,它的存量向量先不参与打分。
266                // 半个领域的向量参会比只用全文更糟——排名会偏向「先补完的那部分」。
267                let (enabled, ready) = {
268                    let state = self.read()?;
269                    let conn = state.conn();
270                    let namespace = &request.filter.namespace;
271                    (embeddings::enabled_kinds(conn, namespace)?, embeddings::ready_kinds(conn, namespace, space_id)?)
272                };
273                let requested: Vec<RecordKind> = request.kinds.iter().copied().filter(|kind| enabled.contains(kind)).collect();
274                vector_kinds = requested.iter().copied().filter(|kind| ready.contains(kind)).collect();
275                if !requested.is_empty() {
276                    // 空间没登记过是配置错误,直接报;登记过但没绑回调才是可降级的情形。
277                    let space = { let state = self.read()?; embeddings::get_space(state.conn(), space_id)? };
278                    if vector_kinds.is_empty() {
279                        // 请求要的类型都启用了,但没有一档补齐:这一轮只给全文。
280                        diagnostics.degraded.push(Degrade::VectorNotReady);
281                    } else {
282                        match self.engine.embedders.get(space_id) {
283                            None => diagnostics.degraded.push(Degrade::NoEmbedder),
284                            Some(entry) => {
285                                let produced = { let mut guard = entry.lock(); guard.embed(&[query.to_string()]) };
286                                match produced {
287                                    Ok(mut values) if values.len() == 1 => embedded_query = Some((space, values.remove(0))),
288                                    _ => diagnostics.degraded.push(Degrade::EmbedFailed),
289                                }
290                            }
291                        }
292                    }
293                }
294            }
295        }
296        if vector_path { mark(&mut stages, "embed"); }
297        let state = self.read()?;
298        let conn = state.conn();
299        let mut text_rank: Vec<(RecordKey, f64)> = Vec::new();
300        if request.text {
301            let text_hits = match self.search_text(conn, query, &request.filter, &request.kinds, limit, request.match_field) {
302                Ok(hits) => Some(hits),
303                // 派生索引查询失败:当场从权威数据重建一次再重试。恢复得了就照常给结果;
304                // 连重建都失败才隔离文本路。不静默退回空文本——那等于把 BM25 地板也丢掉。
305                Err(Error::Index(_)) => match self.rebuild_indexes() {
306                    Ok(_) => match self.search_text(conn, query, &request.filter, &request.kinds, limit, request.match_field) {
307                        Ok(hits) => Some(hits),
308                        Err(_) => { diagnostics.degraded.push(Degrade::TextIndexUnavailable); Some(Vec::new()) }
309                    },
310                    Err(_) => { diagnostics.degraded.push(Degrade::TextIndexUnavailable); Some(Vec::new()) }
311                },
312                // 库已关闭是调用错误,不该被降级吞掉。
313                Err(error) => return Err(error),
314            };
315            text_rank = text_hits.unwrap_or_default();
316            diagnostics.text_used = true;
317        }
318        mark(&mut stages, "text");
319        let mut vector_rank: Vec<(RecordKey, f64)> = Vec::new();
320        let mut vector_space_id: Option<String> = None;
321        if let Some((space, vector)) = &embedded_query {
322            // 分区键与载入查询同源:都用归一化后的 namespace / scope,避免大小写差异导致重复分区。
323            let namespace = text::normalized_tag(&request.filter.namespace);
324            let scopes: Vec<String> = request.filter.scopes.iter().map(|s| text::normalized_tag(s)).collect();
325            let tags: Vec<String> = request.filter.tags.iter().map(|t| text::normalized_tag(t)).collect();
326            let mut scored: Vec<(RecordKey, f64)> = Vec::new();
327            for scope in scopes {
328                let Some(partition) = self.partition(conn, space, &namespace, &scope)? else { continue };
329                scored.extend(partition.search(vector, &vector_kinds, &tags, &request.filter.note_ids, limit, allowed.as_ref())?);
330            }
331            // 跨分区汇总后再统一排名:分区各自从 0 计 rank 会破坏 RRF 融合语义。
332            scored.sort_by(|a,b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
333            scored.truncate(limit);
334            diagnostics.vector_used = true;
335            vector_space_id = Some(space.id.clone());
336            vector_rank = scored;
337        }
338        if diagnostics.vector_used { mark(&mut stages, "vector"); }
339        // 两路各自的分数与 RRF 融合分:融合分只用来充当 `score` 字段与「没法重排」时的兜底排序,
340        // 重排可用时它不参与候选取舍,也不参与名次决定。
341        let mut scores: BTreeMap<RecordKey, (f64, Option<f64>, BTreeMap<String,f64>)> = BTreeMap::new();
342        for (rank, (key, score)) in text_rank.iter().enumerate() {
343            let hit = scores.entry(*key).or_default();
344            hit.0 += request.text_weight / (60.0 + (rank + 1) as f64); hit.1 = Some(*score);
345        }
346        if let Some(space_id) = &vector_space_id {
347            for (rank, (key, score)) in vector_rank.iter().enumerate() {
348                let hit = scores.entry(*key).or_default();
349                hit.0 += 1.0 / (60.0 + (rank + 1) as f64); hit.2.insert(space_id.clone(), *score);
350            }
351        }
352        let total = if request.with_total { Some(storage::count_matches(conn, &request.filter, &request.kinds)?) } else { None };
353        let mut rerank_scores: BTreeMap<RecordKey, f64> = BTreeMap::new();
354        // 候选顺序:重排可用时按两路交替合并(融合分不参与这一步),否则按 RRF 融合分。
355        // 折叠排在这之后、送重排之前——同一篇只留一个代表片,重排预算不重复花在同一篇上。
356        let rerank_entry = if request.rerank { self.engine.rerankers.get() } else { None };
357        let use_rerank = rerank_entry.is_some();
358        let ordered: Vec<RecordKey> = if use_rerank {
359            merge_candidates(&text_rank, &vector_rank).into_iter().map(|(key, _)| key).collect()
360        } else {
361            let mut by_score: Vec<(RecordKey, f64)> = scores.iter().map(|(key, hit)| (*key, hit.0)).collect();
362            by_score.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
363            by_score.into_iter().map(|(key, _)| key).collect()
364        };
365        let candidate_count = ordered.len();
366        mark(&mut stages, "fuse");
367        // 同一篇笔记的多个命中切片折叠成一条:只保留最靠前的那一片,
368        // 否则一篇对话体文档会用自己几十个片段占满整个结果列表,把别的文档全挤出去。
369        // 折叠时顺手把同篇在窗口内的其余命中片段收进 `top_chunks`(第 0 条是代表片自身)。
370        // 聚合与计数只在按正文列(或默认的正文+名字列)检索时做:名字与目录列是笔记级的,
371        // 一篇至多一条命中,聚无可聚、也数无可数。
372        let aggregate = request.text && matches!(request.match_field, MatchField::All | MatchField::Text);
373        let per_note = if aggregate { request.top_chunks_per_note } else { 0 };
374        let ordered_ids: Vec<i64> = ordered.iter().map(|key| key.id).collect();
375        let chunk_notes = storage::chunk_notes(conn, &ordered_ids)?;
376        let mut folded: Vec<RecordKey> = Vec::new();
377        let mut top_chunks: HashMap<i64, Vec<ChunkRef>> = HashMap::new();
378        if chunk_notes.is_empty() {
379            folded = ordered;
380        } else {
381            let mut seen_notes: HashSet<i64> = HashSet::new();
382            for key in ordered {
383                let Some(&(note_id, offset)) = chunk_notes.get(&key.id) else { folded.push(key); continue };
384                if seen_notes.insert(note_id) {
385                    if per_note > 0 { top_chunks.insert(note_id, vec![ChunkRef { id: key.id, offset }]); }
386                    folded.push(key);
387                } else if per_note > 0 {
388                    let list = top_chunks.entry(note_id).or_default();
389                    if list.len() < per_note { list.push(ChunkRef { id: key.id, offset }); }
390                }
391            }
392        }
393        let folded_count = folded.len();
394        mark(&mut stages, "fold");
395        // 候选取舍:重排可用时从前往后取,条数上限与 token 总预算任一先到顶就停——短候选靠
396        // 条数封顶(批次条数才是重排的成本),长正文靠 token 封顶。没有重排时直接取前 limit 条。
397        let mut selected: Vec<RecordKey> = Vec::new();
398        let mut reranked = false;
399        let mut rerank_docs = 0usize;
400        let mut rerank_tokens = 0usize;
401        if let Some(entry) = rerank_entry {
402            let options = entry.lock().options;
403            let budgeted_query = options.max_tokens_query.map(|budget| text::truncate_to_tokens(query, budget)).unwrap_or_else(|| query.to_string());
404            let ids: Vec<i64> = folded.iter().map(|key| key.id).collect();
405            let bodies = match self.index() { Ok(index) => index.bodies(&ids)?, Err(_) => BTreeMap::new() };
406            let names = storage::entity_names(conn, &ids).unwrap_or_default();
407            // 候选正文取自索引的 stored 字段(切片正文在这里)。实体的正文列不含规范名,
408            // 这里给实体把规范名拼在正文前——否则纯名实体送到重排的文档是空的,等于没内容可判。
409            let mut used = text::count_tokens(&budgeted_query);
410            let mut candidates: Vec<RecordKey> = Vec::new();
411            let mut documents: Vec<String> = Vec::new();
412            for key in &folded {
413                // 与门:条数先到顶就停,不必再为这一条算 token。
414                if candidates.len() >= options.max_candidates { break; }
415                let body = bodies.get(&key.id).map(String::as_str).unwrap_or("");
416                let full = match names.get(&key.id).filter(|name| !name.is_empty()) {
417                    Some(name) => format!("{name} {body}"),
418                    None => body.to_string(),
419                };
420                let document = text::truncate_to_tokens(&full, options.max_tokens_per_doc);
421                let cost = text::count_tokens(&document);
422                if used + cost > options.max_tokens_total { break; }
423                used += cost;
424                candidates.push(*key);
425                documents.push(document);
426            }
427            diagnostics.rerank_candidates = candidates.len();
428            diagnostics.rerank_truncated = folded.len().saturating_sub(candidates.len());
429            rerank_docs = candidates.len();
430            rerank_tokens = used;
431            // 没有候选就不必调模型:多数重排服务把空文档列表当无效请求,会误标降级。
432            let produced = if candidates.is_empty() { Ok(Vec::new()) }
433            else { let mut guard = entry.lock(); guard.reranker.rerank(&budgeted_query, &documents) };
434            match produced {
435                Ok(values) if values.len() == candidates.len() && values.iter().all(|value| value.is_finite()) => {
436                    let mut pairs: Vec<(RecordKey, f32)> = candidates.into_iter().zip(values).collect();
437                    pairs.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
438                    for (key, score) in pairs { rerank_scores.insert(key, f64::from(score)); selected.push(key); }
439                    diagnostics.reranked = true;
440                    reranked = true;
441                }
442                // 重排产出不符:退回候选顺序,不把整次检索判死。
443                _ => diagnostics.degraded.push(Degrade::RerankFailed),
444            }
445        }
446        if !reranked { selected = folded; }
447        if use_rerank { mark(&mut stages, "rerank"); }
448        selected.truncate(request.limit);
449        // 「这一篇命中多少片」只对最终返回的条目算:它是结果上的字段,不是候选窗口的属性。
450        // 折叠后的候选常有几百篇,而返回只有 limit 条——按候选窗口算等于白算几十倍。
451        // 计数走一次遍历、按笔记分桶,不为每篇各发一次查询。
452        let note_counts: BTreeMap<i64, usize> = if aggregate {
453            match index_filter(conn, &request.filter, &request.kinds)? {
454                Some(ifilter) => {
455                    let mut targets: Vec<i64> = selected.iter().filter_map(|key| chunk_notes.get(&key.id).map(|(note_id, _)| *note_id)).collect();
456                    targets.sort_unstable();
457                    targets.dedup();
458                    self.index()?.count_in_many(query, &ifilter, request.match_field, &targets)?.into_iter().collect()
459                }
460                None => BTreeMap::new(),
461            }
462        } else { BTreeMap::new() };
463        mark(&mut stages, "count");
464        // 一次批量取回全部命中本体:`load_many` 的语义等同于逐条 `get`(同样的过滤、同样的装配),
465        // 但把每条命中的两轮 SQL(matches_filter + record_value)压成固定三条。
466        // 用 id 做键再按 `selected` 顺序装配,批量取回的顺序不会影响名次。
467        let ids: Vec<i64> = selected.iter().map(|key| key.id).collect();
468        let mut records: BTreeMap<i64, Value> = storage::load_many(conn, &ids, &request.filter)?;
469        let mut hits = Vec::new();
470        for key in selected {
471            // 索引里还有文档、库里记录已不在(并发删除、或索引尚未追上)时跳过这一条:
472            // 为一条已消失的记录让别人的整次检索失败,代价不对等。
473            let Some(record) = records.remove(&key.id) else { continue };
474            let note_id = chunk_notes.get(&key.id).map(|(note_id, _)| *note_id);
475            let note_chunks = note_id.and_then(|id| note_counts.get(&id).copied());
476            let top_chunks = note_id.and_then(|id| top_chunks.remove(&id)).unwrap_or_default();
477            let (score, text_score, vector_scores) = scores.remove(&key).unwrap_or((0.0, None, BTreeMap::new()));
478            hits.push(SearchHit { record, key, score, text_score, vector_scores, rerank_score: rerank_scores.get(&key).copied(), note_chunks, top_chunks });
479        }
480        for degrade in &diagnostics.degraded { self.note_degrade(*degrade); }
481        if let (Some(sink), Some(stages)) = (sink, stages) {
482            let mut event = crate::events::LogEvent::new("search");
483            event.ms = started.map(|started| started.elapsed().as_millis() as u64).unwrap_or(0);
484            event.stages = stages.finish("load");
485            event.candidates = Some(candidate_count);
486            event.folded = Some(folded_count);
487            event.rerank_docs = Some(rerank_docs);
488            event.rerank_tokens = Some(rerank_tokens);
489            event.hits = Some(hits.len());
490            event.degraded = diagnostics.degraded.clone();
491            // 日志回调不该有能力打断检索:宿主 sink 里 panic 只丢这一条事件。
492            let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| sink(&event)));
493        }
494        Ok(SearchResult { hits, revision: storage::current_revision(conn)?, indexed_revision: storage::meta(conn, "indexed_revision")?, total, diagnostics })
495    }
496
497    /// 先按 `request` 检索,再给每条命中挂上它在图上所在实体的邻域(实体 + 关系)。
498    ///
499    /// 「挂载」靠标签文本:命中的记录所带的 tag,与同一 namespace / scope 内某实体的名字或别名
500    /// 文本相同,就认为该记录挂在这个实体上(两者都落在 `strings` 表,同一套归一化)。
501    /// `limit` 是每个实体邻域的规模上限。命中记录若没有命中任何实体,返回空邻域。
502    pub fn search_with_context(&self, request: &SearchRequest, limit: usize) -> Result<Vec<ContextualHit>> {
503        storage::validate_limit(limit)?;
504        // 只用 namespace / scope 定领域边界;tag 是「挂在哪个实体」的线索,不能反过来筛实体。
505        let scope = ReadFilter { tags: vec![], ..request.filter.clone() };
506        let hits = self.search(request)?.hits;
507        let keys: Vec<i64> = hits.iter().map(|hit| hit.key.id).collect();
508        let mut contexts = self.entity_contexts(&keys, &scope, limit)?;
509        Ok(hits.into_iter().map(|hit| ContextualHit {
510            context: contexts.remove(&hit.key.id).unwrap_or_else(|| Neighborhood { entities: vec![], relations: vec![] }),
511            hit,
512        }).collect())
513    }
514
515    /// 批量解析命中记录的实体邻域,一次取锁、把原来每条命中各自的 N+1 查询压成固定 4 条批量查询。
516    /// 结果与逐条展开一致:记录 → tags 文本匹配同领域实体 → 每实体 1 跳关系(各自独立按 `limit` 截断)→ 汇总去重后取 `limit`。
517    fn entity_contexts(&self, record_ids: &[i64], filter: &ReadFilter, limit: usize) -> Result<BTreeMap<i64, Neighborhood>> {
518        let empty = || Neighborhood { entities: vec![], relations: vec![] };
519        let mut out: BTreeMap<i64, Neighborhood> = record_ids.iter().map(|id| (*id, empty())).collect();
520        if record_ids.is_empty() { return Ok(out); }
521        let state = self.read()?;
522        let conn = state.conn();
523        // 1. 记录 → 实体:所有命中记录一次查完。
524        let (entity_condition, entity_values) = storage::filter_sql(filter, &[RecordKind::Entity], false)?;
525        let record_placeholders = vec!["?"; record_ids.len()].join(",");
526        let mut stmt = conn.prepare(&format!(
527            "SELECT DISTINCT rt.record_id, ea.entity_id FROM record_tags rt \
528             JOIN entity_aliases ea ON ea.alias_id=rt.tag_id \
529             JOIN records r ON r.id=ea.entity_id \
530             WHERE rt.record_id IN ({record_placeholders}) AND {entity_condition} ORDER BY rt.record_id, ea.entity_id"
531        ))?;
532        let params = record_ids.iter().map(|id| SqlValue::Integer(*id)).chain(entity_values).collect::<Vec<_>>();
533        let mut seeds: BTreeMap<i64, BTreeSet<i64>> = BTreeMap::new();
534        for row in stmt.query_map(params_from_iter(params), |r| Ok((r.get::<_, i64>(0)?, r.get::<_, i64>(1)?)))? {
535            let (record_id, entity_id) = row?;
536            seeds.entry(record_id).or_default().insert(entity_id);
537        }
538        let roots: Vec<i64> = seeds.values().flatten().copied().collect::<BTreeSet<_>>().into_iter().collect();
539        if roots.is_empty() { return Ok(out); }
540        // 2. 一次查出这批实体上的全部 1 跳关系记录。关系记录按 `filter` 过滤,端点按去掉 tags 的
541        // `entity_filter` 过滤——与逐条版保持一致(tags 约束关系,scope 约束端点)。
542        let (relation_condition, relation_values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
543        let root_placeholders = vec!["?"; roots.len()].join(",");
544        let mut stmt = conn.prepare(&format!(
545            "SELECT rl.record_id, rl.subject_id, rl.object_id FROM relations rl JOIN records r ON r.id=rl.record_id \
546             WHERE (rl.subject_id IN ({root_placeholders}) OR rl.object_id IN ({root_placeholders})) AND {relation_condition} \
547             ORDER BY rl.record_id"
548        ))?;
549        let params = roots.iter().map(|id| SqlValue::Integer(*id)).chain(roots.iter().map(|id| SqlValue::Integer(*id))).chain(relation_values).collect::<Vec<_>>();
550        let mut edges: Vec<(i64, i64, i64)> = Vec::new();
551        for row in stmt.query_map(params_from_iter(params), |r| Ok((r.get::<_, i64>(0)?, r.get::<_, i64>(1)?, r.get::<_, i64>(2)?)))? {
552            edges.push(row?);
553        }
554        // 3. 批量取种子实体、端点实体和关系记录。
555        let entity_filter = ReadFilter { tags: vec![], ..filter.clone() };
556        let root_entities: BTreeMap<i64, Entity> = storage::load_many(conn, &roots, filter)?;
557        let endpoint_ids: Vec<i64> = edges.iter().flat_map(|(_, subject, object)| [*subject, *object]).collect::<BTreeSet<_>>().into_iter().collect();
558        let endpoint_entities: BTreeMap<i64, Entity> = storage::load_many(conn, &endpoint_ids, &entity_filter)?;
559        let relation_ids: Vec<i64> = edges.iter().map(|(id, _, _)| *id).collect();
560        let relation_records: BTreeMap<i64, Relation> = storage::load_many(conn, &relation_ids, filter)?;
561        // 4. 在内存里按记录分发,并复现「每个实体各自按 limit 截断」的语义。
562        let mut incident: BTreeMap<i64, Vec<usize>> = BTreeMap::new();
563        for (i, &(_, subject, object)) in edges.iter().enumerate() {
564            incident.entry(subject).or_default().push(i);
565            if object != subject { incident.entry(object).or_default().push(i); }
566        }
567        for (record_id, root_ids) in &seeds {
568            let mut entities: BTreeMap<i64, Entity> = BTreeMap::new();
569            for id in root_ids { if let Some(entity) = root_entities.get(id) { entities.insert(*id, entity.clone()); } }
570            let mut relations: BTreeMap<i64, Relation> = BTreeMap::new();
571            for root in root_ids {
572                let mut count = 0usize;
573                for &i in incident.get(root).map(Vec::as_slice).unwrap_or(&[]) {
574                    let (relation_id, subject, object) = edges[i];
575                    let Some(relation) = relation_records.get(&relation_id) else { continue };
576                    let endpoint = if subject == *root { object } else { subject };
577                    let Some(entity) = endpoint_entities.get(&endpoint) else { continue };
578                    entities.entry(endpoint).or_insert_with(|| entity.clone());
579                    relations.entry(relation_id).or_insert_with(|| relation.clone());
580                    count += 1;
581                    if count == limit { break; }
582                }
583            }
584            out.insert(*record_id, Neighborhood { entities: entities.into_values().collect(), relations: relations.into_values().take(limit).collect() });
585        }
586        Ok(out)
587    }
588}
589
590#[cfg(test)]
591mod tests {
592    use super::*;
593    use crate::MemoryInput;
594    use std::sync::atomic::Ordering;
595
596    /// 索引查询失败时,检索应触发一次重建并重试:恢复得了就照常出结果,
597    /// 连重建都失败才隔离文本路——不再静默退回空文本。
598    /// token 计数按字符密度估算:ASCII ≈ 4 字符/token,非 ASCII ≈ 2 字符/token,向上取整。
599    #[test]
600    fn token_count_follows_character_density() {
601        assert_eq!(text::count_tokens(""), 0);
602        assert_eq!(text::count_tokens("abcd"), 1);
603        assert_eq!(text::count_tokens("abcde"), 2);
604        assert_eq!(text::count_tokens("中"), 1);
605        assert_eq!(text::count_tokens("中国"), 1);
606        assert_eq!(text::count_tokens("中国人"), 2);
607        // 截断与计数同源:截断后的文本不超预算。
608        assert_eq!(text::truncate_to_tokens("abcd", 1), "abcd");
609        assert_eq!(text::truncate_to_tokens("abcde", 1), "abcd");
610        assert_eq!(text::truncate_to_tokens("中国人", 1), "中国");
611        assert!(text::count_tokens(&text::truncate_to_tokens("中国人", 1)) <= 1);
612    }
613
614    #[test]
615    fn index_query_failure_rebuilds_then_recovers() {
616        let dir = tempfile::tempdir().unwrap();
617        let kb = KnowledgeBase::open(dir.path()).unwrap();
618        kb.memories().upsert(MemoryInput::new("索引故障恢复的独有措辞")).unwrap();
619        let index = kb.index().unwrap();
620        let request = SearchRequest {
621            query: "索引故障恢复的独有措辞".into(), kinds: vec![RecordKind::Memory],
622            vector: false, rerank: false, ..Default::default()
623        };
624
625        index.fail_search.store(true, Ordering::SeqCst);
626        let degraded = kb.search(&request).unwrap();
627        assert!(index.rebuilds.load(Ordering::SeqCst) >= 1, "索引查询失败必须触发重建");
628        assert!(degraded.diagnostics.degraded.contains(&Degrade::TextIndexUnavailable),
629            "重建之后仍失败,才隔离文本路");
630
631        index.fail_search.store(false, Ordering::SeqCst);
632        let recovered = kb.search(&request).unwrap();
633        assert_eq!(recovered.hits.len(), 1, "故障排除后索引可用,照常命中");
634        assert!(!recovered.diagnostics.degraded.contains(&Degrade::TextIndexUnavailable));
635    }
636}