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)>> {
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(_)) => { diagnostics.degraded.push(Degrade::TextIndexUnavailable); Some(Vec::new()) }
306 Err(error) => return Err(error),
308 };
309 text_rank = text_hits.unwrap_or_default();
310 diagnostics.text_used = true;
311 }
312 mark(&mut stages, "text");
313 let mut vector_rank: Vec<(RecordKey, f64)> = Vec::new();
314 let mut vector_space_id: Option<String> = None;
315 if let Some((space, vector)) = &embedded_query {
316 let namespace = text::normalized_tag(&request.filter.namespace);
318 let scopes: Vec<String> = request.filter.scopes.iter().map(|s| text::normalized_tag(s)).collect();
319 let tags: Vec<String> = request.filter.tags.iter().map(|t| text::normalized_tag(t)).collect();
320 let mut scored: Vec<(RecordKey, f64)> = Vec::new();
321 for scope in scopes {
322 let Some(partition) = self.partition(space, &namespace, &scope)? else { continue };
323 scored.extend(partition.search(vector, &vector_kinds, &tags, &request.filter.note_ids, limit, allowed.as_ref())?);
324 }
325 scored.sort_by(|a,b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
327 scored.truncate(limit);
328 diagnostics.vector_used = true;
329 vector_space_id = Some(space.id.clone());
330 vector_rank = scored;
331 }
332 if diagnostics.vector_used { mark(&mut stages, "vector"); }
333 let mut scores: BTreeMap<RecordKey, (f64, Option<f64>, BTreeMap<String,f64>)> = BTreeMap::new();
336 for (rank, (key, score)) in text_rank.iter().enumerate() {
337 let hit = scores.entry(*key).or_default();
338 hit.0 += request.text_weight / (60.0 + (rank + 1) as f64); hit.1 = Some(*score);
339 }
340 if let Some(space_id) = &vector_space_id {
341 for (rank, (key, score)) in vector_rank.iter().enumerate() {
342 let hit = scores.entry(*key).or_default();
343 hit.0 += 1.0 / (60.0 + (rank + 1) as f64); hit.2.insert(space_id.clone(), *score);
344 }
345 }
346 let total = if request.with_total { Some(storage::count_matches(conn, &request.filter, &request.kinds)?) } else { None };
347 let mut rerank_scores: BTreeMap<RecordKey, f64> = BTreeMap::new();
348 let rerank_entry = if request.rerank { self.engine.rerankers.get() } else { None };
351 let use_rerank = rerank_entry.is_some();
352 let ordered: Vec<RecordKey> = if use_rerank {
353 merge_candidates(&text_rank, &vector_rank).into_iter().map(|(key, _)| key).collect()
354 } else {
355 let mut by_score: Vec<(RecordKey, f64)> = scores.iter().map(|(key, hit)| (*key, hit.0)).collect();
356 by_score.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
357 by_score.into_iter().map(|(key, _)| key).collect()
358 };
359 let candidate_count = ordered.len();
360 mark(&mut stages, "fuse");
361 let aggregate = request.text && matches!(request.match_field, MatchField::All | MatchField::Text);
367 let per_note = if aggregate { request.top_chunks_per_note } else { 0 };
368 let ordered_ids: Vec<i64> = ordered.iter().map(|key| key.id).collect();
369 let chunk_notes = storage::chunk_notes(conn, &ordered_ids)?;
370 let mut folded: Vec<RecordKey> = Vec::new();
371 let mut top_chunks: HashMap<i64, Vec<ChunkRef>> = HashMap::new();
372 if chunk_notes.is_empty() {
373 folded = ordered;
374 } else {
375 let mut seen_notes: HashSet<i64> = HashSet::new();
376 for key in ordered {
377 let Some(&(note_id, offset)) = chunk_notes.get(&key.id) else { folded.push(key); continue };
378 if seen_notes.insert(note_id) {
379 if per_note > 0 { top_chunks.insert(note_id, vec![ChunkRef { id: key.id, offset }]); }
380 folded.push(key);
381 } else if per_note > 0 {
382 let list = top_chunks.entry(note_id).or_default();
383 if list.len() < per_note { list.push(ChunkRef { id: key.id, offset }); }
384 }
385 }
386 }
387 let folded_count = folded.len();
388 mark(&mut stages, "fold");
389 let mut selected: Vec<RecordKey> = Vec::new();
392 let mut reranked = false;
393 let mut rerank_docs = 0usize;
394 let mut rerank_tokens = 0usize;
395 if let Some(entry) = rerank_entry {
396 let options = entry.lock().options;
397 let budgeted_query = options.max_tokens_query.map(|budget| text::truncate_to_tokens(query, budget)).unwrap_or_else(|| query.to_string());
398 let ids: Vec<i64> = folded.iter().map(|key| key.id).collect();
399 let bodies = match self.index() { Ok(index) => index.bodies(&ids)?, Err(_) => BTreeMap::new() };
400 let names = storage::entity_names(conn, &ids).unwrap_or_default();
401 let mut used = text::count_tokens(&budgeted_query);
404 let mut candidates: Vec<RecordKey> = Vec::new();
405 let mut documents: Vec<String> = Vec::new();
406 for key in &folded {
407 if candidates.len() >= options.max_candidates { break; }
409 let body = bodies.get(&key.id).map(String::as_str).unwrap_or("");
410 let full = match names.get(&key.id).filter(|name| !name.is_empty()) {
411 Some(name) => format!("{name} {body}"),
412 None => body.to_string(),
413 };
414 let document = text::truncate_to_tokens(&full, options.max_tokens_per_doc);
415 let cost = text::count_tokens(&document);
416 if used + cost > options.max_tokens_total { break; }
417 used += cost;
418 candidates.push(*key);
419 documents.push(document);
420 }
421 diagnostics.rerank_candidates = candidates.len();
422 diagnostics.rerank_truncated = folded.len().saturating_sub(candidates.len());
423 rerank_docs = candidates.len();
424 rerank_tokens = used;
425 let produced = if candidates.is_empty() { Ok(Vec::new()) }
427 else { let mut guard = entry.lock(); guard.reranker.rerank(&budgeted_query, &documents) };
428 match produced {
429 Ok(values) if values.len() == candidates.len() && values.iter().all(|value| value.is_finite()) => {
430 let mut pairs: Vec<(RecordKey, f32)> = candidates.into_iter().zip(values).collect();
431 pairs.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
432 for (key, score) in pairs { rerank_scores.insert(key, f64::from(score)); selected.push(key); }
433 diagnostics.reranked = true;
434 reranked = true;
435 }
436 _ => diagnostics.degraded.push(Degrade::RerankFailed),
438 }
439 }
440 if !reranked { selected = folded; }
441 if use_rerank { mark(&mut stages, "rerank"); }
442 selected.truncate(request.limit);
443 let note_counts: BTreeMap<i64, usize> = if aggregate {
447 match index_filter(conn, &request.filter, &request.kinds)? {
448 Some(ifilter) => {
449 let mut targets: Vec<i64> = selected.iter().filter_map(|key| chunk_notes.get(&key.id).map(|(note_id, _)| *note_id)).collect();
450 targets.sort_unstable();
451 targets.dedup();
452 self.index()?.count_in_many(query, &ifilter, request.match_field, &targets)?.into_iter().collect()
453 }
454 None => BTreeMap::new(),
455 }
456 } else { BTreeMap::new() };
457 mark(&mut stages, "count");
458 let ids: Vec<i64> = selected.iter().map(|key| key.id).collect();
462 let mut records: BTreeMap<i64, Value> = storage::load_many(conn, &ids, &request.filter)?;
463 let mut hits = Vec::new();
464 for key in selected {
465 let Some(record) = records.remove(&key.id) else { continue };
468 let note_id = chunk_notes.get(&key.id).map(|(note_id, _)| *note_id);
469 let note_chunks = note_id.and_then(|id| note_counts.get(&id).copied());
470 let top_chunks = note_id.and_then(|id| top_chunks.remove(&id)).unwrap_or_default();
471 let (score, text_score, vector_scores) = scores.remove(&key).unwrap_or((0.0, None, BTreeMap::new()));
472 hits.push(SearchHit { record, key, score, text_score, vector_scores, rerank_score: rerank_scores.get(&key).copied(), note_chunks, top_chunks });
473 }
474 for degrade in &diagnostics.degraded { self.note_degrade(*degrade); }
475 if let (Some(sink), Some(stages)) = (sink, stages) {
476 let mut event = crate::events::LogEvent::new("search");
477 event.ms = started.map(|started| started.elapsed().as_millis() as u64).unwrap_or(0);
478 event.stages = stages.finish("load");
479 event.candidates = Some(candidate_count);
480 event.folded = Some(folded_count);
481 event.rerank_docs = Some(rerank_docs);
482 event.rerank_tokens = Some(rerank_tokens);
483 event.hits = Some(hits.len());
484 event.degraded = diagnostics.degraded.clone();
485 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| sink(&event)));
487 }
488 Ok(SearchResult { hits, revision: storage::current_revision(conn)?, indexed_revision: storage::meta(conn, "indexed_revision")?, total, diagnostics })
489 }
490
491 pub fn search_with_context(&self, request: &SearchRequest, limit: usize) -> Result<Vec<ContextualHit>> {
497 storage::validate_limit(limit)?;
498 let scope = ReadFilter { tags: vec![], ..request.filter.clone() };
500 let hits = self.search(request)?.hits;
501 let keys: Vec<i64> = hits.iter().map(|hit| hit.key.id).collect();
502 let mut contexts = self.entity_contexts(&keys, &scope, limit)?;
503 Ok(hits.into_iter().map(|hit| ContextualHit {
504 context: contexts.remove(&hit.key.id).unwrap_or_else(|| Neighborhood { entities: vec![], relations: vec![] }),
505 hit,
506 }).collect())
507 }
508
509 fn entity_contexts(&self, record_ids: &[i64], filter: &ReadFilter, limit: usize) -> Result<BTreeMap<i64, Neighborhood>> {
512 let empty = || Neighborhood { entities: vec![], relations: vec![] };
513 let mut out: BTreeMap<i64, Neighborhood> = record_ids.iter().map(|id| (*id, empty())).collect();
514 if record_ids.is_empty() { return Ok(out); }
515 let state = self.read()?;
516 let conn = state.conn();
517 let (entity_condition, entity_values) = storage::filter_sql(filter, &[RecordKind::Entity], false)?;
519 let record_placeholders = vec!["?"; record_ids.len()].join(",");
520 let mut stmt = conn.prepare(&format!(
521 "SELECT DISTINCT rt.record_id, ea.entity_id FROM record_tags rt \
522 JOIN entity_aliases ea ON ea.alias_id=rt.tag_id \
523 JOIN records r ON r.id=ea.entity_id \
524 WHERE rt.record_id IN ({record_placeholders}) AND {entity_condition} ORDER BY rt.record_id, ea.entity_id"
525 ))?;
526 let params = record_ids.iter().map(|id| SqlValue::Integer(*id)).chain(entity_values).collect::<Vec<_>>();
527 let mut seeds: BTreeMap<i64, BTreeSet<i64>> = BTreeMap::new();
528 for row in stmt.query_map(params_from_iter(params), |r| Ok((r.get::<_, i64>(0)?, r.get::<_, i64>(1)?)))? {
529 let (record_id, entity_id) = row?;
530 seeds.entry(record_id).or_default().insert(entity_id);
531 }
532 let roots: Vec<i64> = seeds.values().flatten().copied().collect::<BTreeSet<_>>().into_iter().collect();
533 if roots.is_empty() { return Ok(out); }
534 let (relation_condition, relation_values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
537 let root_placeholders = vec!["?"; roots.len()].join(",");
538 let mut stmt = conn.prepare(&format!(
539 "SELECT rl.record_id, rl.subject_id, rl.object_id FROM relations rl JOIN records r ON r.id=rl.record_id \
540 WHERE (rl.subject_id IN ({root_placeholders}) OR rl.object_id IN ({root_placeholders})) AND {relation_condition} \
541 ORDER BY rl.record_id"
542 ))?;
543 let params = roots.iter().map(|id| SqlValue::Integer(*id)).chain(roots.iter().map(|id| SqlValue::Integer(*id))).chain(relation_values).collect::<Vec<_>>();
544 let mut edges: Vec<(i64, i64, i64)> = Vec::new();
545 for row in stmt.query_map(params_from_iter(params), |r| Ok((r.get::<_, i64>(0)?, r.get::<_, i64>(1)?, r.get::<_, i64>(2)?)))? {
546 edges.push(row?);
547 }
548 let entity_filter = ReadFilter { tags: vec![], ..filter.clone() };
550 let root_entities: BTreeMap<i64, Entity> = storage::load_many(conn, &roots, filter)?;
551 let endpoint_ids: Vec<i64> = edges.iter().flat_map(|(_, subject, object)| [*subject, *object]).collect::<BTreeSet<_>>().into_iter().collect();
552 let endpoint_entities: BTreeMap<i64, Entity> = storage::load_many(conn, &endpoint_ids, &entity_filter)?;
553 let relation_ids: Vec<i64> = edges.iter().map(|(id, _, _)| *id).collect();
554 let relation_records: BTreeMap<i64, Relation> = storage::load_many(conn, &relation_ids, filter)?;
555 let mut incident: BTreeMap<i64, Vec<usize>> = BTreeMap::new();
557 for (i, &(_, subject, object)) in edges.iter().enumerate() {
558 incident.entry(subject).or_default().push(i);
559 if object != subject { incident.entry(object).or_default().push(i); }
560 }
561 for (record_id, root_ids) in &seeds {
562 let mut entities: BTreeMap<i64, Entity> = BTreeMap::new();
563 for id in root_ids { if let Some(entity) = root_entities.get(id) { entities.insert(*id, entity.clone()); } }
564 let mut relations: BTreeMap<i64, Relation> = BTreeMap::new();
565 for root in root_ids {
566 let mut count = 0usize;
567 for &i in incident.get(root).map(Vec::as_slice).unwrap_or(&[]) {
568 let (relation_id, subject, object) = edges[i];
569 let Some(relation) = relation_records.get(&relation_id) else { continue };
570 let endpoint = if subject == *root { object } else { subject };
571 let Some(entity) = endpoint_entities.get(&endpoint) else { continue };
572 entities.entry(endpoint).or_insert_with(|| entity.clone());
573 relations.entry(relation_id).or_insert_with(|| relation.clone());
574 count += 1;
575 if count == limit { break; }
576 }
577 }
578 out.insert(*record_id, Neighborhood { entities: entities.into_values().collect(), relations: relations.into_values().take(limit).collect() });
579 }
580 Ok(out)
581 }
582}
583
584#[cfg(test)]
585mod tests {
586 use super::*;
587 use crate::MemoryInput;
588 use std::sync::atomic::Ordering;
589
590 #[test]
594 fn token_count_follows_character_density() {
595 assert_eq!(text::count_tokens(""), 0);
596 assert_eq!(text::count_tokens("abcd"), 1);
597 assert_eq!(text::count_tokens("abcde"), 2);
598 assert_eq!(text::count_tokens("中"), 1);
599 assert_eq!(text::count_tokens("中国"), 1);
600 assert_eq!(text::count_tokens("中国人"), 2);
601 assert_eq!(text::truncate_to_tokens("abcd", 1), "abcd");
603 assert_eq!(text::truncate_to_tokens("abcde", 1), "abcd");
604 assert_eq!(text::truncate_to_tokens("中国人", 1), "中国");
605 assert!(text::count_tokens(&text::truncate_to_tokens("中国人", 1)) <= 1);
606 }
607
608 #[test]
609 fn index_query_failure_degrades_without_rebuilding() {
610 let dir = tempfile::tempdir().unwrap();
611 let kb = KnowledgeBase::open(dir.path()).unwrap();
612 kb.memories().upsert(MemoryInput::new("索引故障恢复的独有措辞")).unwrap();
613 kb.update_index().unwrap();
615 let index = kb.index().unwrap();
616 let request = SearchRequest {
617 query: "索引故障恢复的独有措辞".into(), kinds: vec![RecordKind::Memory],
618 vector: false, rerank: false, ..Default::default()
619 };
620
621 index.fail_search.store(true, Ordering::SeqCst);
622 let degraded = kb.search(&request).unwrap();
623 assert!(degraded.diagnostics.degraded.contains(&Degrade::TextIndexUnavailable),
624 "索引查询失败应隔离文本路");
625
626 index.fail_search.store(false, Ordering::SeqCst);
627 let recovered = kb.search(&request).unwrap();
628 assert_eq!(recovered.hits.len(), 1, "故障排除后索引可用,照常命中");
629 assert!(!recovered.diagnostics.degraded.contains(&Degrade::TextIndexUnavailable));
630 }
631}