1use rusqlite::{params_from_iter, types::Value as SqlValue, Connection};
8use serde::{de::DeserializeOwned, Deserialize, Serialize};
9use serde_json::Value;
10use std::collections::{BTreeMap, BTreeSet};
11
12use crate::graph::{Entity, Event, Relation};
13use crate::search::{SearchHit, SearchRequest, SearchResult};
14use crate::storage::{self, KnowledgeBase};
15use crate::types::{MatchField, ReadFilter, RecordKind, SearchDiagnostics};
16use crate::{text, Error, Result};
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
20#[serde(rename_all = "snake_case")]
21pub enum SearchPreset { Memory, Graph, Notes, Rag, Broad }
22
23impl SearchPreset {
24 pub const ALL: [Self; 5] = [Self::Memory, Self::Graph, Self::Notes, Self::Rag, Self::Broad];
25
26 pub fn as_str(self) -> &'static str {
27 match self {
28 Self::Memory => "memory", Self::Graph => "graph", Self::Notes => "notes",
29 Self::Rag => "rag", Self::Broad => "broad",
30 }
31 }
32
33 pub fn parse(value: &str) -> Result<Self> {
34 Self::ALL.into_iter().find(|preset| preset.as_str() == value)
35 .ok_or_else(|| Error::Validation(format!("preset must be memory, graph, notes, rag or broad, got {value}")))
36 }
37
38 pub fn uses_memory(self) -> bool { matches!(self, Self::Memory | Self::Rag | Self::Broad) }
40 pub fn uses_graph(self) -> bool { matches!(self, Self::Graph | Self::Rag | Self::Broad) }
42 pub fn uses_notes(self) -> bool { matches!(self, Self::Notes | Self::Broad) }
44}
45
46#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
49#[serde(default)]
50pub struct PresetBudget {
51 pub memory_chars: usize,
53 pub notes_chars: usize,
55 pub seed_entities: usize,
57 pub graph_relations_chars: usize,
59 pub graph_context_chars: usize,
61 pub note_titles: usize,
63}
64
65impl Default for PresetBudget {
66 fn default() -> Self {
67 Self { memory_chars: 2000, notes_chars: 3000, seed_entities: 4,
68 graph_relations_chars: 1000, graph_context_chars: 2000, note_titles: 5 }
69 }
70}
71
72#[derive(Debug, Clone, Serialize, Deserialize)]
73#[serde(default)]
74pub struct PresetRequest {
75 pub preset: SearchPreset,
76 pub query: String,
77 pub filter: ReadFilter,
78 pub embed_space: Option<String>,
80 pub text: bool,
82 pub vector: bool,
84 pub rerank: bool,
86 pub budget: PresetBudget,
88 pub candidate_limit: usize,
90}
91
92impl Default for PresetRequest {
93 fn default() -> Self {
94 Self { preset: SearchPreset::Rag, query: String::new(), filter: ReadFilter::default(),
95 embed_space: None, text: true, vector: true, rerank: true,
96 budget: PresetBudget::default(), candidate_limit: 64 }
97 }
98}
99
100#[derive(Debug, Clone, Default, Serialize, Deserialize)]
102pub struct GraphSection {
103 pub entities: Vec<Entity>,
105 pub relations: Vec<Relation>,
107 pub context_relations: Vec<Relation>,
109 pub context_events: Vec<Event>,
111}
112
113#[derive(Debug, Clone, Default, Serialize, Deserialize)]
115pub struct NoteSection {
116 pub titles: Vec<SearchHit>,
118 pub contents: Vec<SearchHit>,
120 pub paths: Vec<SearchHit>,
122}
123
124#[derive(Debug, Clone, Serialize, Deserialize)]
126pub struct PresetResult {
127 pub preset: SearchPreset,
128 pub memories: Vec<SearchHit>,
129 pub graph: GraphSection,
130 pub notes: NoteSection,
131 pub revision: i64,
132 pub indexed_revision: i64,
133 pub diagnostics: SearchDiagnostics,
134}
135
136impl KnowledgeBase {
137 pub fn search_preset(&self, request: &PresetRequest) -> Result<PresetResult> {
139 let query = request.query.trim();
140 if query.is_empty() { return Err(Error::Validation("a text query is required".into())); }
141 storage::validate_filter(&request.filter)?;
142 if request.budget.seed_entities == 0 { return Err(Error::Validation("seed_entities must be at least 1".into())); }
143 if request.candidate_limit == 0 { return Err(Error::Validation("candidate_limit must be at least 1".into())); }
144
145 let mut diagnostics = SearchDiagnostics::default();
146 let memories = if request.preset.uses_memory() {
147 let result = self.preset_hits(query, request, &[RecordKind::Memory], MatchField::All)?;
148 merge_diagnostics(&mut diagnostics, &result.diagnostics);
149 self.truncate_hits_by_chars(result.hits, request.budget.memory_chars)?
150 } else { Vec::new() };
151 let notes = if request.preset.uses_notes() {
152 self.preset_notes(query, request, &mut diagnostics)?
153 } else { NoteSection::default() };
154 let graph = if request.preset.uses_graph() {
155 self.preset_graph(query, request, &mut diagnostics)?
156 } else { GraphSection::default() };
157
158 let (revision, indexed_revision) = {
159 let state = self.read()?;
160 (storage::current_revision(state.conn())?, storage::meta(state.conn(), "indexed_revision")?)
161 };
162 Ok(PresetResult { preset: request.preset, memories, graph, notes, revision, indexed_revision, diagnostics })
163 }
164
165 fn preset_hits(&self, query: &str, request: &PresetRequest, kinds: &[RecordKind], match_field: MatchField) -> Result<SearchResult> {
167 let inner = SearchRequest {
168 query: query.to_string(), filter: request.filter.clone(), kinds: kinds.to_vec(),
169 limit: request.candidate_limit, embed_space: request.embed_space.clone(),
170 text: request.text, vector: request.vector, rerank: request.rerank,
171 match_field, ..Default::default()
172 };
173 self.search(&inner)
174 }
175
176 fn preset_notes(&self, query: &str, request: &PresetRequest, diagnostics: &mut SearchDiagnostics) -> Result<NoteSection> {
180 let title_result = self.preset_field_hits(query, request, MatchField::Name)?;
182 merge_diagnostics(diagnostics, &title_result.diagnostics);
183 let titles: Vec<SearchHit> = title_result.hits.into_iter().take(request.budget.note_titles).collect();
184 let title_ids: BTreeSet<i64> = titles.iter().map(|hit| hit.key.id).collect();
185
186 let content_result = self.preset_hits(query, request, &[RecordKind::Chunk], MatchField::Text)?;
188 merge_diagnostics(diagnostics, &content_result.diagnostics);
189 let content_hits: Vec<SearchHit> = content_result.hits.into_iter()
190 .filter(|hit| !title_ids.contains(&hit.key.id)).collect();
191 let contents = self.truncate_hits_by_chars(content_hits, request.budget.notes_chars)?;
192 let mut surfaced: BTreeSet<i64> = title_ids;
193 surfaced.extend(contents.iter().map(|hit| hit.key.id));
194
195 let mut paths = Vec::new();
197 if titles.len() < request.budget.note_titles {
198 let need = request.budget.note_titles - titles.len();
199 let path_result = self.preset_field_hits(query, request, MatchField::Path)?;
200 merge_diagnostics(diagnostics, &path_result.diagnostics);
201 paths = path_result.hits.into_iter()
202 .filter(|hit| !surfaced.contains(&hit.key.id)).take(need).collect();
203 }
204 Ok(NoteSection { titles, contents, paths })
205 }
206
207 fn preset_field_hits(&self, query: &str, request: &PresetRequest, match_field: MatchField) -> Result<SearchResult> {
209 let inner = SearchRequest {
210 query: query.to_string(), filter: request.filter.clone(), kinds: vec![RecordKind::Chunk],
211 limit: request.candidate_limit.max(request.budget.note_titles),
212 text: true, vector: false, rerank: false, match_field, ..Default::default()
213 };
214 self.search(&inner)
215 }
216
217 fn truncate_hits_by_chars(&self, hits: Vec<SearchHit>, chars: usize) -> Result<Vec<SearchHit>> {
220 if chars == 0 || hits.is_empty() { return Ok(Vec::new()); }
221 let ids: Vec<i64> = hits.iter().map(|hit| hit.key.id).collect();
222 let bodies = match self.index() { Ok(index) => index.bodies(&ids).unwrap_or_default(), Err(_) => BTreeMap::new() };
224 let mut used = 0usize;
225 let mut out = Vec::new();
226 for hit in hits {
227 let len = bodies.get(&hit.key.id).map(|body| body.chars().count()).unwrap_or(0);
228 if !out.is_empty() && used + len > chars { break; }
229 used += len;
230 out.push(hit);
231 }
232 Ok(out)
233 }
234
235 fn preset_graph(&self, query: &str, request: &PresetRequest, diagnostics: &mut SearchDiagnostics) -> Result<GraphSection> {
237 let entity_result = self.preset_hits(query, request, &[RecordKind::Entity], MatchField::All)?;
239 merge_diagnostics(diagnostics, &entity_result.diagnostics);
240 let seeds: Vec<i64> = entity_result.hits.iter().take(request.budget.seed_entities).map(|hit| hit.key.id).collect();
241 if seeds.is_empty() { return Ok(GraphSection::default()); }
242
243 let state = self.read()?;
244 let conn = state.conn();
245 let entity_filter = ReadFilter { tags: vec![], ..request.filter.clone() };
247 let loaded: BTreeMap<i64, Entity> = storage::load_many(conn, &seeds, &entity_filter)?;
248 let entities: Vec<Entity> = seeds.iter().filter_map(|id| loaded.get(id).cloned()).collect();
249
250 let candidate_ids = incident_relations(conn, &seeds, &request.filter)?;
255 let candidates = storage::record_values(conn, &candidate_ids)?;
256 let remainder = remaining_query_after_entities(query, &entities);
258 let graph_query = crate::graph::match_predicate_synonyms(conn, &request.filter.namespace, &remainder)?;
259 let ranked = rank_relations_by_query(&candidates, &graph_query);
260 let hit_ids = truncate_values_by_chars(&candidates, &ranked, RecordKind::Relation, request.budget.graph_relations_chars);
261 let relations: Vec<Relation> = decode_all(&candidates, &hit_ids)?;
262
263 let mut members: BTreeSet<i64> = seeds.iter().copied().collect();
265 for relation in &relations { members.insert(relation.subject_id); members.insert(relation.object_id); }
266 let member_ids: Vec<i64> = members.into_iter().collect();
267 let hit_set: BTreeSet<i64> = hit_ids.iter().copied().collect();
268
269 let between_ids: Vec<i64> = relations_between(conn, &member_ids, &request.filter)?
270 .into_iter().filter(|id| !hit_set.contains(id)).collect();
271 let between_values = storage::record_values(conn, &between_ids)?;
272 let between_order: Vec<i64> = between_ids.iter().copied().filter(|id| between_values.contains_key(id)).collect();
273 let kept_relations = truncate_values_by_chars(&between_values, &between_order, RecordKind::Relation, request.budget.graph_context_chars);
274 let used = chars_of(&between_values, &kept_relations, RecordKind::Relation);
275 let context_relations: Vec<Relation> = decode_all(&between_values, &kept_relations)?;
276
277 let event_ids = events_between(conn, &member_ids, &request.filter)?;
278 let remaining = request.budget.graph_context_chars.saturating_sub(used);
279 let lengths = storage::event_text_lengths(conn, &event_ids)?;
282 let event_order: Vec<i64> = event_ids.iter().copied().filter(|id| lengths.contains_key(id)).collect();
283 let kept_events = truncate_ids_by_lengths(&event_order, &lengths, remaining);
284 let event_values = storage::record_values(conn, &kept_events)?;
285 let context_events: Vec<Event> = decode_all(&event_values, &kept_events)?;
286
287 Ok(GraphSection { entities, relations, context_relations, context_events })
288 }
289}
290
291fn merge_diagnostics(target: &mut SearchDiagnostics, source: &SearchDiagnostics) {
293 target.text_used |= source.text_used;
294 target.vector_used |= source.vector_used;
295 target.reranked |= source.reranked;
296 target.rerank_candidates += source.rerank_candidates;
297 target.rerank_truncated += source.rerank_truncated;
298 for degrade in &source.degraded {
299 if !target.degraded.contains(degrade) { target.degraded.push(*degrade); }
300 }
301}
302
303fn token_weight(body: &str, tokens: &[String]) -> usize {
305 let present: BTreeSet<String> = text::tokenize(body).into_iter().collect();
306 tokens.iter().map(|token| {
307 if !present.contains(token) { 0 } else if token.chars().count() > 1 { 2 } else { 1 }
308 }).sum()
309}
310
311fn remaining_query_after_entities(query: &str, entities: &[Entity]) -> String {
316 let mut remainder = query.to_string();
317 for entity in entities {
318 for term in std::iter::once(&entity.name).chain(entity.aliases.iter()) {
319 if !term.is_empty() { remainder = remainder.replace(term.as_str(), " "); }
320 }
321 }
322 remainder
323}
324
325fn rank_relations_by_query(values: &BTreeMap<i64, Value>, query: &str) -> Vec<i64> {
327 let tokens = text::query_terms(query, true);
328 if tokens.is_empty() { return values.keys().copied().collect(); }
329 let mut scored: Vec<(usize, i64)> = values.iter()
330 .map(|(id, payload)| (token_weight(&storage::record_text(RecordKind::Relation, payload), &tokens), *id))
331 .filter(|(weight, _)| *weight > 0)
332 .collect();
333 scored.sort_by(|a, b| b.0.cmp(&a.0).then_with(|| a.1.cmp(&b.1)));
334 scored.into_iter().map(|(_, id)| id).collect()
335}
336
337fn truncate_values_by_chars(values: &BTreeMap<i64, Value>, order: &[i64], kind: RecordKind, chars: usize) -> Vec<i64> {
339 if chars == 0 { return Vec::new(); }
340 let mut used = 0usize;
341 let mut kept = Vec::new();
342 for id in order {
343 let Some(payload) = values.get(id) else { continue };
344 let len = storage::record_text(kind, payload).chars().count();
345 if !kept.is_empty() && used + len > chars { break; }
346 used += len;
347 kept.push(*id);
348 }
349 kept
350}
351
352fn truncate_ids_by_lengths(order: &[i64], lengths: &BTreeMap<i64, usize>, chars: usize) -> Vec<i64> {
355 if chars == 0 { return Vec::new(); }
356 let mut used = 0usize;
357 let mut kept = Vec::new();
358 for id in order {
359 let Some(len) = lengths.get(id) else { continue };
360 if !kept.is_empty() && used + len > chars { break; }
361 used += len;
362 kept.push(*id);
363 }
364 kept
365}
366
367fn chars_of(values: &BTreeMap<i64, Value>, ids: &[i64], kind: RecordKind) -> usize {
368 ids.iter().filter_map(|id| values.get(id)).map(|payload| storage::record_text(kind, payload).chars().count()).sum()
369}
370
371fn decode_all<T: DeserializeOwned>(values: &BTreeMap<i64, Value>, ids: &[i64]) -> Result<Vec<T>> {
372 let mut out = Vec::new();
373 for id in ids {
374 if let Some(payload) = values.get(id) { out.push(serde_json::from_value(payload.clone())?); }
375 }
376 Ok(out)
377}
378
379fn incident_relations(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
381 let (condition, values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
382 let placeholders = vec!["?"; entity_ids.len()].join(",");
383 let sql = format!("SELECT rl.record_id FROM relations rl JOIN records r ON r.id=rl.record_id \
384 WHERE (rl.subject_id IN ({placeholders}) OR rl.object_id IN ({placeholders})) AND {condition} ORDER BY rl.record_id");
385 let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id))
386 .chain(entity_ids.iter().map(|id| SqlValue::Integer(*id))).chain(values).collect();
387 let mut stmt = conn.prepare(&sql)?;
388 let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
389 Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
390}
391
392fn relations_between(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
394 let (condition, values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
395 let placeholders = vec!["?"; entity_ids.len()].join(",");
396 let sql = format!("SELECT rl.record_id FROM relations rl JOIN records r ON r.id=rl.record_id \
397 WHERE rl.subject_id IN ({placeholders}) AND rl.object_id IN ({placeholders}) AND {condition} ORDER BY rl.record_id");
398 let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id))
399 .chain(entity_ids.iter().map(|id| SqlValue::Integer(*id))).chain(values).collect();
400 let mut stmt = conn.prepare(&sql)?;
401 let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
402 Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
403}
404
405fn events_between(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
411 let (condition, values) = storage::filter_sql(filter, &[RecordKind::Event], false)?;
412 let placeholders = vec!["?"; entity_ids.len()].join(",");
413 let sql = format!("SELECT e.event_id FROM (SELECT event_id FROM event_participants \
414 WHERE entity_id IN ({placeholders}) GROUP BY event_id HAVING COUNT(DISTINCT entity_id) >= 2) e \
415 JOIN records r ON r.id=e.event_id WHERE {condition} ORDER BY e.event_id");
416 let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id)).chain(values).collect();
417 let mut stmt = conn.prepare(&sql)?;
418 let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
419 Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
420}