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 }
10fn 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
17fn 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
22fn 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#[derive(Debug, Clone, Serialize, Deserialize)]
40pub struct GraphPrune {
41 pub root: i64,
43 pub depth: usize,
45 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 pub embed_space: Option<String>,
56 pub text_weight: f64, pub prune: Option<GraphPrune>,
57 #[serde(default = "yes")] pub text: bool,
59 #[serde(default = "yes")] pub vector: bool,
61 #[serde(default = "yes")] pub rerank: bool,
63 #[serde(default)] pub with_total: bool,
65 #[serde(default)] pub match_field: MatchField,
67 #[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#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
83pub struct ChunkRef {
84 pub id: i64,
86 pub offset: usize,
88}
89#[derive(Debug, Clone, Serialize, Deserialize)]
90pub struct SearchHit {
91 pub key: RecordKey,
92 pub score: f64, pub text_score: Option<f64>, pub vector_scores: BTreeMap<String, f64>,
94 #[serde(default)] pub rerank_score: Option<f64>,
96 #[serde(default)] pub note_chunks: Option<usize>,
100 #[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 #[serde(default)] pub total: Option<usize>,
113 #[serde(default)] pub diagnostics: SearchDiagnostics,
114}
115
116#[derive(Debug, Clone, Serialize, Deserialize)]
118pub struct ContextualHit { pub hit: SearchHit, pub context: Neighborhood }
119
120pub 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
137fn 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
155pub struct RerankerOptions {
156 #[serde(default = "default_max_tokens_total")] pub max_tokens_total: usize,
158 #[serde(default = "default_max_candidates")] pub max_candidates: usize,
160 #[serde(default = "default_max_tokens_per_doc")] pub max_tokens_per_doc: usize,
161 #[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#[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
182const RERANK_SAMPLE_DOCS: [&str; 2] = ["重排校验样本一", "rerank probe two"];
184
185impl KnowledgeBase {
186 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 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 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 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 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 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 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 let mut embedded_query: Option<(embeddings::EmbeddingSpace, Vec<f32>)> = None;
255 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 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 let space = { let state = self.read()?; embeddings::get_space(state.conn(), space_id)? };
278 if vector_kinds.is_empty() {
279 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 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 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 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 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 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 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 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 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 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 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 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 _ => 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 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 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 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 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 pub fn search_with_context(&self, request: &SearchRequest, limit: usize) -> Result<Vec<ContextualHit>> {
503 storage::validate_limit(limit)?;
504 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 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 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 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 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 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 #[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 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}