use crate::{embeddings, graph::{Entity, Neighborhood, Relation}, storage::{self, KnowledgeBase}, text, types::*, Error, Result};
use parking_lot::Mutex;
use rusqlite::{params_from_iter, types::Value as SqlValue, OptionalExtension};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
use std::sync::Arc;
fn default_limit() -> usize { 10 }
fn mark(stages: &mut Option<crate::events::StageTimer>, name: &str) {
if let Some(stages) = stages.as_mut() { stages.mark(name); }
}
fn yes() -> bool { true }
pub fn default_kinds() -> Vec<RecordKind> { vec![RecordKind::Memory, RecordKind::Entity, RecordKind::Relation, RecordKind::Event, RecordKind::Chunk] }
fn string_id(conn: &rusqlite::Connection, value: &str) -> Result<Option<i64>> {
Ok(conn.query_row("SELECT id FROM strings WHERE text=?1", [text::normalized_tag(value)], |r| r.get(0)).optional()?)
}
fn index_filter(conn: &rusqlite::Connection, filter: &ReadFilter, kinds: &[RecordKind]) -> Result<Option<crate::index::IndexFilter>> {
let Some(namespace) = string_id(conn, &filter.namespace)? else { return Ok(None) };
let mut scopes = Vec::with_capacity(filter.scopes.len());
for scope in &filter.scopes {
match string_id(conn, scope)? { Some(id) => scopes.push(id), None => return Ok(None) }
}
let mut tags = Vec::with_capacity(filter.tags.len());
for tag in &filter.tags {
match string_id(conn, tag)? { Some(id) => tags.push(id), None => return Ok(None) }
}
Ok(Some(crate::index::IndexFilter { namespace, scopes, kinds: kinds.iter().map(|kind| kind.code()).collect(), tags, note_ids: filter.note_ids.clone() }))
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GraphPrune {
pub root: i64,
pub depth: usize,
pub limit: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default)]
pub struct SearchRequest {
pub query: String, pub filter: ReadFilter, pub kinds: Vec<RecordKind>,
pub limit: usize, pub candidate_limit: Option<usize>,
pub embed_space: Option<String>,
pub text_weight: f64, pub prune: Option<GraphPrune>,
#[serde(default = "yes")] pub text: bool,
#[serde(default = "yes")] pub vector: bool,
#[serde(default = "yes")] pub rerank: bool,
#[serde(default)] pub with_total: bool,
#[serde(default)] pub match_field: MatchField,
#[serde(default = "default_top_chunks_per_note")] pub top_chunks_per_note: usize,
}
fn default_top_chunks_per_note() -> usize { 3 }
impl Default for SearchRequest {
fn default() -> Self {
Self { query: String::new(), filter: ReadFilter::default(), kinds: default_kinds(), limit: default_limit(),
candidate_limit: None, embed_space: None, text_weight: 1.0, prune: None,
text: true, vector: true, rerank: true, with_total: false, match_field: MatchField::All,
top_chunks_per_note: default_top_chunks_per_note() }
}
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct ChunkRef {
pub id: i64,
pub offset: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SearchHit {
pub key: RecordKey,
pub score: f64, pub text_score: Option<f64>, pub vector_scores: BTreeMap<String, f64>,
#[serde(default)] pub rerank_score: Option<f64>,
#[serde(default)] pub note_chunks: Option<usize>,
#[serde(default)] pub top_chunks: Vec<ChunkRef>,
pub record: Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SearchResult {
pub hits: Vec<SearchHit>, pub revision: i64, pub indexed_revision: i64,
#[serde(default)] pub total: Option<usize>,
#[serde(default)] pub diagnostics: SearchDiagnostics,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ContextualHit { pub hit: SearchHit, pub context: Neighborhood }
pub trait Reranker: Send {
fn rerank(&mut self, query: &str, documents: &[String]) -> std::result::Result<Vec<f32>, String>;
}
impl<F> Reranker for F
where F: FnMut(&str, &[String]) -> std::result::Result<Vec<f32>, String> + Send {
fn rerank(&mut self, query: &str, documents: &[String]) -> std::result::Result<Vec<f32>, String> { self(query, documents) }
}
fn default_max_tokens_total() -> usize { 8192 }
fn default_max_tokens_per_doc() -> usize { 1024 }
fn default_max_candidates() -> usize { 50 }
fn merge_candidates(text: &[(RecordKey, f64)], vector: &[(RecordKey, f64)]) -> Vec<(RecordKey, f64)> {
let mut seen = HashSet::new();
let mut merged = Vec::with_capacity(text.len() + vector.len());
let mut index = 0;
while index < text.len() || index < vector.len() {
if let Some(item) = text.get(index) { if seen.insert(item.0) { merged.push(*item); } }
if let Some(item) = vector.get(index) { if seen.insert(item.0) { merged.push(*item); } }
index += 1;
}
merged
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct RerankerOptions {
#[serde(default = "default_max_tokens_total")] pub max_tokens_total: usize,
#[serde(default = "default_max_candidates")] pub max_candidates: usize,
#[serde(default = "default_max_tokens_per_doc")] pub max_tokens_per_doc: usize,
#[serde(default)] pub max_tokens_query: Option<usize>,
}
impl Default for RerankerOptions {
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 } }
}
pub(crate) struct RerankerEntry { pub options: RerankerOptions, pub reranker: Box<dyn Reranker> }
#[derive(Default)]
pub(crate) struct RerankerRegistry { entry: Mutex<Option<Arc<Mutex<RerankerEntry>>>> }
impl RerankerRegistry {
pub fn new() -> Self { Self::default() }
pub fn is_registered(&self) -> bool { self.entry.lock().is_some() }
pub fn get(&self) -> Option<Arc<Mutex<RerankerEntry>>> { self.entry.lock().clone() }
pub fn register(&self, entry: RerankerEntry) { *self.entry.lock() = Some(Arc::new(Mutex::new(entry))); }
pub fn remove(&self) -> bool { self.entry.lock().take().is_some() }
}
const RERANK_SAMPLE_DOCS: [&str; 2] = ["重排校验样本一", "rerank probe two"];
impl KnowledgeBase {
pub fn register_reranker<F: Reranker + 'static>(&self, reranker: F) -> Result<()> {
self.register_reranker_with(reranker, RerankerOptions::default())
}
pub fn register_reranker_with<F: Reranker + 'static>(&self, reranker: F, options: RerankerOptions) -> Result<()> {
if options.max_tokens_total == 0 { return Err(Error::Validation("max_tokens_total must be at least 1".into())); }
if options.max_candidates == 0 { return Err(Error::Validation("max_candidates must be at least 1".into())); }
if options.max_tokens_per_doc == 0 { return Err(Error::Validation("max_tokens_per_doc must be at least 1".into())); }
if options.max_tokens_query == Some(0) { return Err(Error::Validation("max_tokens_query must be positive".into())); }
let mut entry = RerankerEntry { options, reranker: Box::new(reranker) };
let documents: Vec<String> = RERANK_SAMPLE_DOCS.iter().map(|sample| (*sample).to_string()).collect();
if let Ok(produced) = entry.reranker.rerank("校验样本", &documents) {
if produced.len() != documents.len() {
return Err(Error::Validation(format!("reranker returned {} scores for {} documents", produced.len(), documents.len())));
}
if produced.iter().any(|score| !score.is_finite()) { return Err(Error::Validation("reranker scores must be finite".into())); }
}
self.engine.rerankers.register(entry);
Ok(())
}
pub fn unregister_reranker(&self) -> bool { self.engine.rerankers.remove() }
pub fn reranker_registered(&self) -> bool { self.engine.rerankers.is_registered() }
fn search_text(&self, conn: &rusqlite::Connection, query: &str, filter: &ReadFilter, kinds: &[RecordKind], limit: usize, field: MatchField) -> Result<Vec<(RecordKey, f64)>> {
self.sync_index_if_behind(conn)?;
let Some(index_filter) = index_filter(conn, filter, kinds)? else { return Ok(Vec::new()) };
let expanded = crate::graph::match_predicate_synonyms(conn, &filter.namespace, query)?;
self.index()?.search_in(&expanded, &index_filter, limit, field)
}
pub fn search(&self, request: &SearchRequest) -> Result<SearchResult> {
storage::validate_filter(&request.filter)?;
storage::validate_limit(request.limit)?;
let query = request.query.trim();
if query.is_empty() { return Err(Error::Validation("a text query is required".into())); }
let vector_path = request.vector && request.embed_space.is_some();
if !request.text && !vector_path { return Err(Error::Validation("enable at least one of text or vector".into())); }
if !request.text_weight.is_finite() || request.text_weight <= 0.0 { return Err(Error::Validation("text_weight must be finite and positive".into())); }
let limit = request.candidate_limit.unwrap_or((request.limit * 5).max(1_000).min(10_000));
storage::validate_limit(limit)?;
if limit < request.limit { return Err(Error::Validation("candidate_limit must be at least limit".into())); }
let sink = self.engine.events.get();
let started = sink.is_some().then(|| std::time::Instant::now());
let mut stages = sink.is_some().then(crate::events::StageTimer::start);
let allowed = match &request.prune {
Some(prune) => {
if prune.depth == 0 { return Err(Error::Validation("graph prune depth must be at least 1".into())); }
storage::validate_limit(prune.limit)?;
let mut ids: HashSet<i64> = self.graph().build_graph(&request.filter)?.ego_ids(prune.root, prune.depth, prune.limit).into_iter().collect();
ids.insert(prune.root);
Some(ids)
}
None => None,
};
mark(&mut stages, "prepare");
let mut diagnostics = SearchDiagnostics::default();
let mut embedded_query: Option<(embeddings::EmbeddingSpace, Vec<f32>)> = None;
let mut vector_kinds: Vec<RecordKind> = Vec::new();
if let Some(space_id) = request.embed_space.as_deref().filter(|_| request.vector) {
let gated = { let state = self.read()?; embeddings::namespace_vectorization(state.conn(), &request.filter.namespace)? };
if !gated {
diagnostics.degraded.push(Degrade::NamespaceDisabled);
} else {
let (enabled, ready) = {
let state = self.read()?;
let conn = state.conn();
let namespace = &request.filter.namespace;
(embeddings::enabled_kinds(conn, namespace)?, embeddings::ready_kinds(conn, namespace, space_id)?)
};
let requested: Vec<RecordKind> = request.kinds.iter().copied().filter(|kind| enabled.contains(kind)).collect();
vector_kinds = requested.iter().copied().filter(|kind| ready.contains(kind)).collect();
if !requested.is_empty() {
let space = { let state = self.read()?; embeddings::get_space(state.conn(), space_id)? };
if vector_kinds.is_empty() {
diagnostics.degraded.push(Degrade::VectorNotReady);
} else {
match self.engine.embedders.get(space_id) {
None => diagnostics.degraded.push(Degrade::NoEmbedder),
Some(entry) => {
let produced = { let mut guard = entry.lock(); guard.embed(&[query.to_string()]) };
match produced {
Ok(mut values) if values.len() == 1 => embedded_query = Some((space, values.remove(0))),
_ => diagnostics.degraded.push(Degrade::EmbedFailed),
}
}
}
}
}
}
}
if vector_path { mark(&mut stages, "embed"); }
let state = self.read()?;
let conn = state.conn();
let mut text_rank: Vec<(RecordKey, f64)> = Vec::new();
if request.text {
let text_hits = match self.search_text(conn, query, &request.filter, &request.kinds, limit, request.match_field) {
Ok(hits) => Some(hits),
Err(Error::Index(_)) => match self.rebuild_indexes() {
Ok(_) => match self.search_text(conn, query, &request.filter, &request.kinds, limit, request.match_field) {
Ok(hits) => Some(hits),
Err(_) => { diagnostics.degraded.push(Degrade::TextIndexUnavailable); Some(Vec::new()) }
},
Err(_) => { diagnostics.degraded.push(Degrade::TextIndexUnavailable); Some(Vec::new()) }
},
Err(error) => return Err(error),
};
text_rank = text_hits.unwrap_or_default();
diagnostics.text_used = true;
}
mark(&mut stages, "text");
let mut vector_rank: Vec<(RecordKey, f64)> = Vec::new();
let mut vector_space_id: Option<String> = None;
if let Some((space, vector)) = &embedded_query {
let namespace = text::normalized_tag(&request.filter.namespace);
let scopes: Vec<String> = request.filter.scopes.iter().map(|s| text::normalized_tag(s)).collect();
let tags: Vec<String> = request.filter.tags.iter().map(|t| text::normalized_tag(t)).collect();
let mut scored: Vec<(RecordKey, f64)> = Vec::new();
for scope in scopes {
let Some(partition) = self.partition(conn, space, &namespace, &scope)? else { continue };
scored.extend(partition.search(vector, &vector_kinds, &tags, &request.filter.note_ids, limit, allowed.as_ref())?);
}
scored.sort_by(|a,b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
scored.truncate(limit);
diagnostics.vector_used = true;
vector_space_id = Some(space.id.clone());
vector_rank = scored;
}
if diagnostics.vector_used { mark(&mut stages, "vector"); }
let mut scores: BTreeMap<RecordKey, (f64, Option<f64>, BTreeMap<String,f64>)> = BTreeMap::new();
for (rank, (key, score)) in text_rank.iter().enumerate() {
let hit = scores.entry(*key).or_default();
hit.0 += request.text_weight / (60.0 + (rank + 1) as f64); hit.1 = Some(*score);
}
if let Some(space_id) = &vector_space_id {
for (rank, (key, score)) in vector_rank.iter().enumerate() {
let hit = scores.entry(*key).or_default();
hit.0 += 1.0 / (60.0 + (rank + 1) as f64); hit.2.insert(space_id.clone(), *score);
}
}
let total = if request.with_total { Some(storage::count_matches(conn, &request.filter, &request.kinds)?) } else { None };
let mut rerank_scores: BTreeMap<RecordKey, f64> = BTreeMap::new();
let rerank_entry = if request.rerank { self.engine.rerankers.get() } else { None };
let use_rerank = rerank_entry.is_some();
let ordered: Vec<RecordKey> = if use_rerank {
merge_candidates(&text_rank, &vector_rank).into_iter().map(|(key, _)| key).collect()
} else {
let mut by_score: Vec<(RecordKey, f64)> = scores.iter().map(|(key, hit)| (*key, hit.0)).collect();
by_score.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
by_score.into_iter().map(|(key, _)| key).collect()
};
let candidate_count = ordered.len();
mark(&mut stages, "fuse");
let aggregate = request.text && matches!(request.match_field, MatchField::All | MatchField::Text);
let per_note = if aggregate { request.top_chunks_per_note } else { 0 };
let ordered_ids: Vec<i64> = ordered.iter().map(|key| key.id).collect();
let chunk_notes = storage::chunk_notes(conn, &ordered_ids)?;
let mut folded: Vec<RecordKey> = Vec::new();
let mut top_chunks: HashMap<i64, Vec<ChunkRef>> = HashMap::new();
if chunk_notes.is_empty() {
folded = ordered;
} else {
let mut seen_notes: HashSet<i64> = HashSet::new();
for key in ordered {
let Some(&(note_id, offset)) = chunk_notes.get(&key.id) else { folded.push(key); continue };
if seen_notes.insert(note_id) {
if per_note > 0 { top_chunks.insert(note_id, vec![ChunkRef { id: key.id, offset }]); }
folded.push(key);
} else if per_note > 0 {
let list = top_chunks.entry(note_id).or_default();
if list.len() < per_note { list.push(ChunkRef { id: key.id, offset }); }
}
}
}
let folded_count = folded.len();
mark(&mut stages, "fold");
let mut selected: Vec<RecordKey> = Vec::new();
let mut reranked = false;
let mut rerank_docs = 0usize;
let mut rerank_tokens = 0usize;
if let Some(entry) = rerank_entry {
let options = entry.lock().options;
let budgeted_query = options.max_tokens_query.map(|budget| text::truncate_to_tokens(query, budget)).unwrap_or_else(|| query.to_string());
let ids: Vec<i64> = folded.iter().map(|key| key.id).collect();
let bodies = match self.index() { Ok(index) => index.bodies(&ids)?, Err(_) => BTreeMap::new() };
let names = storage::entity_names(conn, &ids).unwrap_or_default();
let mut used = text::count_tokens(&budgeted_query);
let mut candidates: Vec<RecordKey> = Vec::new();
let mut documents: Vec<String> = Vec::new();
for key in &folded {
if candidates.len() >= options.max_candidates { break; }
let body = bodies.get(&key.id).map(String::as_str).unwrap_or("");
let full = match names.get(&key.id).filter(|name| !name.is_empty()) {
Some(name) => format!("{name} {body}"),
None => body.to_string(),
};
let document = text::truncate_to_tokens(&full, options.max_tokens_per_doc);
let cost = text::count_tokens(&document);
if used + cost > options.max_tokens_total { break; }
used += cost;
candidates.push(*key);
documents.push(document);
}
diagnostics.rerank_candidates = candidates.len();
diagnostics.rerank_truncated = folded.len().saturating_sub(candidates.len());
rerank_docs = candidates.len();
rerank_tokens = used;
let produced = if candidates.is_empty() { Ok(Vec::new()) }
else { let mut guard = entry.lock(); guard.reranker.rerank(&budgeted_query, &documents) };
match produced {
Ok(values) if values.len() == candidates.len() && values.iter().all(|value| value.is_finite()) => {
let mut pairs: Vec<(RecordKey, f32)> = candidates.into_iter().zip(values).collect();
pairs.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
for (key, score) in pairs { rerank_scores.insert(key, f64::from(score)); selected.push(key); }
diagnostics.reranked = true;
reranked = true;
}
_ => diagnostics.degraded.push(Degrade::RerankFailed),
}
}
if !reranked { selected = folded; }
if use_rerank { mark(&mut stages, "rerank"); }
selected.truncate(request.limit);
let note_counts: BTreeMap<i64, usize> = if aggregate {
match index_filter(conn, &request.filter, &request.kinds)? {
Some(ifilter) => {
let mut targets: Vec<i64> = selected.iter().filter_map(|key| chunk_notes.get(&key.id).map(|(note_id, _)| *note_id)).collect();
targets.sort_unstable();
targets.dedup();
self.index()?.count_in_many(query, &ifilter, request.match_field, &targets)?.into_iter().collect()
}
None => BTreeMap::new(),
}
} else { BTreeMap::new() };
mark(&mut stages, "count");
let ids: Vec<i64> = selected.iter().map(|key| key.id).collect();
let mut records: BTreeMap<i64, Value> = storage::load_many(conn, &ids, &request.filter)?;
let mut hits = Vec::new();
for key in selected {
let Some(record) = records.remove(&key.id) else { continue };
let note_id = chunk_notes.get(&key.id).map(|(note_id, _)| *note_id);
let note_chunks = note_id.and_then(|id| note_counts.get(&id).copied());
let top_chunks = note_id.and_then(|id| top_chunks.remove(&id)).unwrap_or_default();
let (score, text_score, vector_scores) = scores.remove(&key).unwrap_or((0.0, None, BTreeMap::new()));
hits.push(SearchHit { record, key, score, text_score, vector_scores, rerank_score: rerank_scores.get(&key).copied(), note_chunks, top_chunks });
}
for degrade in &diagnostics.degraded { self.note_degrade(*degrade); }
if let (Some(sink), Some(stages)) = (sink, stages) {
let mut event = crate::events::LogEvent::new("search");
event.ms = started.map(|started| started.elapsed().as_millis() as u64).unwrap_or(0);
event.stages = stages.finish("load");
event.candidates = Some(candidate_count);
event.folded = Some(folded_count);
event.rerank_docs = Some(rerank_docs);
event.rerank_tokens = Some(rerank_tokens);
event.hits = Some(hits.len());
event.degraded = diagnostics.degraded.clone();
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| sink(&event)));
}
Ok(SearchResult { hits, revision: storage::current_revision(conn)?, indexed_revision: storage::meta(conn, "indexed_revision")?, total, diagnostics })
}
pub fn search_with_context(&self, request: &SearchRequest, limit: usize) -> Result<Vec<ContextualHit>> {
storage::validate_limit(limit)?;
let scope = ReadFilter { tags: vec![], ..request.filter.clone() };
let hits = self.search(request)?.hits;
let keys: Vec<i64> = hits.iter().map(|hit| hit.key.id).collect();
let mut contexts = self.entity_contexts(&keys, &scope, limit)?;
Ok(hits.into_iter().map(|hit| ContextualHit {
context: contexts.remove(&hit.key.id).unwrap_or_else(|| Neighborhood { entities: vec![], relations: vec![] }),
hit,
}).collect())
}
fn entity_contexts(&self, record_ids: &[i64], filter: &ReadFilter, limit: usize) -> Result<BTreeMap<i64, Neighborhood>> {
let empty = || Neighborhood { entities: vec![], relations: vec![] };
let mut out: BTreeMap<i64, Neighborhood> = record_ids.iter().map(|id| (*id, empty())).collect();
if record_ids.is_empty() { return Ok(out); }
let state = self.read()?;
let conn = state.conn();
let (entity_condition, entity_values) = storage::filter_sql(filter, &[RecordKind::Entity], false)?;
let record_placeholders = vec!["?"; record_ids.len()].join(",");
let mut stmt = conn.prepare(&format!(
"SELECT DISTINCT rt.record_id, ea.entity_id FROM record_tags rt \
JOIN entity_aliases ea ON ea.alias_id=rt.tag_id \
JOIN records r ON r.id=ea.entity_id \
WHERE rt.record_id IN ({record_placeholders}) AND {entity_condition} ORDER BY rt.record_id, ea.entity_id"
))?;
let params = record_ids.iter().map(|id| SqlValue::Integer(*id)).chain(entity_values).collect::<Vec<_>>();
let mut seeds: BTreeMap<i64, BTreeSet<i64>> = BTreeMap::new();
for row in stmt.query_map(params_from_iter(params), |r| Ok((r.get::<_, i64>(0)?, r.get::<_, i64>(1)?)))? {
let (record_id, entity_id) = row?;
seeds.entry(record_id).or_default().insert(entity_id);
}
let roots: Vec<i64> = seeds.values().flatten().copied().collect::<BTreeSet<_>>().into_iter().collect();
if roots.is_empty() { return Ok(out); }
let (relation_condition, relation_values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
let root_placeholders = vec!["?"; roots.len()].join(",");
let mut stmt = conn.prepare(&format!(
"SELECT rl.record_id, rl.subject_id, rl.object_id FROM relations rl JOIN records r ON r.id=rl.record_id \
WHERE (rl.subject_id IN ({root_placeholders}) OR rl.object_id IN ({root_placeholders})) AND {relation_condition} \
ORDER BY rl.record_id"
))?;
let params = roots.iter().map(|id| SqlValue::Integer(*id)).chain(roots.iter().map(|id| SqlValue::Integer(*id))).chain(relation_values).collect::<Vec<_>>();
let mut edges: Vec<(i64, i64, i64)> = Vec::new();
for row in stmt.query_map(params_from_iter(params), |r| Ok((r.get::<_, i64>(0)?, r.get::<_, i64>(1)?, r.get::<_, i64>(2)?)))? {
edges.push(row?);
}
let entity_filter = ReadFilter { tags: vec![], ..filter.clone() };
let root_entities: BTreeMap<i64, Entity> = storage::load_many(conn, &roots, filter)?;
let endpoint_ids: Vec<i64> = edges.iter().flat_map(|(_, subject, object)| [*subject, *object]).collect::<BTreeSet<_>>().into_iter().collect();
let endpoint_entities: BTreeMap<i64, Entity> = storage::load_many(conn, &endpoint_ids, &entity_filter)?;
let relation_ids: Vec<i64> = edges.iter().map(|(id, _, _)| *id).collect();
let relation_records: BTreeMap<i64, Relation> = storage::load_many(conn, &relation_ids, filter)?;
let mut incident: BTreeMap<i64, Vec<usize>> = BTreeMap::new();
for (i, &(_, subject, object)) in edges.iter().enumerate() {
incident.entry(subject).or_default().push(i);
if object != subject { incident.entry(object).or_default().push(i); }
}
for (record_id, root_ids) in &seeds {
let mut entities: BTreeMap<i64, Entity> = BTreeMap::new();
for id in root_ids { if let Some(entity) = root_entities.get(id) { entities.insert(*id, entity.clone()); } }
let mut relations: BTreeMap<i64, Relation> = BTreeMap::new();
for root in root_ids {
let mut count = 0usize;
for &i in incident.get(root).map(Vec::as_slice).unwrap_or(&[]) {
let (relation_id, subject, object) = edges[i];
let Some(relation) = relation_records.get(&relation_id) else { continue };
let endpoint = if subject == *root { object } else { subject };
let Some(entity) = endpoint_entities.get(&endpoint) else { continue };
entities.entry(endpoint).or_insert_with(|| entity.clone());
relations.entry(relation_id).or_insert_with(|| relation.clone());
count += 1;
if count == limit { break; }
}
}
out.insert(*record_id, Neighborhood { entities: entities.into_values().collect(), relations: relations.into_values().take(limit).collect() });
}
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::MemoryInput;
use std::sync::atomic::Ordering;
#[test]
fn token_count_follows_character_density() {
assert_eq!(text::count_tokens(""), 0);
assert_eq!(text::count_tokens("abcd"), 1);
assert_eq!(text::count_tokens("abcde"), 2);
assert_eq!(text::count_tokens("中"), 1);
assert_eq!(text::count_tokens("中国"), 1);
assert_eq!(text::count_tokens("中国人"), 2);
assert_eq!(text::truncate_to_tokens("abcd", 1), "abcd");
assert_eq!(text::truncate_to_tokens("abcde", 1), "abcd");
assert_eq!(text::truncate_to_tokens("中国人", 1), "中国");
assert!(text::count_tokens(&text::truncate_to_tokens("中国人", 1)) <= 1);
}
#[test]
fn index_query_failure_rebuilds_then_recovers() {
let dir = tempfile::tempdir().unwrap();
let kb = KnowledgeBase::open(dir.path()).unwrap();
kb.memories().upsert(MemoryInput::new("索引故障恢复的独有措辞")).unwrap();
let index = kb.index().unwrap();
let request = SearchRequest {
query: "索引故障恢复的独有措辞".into(), kinds: vec![RecordKind::Memory],
vector: false, rerank: false, ..Default::default()
};
index.fail_search.store(true, Ordering::SeqCst);
let degraded = kb.search(&request).unwrap();
assert!(index.rebuilds.load(Ordering::SeqCst) >= 1, "索引查询失败必须触发重建");
assert!(degraded.diagnostics.degraded.contains(&Degrade::TextIndexUnavailable),
"重建之后仍失败,才隔离文本路");
index.fail_search.store(false, Ordering::SeqCst);
let recovered = kb.search(&request).unwrap();
assert_eq!(recovered.hits.len(), 1, "故障排除后索引可用,照常命中");
assert!(!recovered.diagnostics.degraded.contains(&Degrade::TextIndexUnavailable));
}
}