use rusqlite::{params_from_iter, types::Value as SqlValue, Connection};
use serde::{de::DeserializeOwned, Deserialize, Serialize};
use serde_json::Value;
use std::collections::{BTreeMap, BTreeSet};
use crate::graph::{Entity, Event, Relation};
use crate::search::{SearchHit, SearchRequest, SearchResult};
use crate::storage::{self, KnowledgeBase};
use crate::types::{MatchField, ReadFilter, RecordKind, SearchDiagnostics};
use crate::{text, Error, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SearchPreset { Memory, Graph, Notes, Rag, Broad }
impl SearchPreset {
pub const ALL: [Self; 5] = [Self::Memory, Self::Graph, Self::Notes, Self::Rag, Self::Broad];
pub fn as_str(self) -> &'static str {
match self {
Self::Memory => "memory", Self::Graph => "graph", Self::Notes => "notes",
Self::Rag => "rag", Self::Broad => "broad",
}
}
pub fn parse(value: &str) -> Result<Self> {
Self::ALL.into_iter().find(|preset| preset.as_str() == value)
.ok_or_else(|| Error::Validation(format!("preset must be memory, graph, notes, rag or broad, got {value}")))
}
pub fn uses_memory(self) -> bool { matches!(self, Self::Memory | Self::Rag | Self::Broad) }
pub fn uses_graph(self) -> bool { matches!(self, Self::Graph | Self::Rag | Self::Broad) }
pub fn uses_notes(self) -> bool { matches!(self, Self::Notes | Self::Broad) }
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default)]
pub struct PresetBudget {
pub memory_chars: usize,
pub notes_chars: usize,
pub seed_entities: usize,
pub graph_relations_chars: usize,
pub graph_context_chars: usize,
pub note_titles: usize,
}
impl Default for PresetBudget {
fn default() -> Self {
Self { memory_chars: 2000, notes_chars: 3000, seed_entities: 4,
graph_relations_chars: 1000, graph_context_chars: 2000, note_titles: 5 }
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default)]
pub struct PresetRequest {
pub preset: SearchPreset,
pub query: String,
pub filter: ReadFilter,
pub embed_space: Option<String>,
pub text: bool,
pub vector: bool,
pub rerank: bool,
pub budget: PresetBudget,
pub candidate_limit: usize,
}
impl Default for PresetRequest {
fn default() -> Self {
Self { preset: SearchPreset::Rag, query: String::new(), filter: ReadFilter::default(),
embed_space: None, text: true, vector: true, rerank: true,
budget: PresetBudget::default(), candidate_limit: 64 }
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct GraphSection {
pub entities: Vec<Entity>,
pub relations: Vec<Relation>,
pub context_relations: Vec<Relation>,
pub context_events: Vec<Event>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct NoteSection {
pub titles: Vec<SearchHit>,
pub contents: Vec<SearchHit>,
pub paths: Vec<SearchHit>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PresetResult {
pub preset: SearchPreset,
pub memories: Vec<SearchHit>,
pub graph: GraphSection,
pub notes: NoteSection,
pub revision: i64,
pub indexed_revision: i64,
pub diagnostics: SearchDiagnostics,
}
impl KnowledgeBase {
pub fn search_preset(&self, request: &PresetRequest) -> Result<PresetResult> {
let query = request.query.trim();
if query.is_empty() { return Err(Error::Validation("a text query is required".into())); }
storage::validate_filter(&request.filter)?;
if request.budget.seed_entities == 0 { return Err(Error::Validation("seed_entities must be at least 1".into())); }
if request.candidate_limit == 0 { return Err(Error::Validation("candidate_limit must be at least 1".into())); }
let mut diagnostics = SearchDiagnostics::default();
let memories = if request.preset.uses_memory() {
let result = self.preset_hits(query, request, &[RecordKind::Memory], MatchField::All)?;
merge_diagnostics(&mut diagnostics, &result.diagnostics);
self.truncate_hits_by_chars(result.hits, request.budget.memory_chars)?
} else { Vec::new() };
let notes = if request.preset.uses_notes() {
self.preset_notes(query, request, &mut diagnostics)?
} else { NoteSection::default() };
let graph = if request.preset.uses_graph() {
self.preset_graph(query, request, &mut diagnostics)?
} else { GraphSection::default() };
let (revision, indexed_revision) = {
let state = self.read()?;
(storage::current_revision(state.conn())?, storage::meta(state.conn(), "indexed_revision")?)
};
Ok(PresetResult { preset: request.preset, memories, graph, notes, revision, indexed_revision, diagnostics })
}
fn preset_hits(&self, query: &str, request: &PresetRequest, kinds: &[RecordKind], match_field: MatchField) -> Result<SearchResult> {
let inner = SearchRequest {
query: query.to_string(), filter: request.filter.clone(), kinds: kinds.to_vec(),
limit: request.candidate_limit, embed_space: request.embed_space.clone(),
text: request.text, vector: request.vector, rerank: request.rerank,
match_field, ..Default::default()
};
self.search(&inner)
}
fn preset_notes(&self, query: &str, request: &PresetRequest, diagnostics: &mut SearchDiagnostics) -> Result<NoteSection> {
let title_result = self.preset_field_hits(query, request, MatchField::Name)?;
merge_diagnostics(diagnostics, &title_result.diagnostics);
let titles: Vec<SearchHit> = title_result.hits.into_iter().take(request.budget.note_titles).collect();
let title_ids: BTreeSet<i64> = titles.iter().map(|hit| hit.key.id).collect();
let content_result = self.preset_hits(query, request, &[RecordKind::Chunk], MatchField::Text)?;
merge_diagnostics(diagnostics, &content_result.diagnostics);
let content_hits: Vec<SearchHit> = content_result.hits.into_iter()
.filter(|hit| !title_ids.contains(&hit.key.id)).collect();
let contents = self.truncate_hits_by_chars(content_hits, request.budget.notes_chars)?;
let mut surfaced: BTreeSet<i64> = title_ids;
surfaced.extend(contents.iter().map(|hit| hit.key.id));
let mut paths = Vec::new();
if titles.len() < request.budget.note_titles {
let need = request.budget.note_titles - titles.len();
let path_result = self.preset_field_hits(query, request, MatchField::Path)?;
merge_diagnostics(diagnostics, &path_result.diagnostics);
paths = path_result.hits.into_iter()
.filter(|hit| !surfaced.contains(&hit.key.id)).take(need).collect();
}
Ok(NoteSection { titles, contents, paths })
}
fn preset_field_hits(&self, query: &str, request: &PresetRequest, match_field: MatchField) -> Result<SearchResult> {
let inner = SearchRequest {
query: query.to_string(), filter: request.filter.clone(), kinds: vec![RecordKind::Chunk],
limit: request.candidate_limit.max(request.budget.note_titles),
text: true, vector: false, rerank: false, match_field, ..Default::default()
};
self.search(&inner)
}
fn truncate_hits_by_chars(&self, hits: Vec<SearchHit>, chars: usize) -> Result<Vec<SearchHit>> {
if chars == 0 || hits.is_empty() { return Ok(Vec::new()); }
let ids: Vec<i64> = hits.iter().map(|hit| hit.key.id).collect();
let bodies = match self.index() { Ok(index) => index.bodies(&ids).unwrap_or_default(), Err(_) => BTreeMap::new() };
let mut used = 0usize;
let mut out = Vec::new();
for hit in hits {
let len = bodies.get(&hit.key.id).map(|body| body.chars().count()).unwrap_or(0);
if !out.is_empty() && used + len > chars { break; }
used += len;
out.push(hit);
}
Ok(out)
}
fn preset_graph(&self, query: &str, request: &PresetRequest, diagnostics: &mut SearchDiagnostics) -> Result<GraphSection> {
let entity_result = self.preset_hits(query, request, &[RecordKind::Entity], MatchField::All)?;
merge_diagnostics(diagnostics, &entity_result.diagnostics);
let seeds: Vec<i64> = entity_result.hits.iter().take(request.budget.seed_entities).map(|hit| hit.key.id).collect();
if seeds.is_empty() { return Ok(GraphSection::default()); }
let state = self.read()?;
let conn = state.conn();
let entity_filter = ReadFilter { tags: vec![], ..request.filter.clone() };
let loaded: BTreeMap<i64, Entity> = storage::load_many(conn, &seeds, &entity_filter)?;
let entities: Vec<Entity> = seeds.iter().filter_map(|id| loaded.get(id).cloned()).collect();
let candidate_ids = incident_relations(conn, &seeds, &request.filter)?;
let candidates = storage::record_values(conn, &candidate_ids)?;
let remainder = remaining_query_after_entities(query, &entities);
let graph_query = crate::graph::match_predicate_synonyms(conn, &request.filter.namespace, &remainder)?;
let ranked = rank_relations_by_query(&candidates, &graph_query);
let hit_ids = truncate_values_by_chars(&candidates, &ranked, RecordKind::Relation, request.budget.graph_relations_chars);
let relations: Vec<Relation> = decode_all(&candidates, &hit_ids)?;
let mut members: BTreeSet<i64> = seeds.iter().copied().collect();
for relation in &relations { members.insert(relation.subject_id); members.insert(relation.object_id); }
let member_ids: Vec<i64> = members.into_iter().collect();
let hit_set: BTreeSet<i64> = hit_ids.iter().copied().collect();
let between_ids: Vec<i64> = relations_between(conn, &member_ids, &request.filter)?
.into_iter().filter(|id| !hit_set.contains(id)).collect();
let between_values = storage::record_values(conn, &between_ids)?;
let between_order: Vec<i64> = between_ids.iter().copied().filter(|id| between_values.contains_key(id)).collect();
let kept_relations = truncate_values_by_chars(&between_values, &between_order, RecordKind::Relation, request.budget.graph_context_chars);
let used = chars_of(&between_values, &kept_relations, RecordKind::Relation);
let context_relations: Vec<Relation> = decode_all(&between_values, &kept_relations)?;
let event_ids = events_between(conn, &member_ids, &request.filter)?;
let remaining = request.budget.graph_context_chars.saturating_sub(used);
let lengths = storage::event_text_lengths(conn, &event_ids)?;
let event_order: Vec<i64> = event_ids.iter().copied().filter(|id| lengths.contains_key(id)).collect();
let kept_events = truncate_ids_by_lengths(&event_order, &lengths, remaining);
let event_values = storage::record_values(conn, &kept_events)?;
let context_events: Vec<Event> = decode_all(&event_values, &kept_events)?;
Ok(GraphSection { entities, relations, context_relations, context_events })
}
}
fn merge_diagnostics(target: &mut SearchDiagnostics, source: &SearchDiagnostics) {
target.text_used |= source.text_used;
target.vector_used |= source.vector_used;
target.reranked |= source.reranked;
target.rerank_candidates += source.rerank_candidates;
target.rerank_truncated += source.rerank_truncated;
for degrade in &source.degraded {
if !target.degraded.contains(degrade) { target.degraded.push(*degrade); }
}
}
fn token_weight(body: &str, tokens: &[String]) -> usize {
let present: BTreeSet<String> = text::tokenize(body).into_iter().collect();
tokens.iter().map(|token| {
if !present.contains(token) { 0 } else if token.chars().count() > 1 { 2 } else { 1 }
}).sum()
}
fn remaining_query_after_entities(query: &str, entities: &[Entity]) -> String {
let mut remainder = query.to_string();
for entity in entities {
for term in std::iter::once(&entity.name).chain(entity.aliases.iter()) {
if !term.is_empty() { remainder = remainder.replace(term.as_str(), " "); }
}
}
remainder
}
fn rank_relations_by_query(values: &BTreeMap<i64, Value>, query: &str) -> Vec<i64> {
let tokens = text::query_terms(query, true);
if tokens.is_empty() { return values.keys().copied().collect(); }
let mut scored: Vec<(usize, i64)> = values.iter()
.map(|(id, payload)| (token_weight(&storage::record_text(RecordKind::Relation, payload), &tokens), *id))
.filter(|(weight, _)| *weight > 0)
.collect();
scored.sort_by(|a, b| b.0.cmp(&a.0).then_with(|| a.1.cmp(&b.1)));
scored.into_iter().map(|(_, id)| id).collect()
}
fn truncate_values_by_chars(values: &BTreeMap<i64, Value>, order: &[i64], kind: RecordKind, chars: usize) -> Vec<i64> {
if chars == 0 { return Vec::new(); }
let mut used = 0usize;
let mut kept = Vec::new();
for id in order {
let Some(payload) = values.get(id) else { continue };
let len = storage::record_text(kind, payload).chars().count();
if !kept.is_empty() && used + len > chars { break; }
used += len;
kept.push(*id);
}
kept
}
fn truncate_ids_by_lengths(order: &[i64], lengths: &BTreeMap<i64, usize>, chars: usize) -> Vec<i64> {
if chars == 0 { return Vec::new(); }
let mut used = 0usize;
let mut kept = Vec::new();
for id in order {
let Some(len) = lengths.get(id) else { continue };
if !kept.is_empty() && used + len > chars { break; }
used += len;
kept.push(*id);
}
kept
}
fn chars_of(values: &BTreeMap<i64, Value>, ids: &[i64], kind: RecordKind) -> usize {
ids.iter().filter_map(|id| values.get(id)).map(|payload| storage::record_text(kind, payload).chars().count()).sum()
}
fn decode_all<T: DeserializeOwned>(values: &BTreeMap<i64, Value>, ids: &[i64]) -> Result<Vec<T>> {
let mut out = Vec::new();
for id in ids {
if let Some(payload) = values.get(id) { out.push(serde_json::from_value(payload.clone())?); }
}
Ok(out)
}
fn incident_relations(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
let (condition, values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
let placeholders = vec!["?"; entity_ids.len()].join(",");
let sql = format!("SELECT rl.record_id FROM relations rl JOIN records r ON r.id=rl.record_id \
WHERE (rl.subject_id IN ({placeholders}) OR rl.object_id IN ({placeholders})) AND {condition} ORDER BY rl.record_id");
let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id))
.chain(entity_ids.iter().map(|id| SqlValue::Integer(*id))).chain(values).collect();
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
}
fn relations_between(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
let (condition, values) = storage::filter_sql(filter, &[RecordKind::Relation], false)?;
let placeholders = vec!["?"; entity_ids.len()].join(",");
let sql = format!("SELECT rl.record_id FROM relations rl JOIN records r ON r.id=rl.record_id \
WHERE rl.subject_id IN ({placeholders}) AND rl.object_id IN ({placeholders}) AND {condition} ORDER BY rl.record_id");
let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id))
.chain(entity_ids.iter().map(|id| SqlValue::Integer(*id))).chain(values).collect();
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
}
fn events_between(conn: &Connection, entity_ids: &[i64], filter: &ReadFilter) -> Result<Vec<i64>> {
let (condition, values) = storage::filter_sql(filter, &[RecordKind::Event], false)?;
let placeholders = vec!["?"; entity_ids.len()].join(",");
let sql = format!("SELECT e.event_id FROM (SELECT event_id FROM event_participants \
WHERE entity_id IN ({placeholders}) GROUP BY event_id HAVING COUNT(DISTINCT entity_id) >= 2) e \
JOIN records r ON r.id=e.event_id WHERE {condition} ORDER BY e.event_id");
let params: Vec<SqlValue> = entity_ids.iter().map(|id| SqlValue::Integer(*id)).chain(values).collect();
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map(params_from_iter(params), |row| row.get::<_, i64>(0))?;
Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
}