use std::collections::HashMap;
use serde_json::Value;
use super::{
reject_reserved_keys, strip_reserved_keys, MemoryService, Metadata, HUB_FIELD,
MENTIONS_RELATION,
};
use crate::embedder::Embedder;
use crate::error::MemoryError;
use crate::fusion::{self, Candidate};
use crate::model::{FusionOptions, MemoryEdge, MemoryNode, Recollection};
use crate::rerank::Reranker;
use crate::storage::MemoryStore;
impl<E: Embedder, S: MemoryStore> MemoryService<E, S> {
pub fn recall_fused(
&self,
query: &str,
k: usize,
filter: Option<&Metadata>,
opts: FusionOptions,
) -> Result<Vec<Recollection>, MemoryError> {
let query = query.trim();
if query.is_empty() || k == 0 {
return Ok(Vec::new());
}
let opts = opts.sanitized();
reject_reserved_keys(filter)?;
let embedding = self.embedder.embed(query)?;
let pool = self.fused_pool(&embedding, pool_depth(k, opts), filter)?;
let reached = self.graph_reached(&embedding, filter, opts.hops)?;
Ok(fusion::fuse(pool, &reached, k, opts.graph_boost))
}
pub fn recall_fused_dated(
&self,
query: &str,
k: usize,
filter: Option<&Metadata>,
opts: FusionOptions,
date_field: &str,
) -> Result<(Vec<Recollection>, crate::DatedContext), MemoryError> {
let hits = self.recall_fused(query, k, filter, opts)?;
let ctx = crate::format_dated_context(&hits, date_field);
Ok((hits, ctx))
}
pub fn recall_fused_reranked<R: Reranker>(
&self,
query: &str,
k: usize,
filter: Option<&Metadata>,
opts: FusionOptions,
reranker: &R,
) -> Result<Vec<Recollection>, MemoryError> {
let query = query.trim();
if query.is_empty() || k == 0 {
return Ok(Vec::new());
}
let opts = opts.sanitized();
reject_reserved_keys(filter)?;
let embedding = self.embedder.embed(query)?;
let depth = pool_depth(k, opts);
let pool = self.fused_pool(&embedding, depth, filter)?;
let reached = self.graph_reached(&embedding, filter, opts.hops)?;
let fused = fusion::fuse(pool, &reached, depth, opts.graph_boost);
let ranked = reranker.rerank(query, fused)?;
Ok(ranked.into_iter().take(k).collect())
}
fn fused_pool(
&self,
embedding: &[f32],
depth: usize,
filter: Option<&Metadata>,
) -> Result<Vec<Candidate>, MemoryError> {
let hits = self.search(embedding, depth, filter)?;
let ids: Vec<u64> = hits.iter().map(|(id, _, _)| *id).collect();
let metadata = self.recall_metadata_batch(&ids)?;
Ok(hits
.into_iter()
.zip(metadata)
.map(|((id, score, content), metadata)| Candidate {
recollection: Recollection {
id,
score,
content,
metadata,
},
vector_score: f64::from(score),
graph_weight: 0.0,
})
.collect())
}
pub(crate) fn recall_metadata_batch(
&self,
ids: &[u64],
) -> Result<Vec<Option<Metadata>>, MemoryError> {
Ok(self
.store
.get_metadata_batch(ids)?
.into_iter()
.map(strip_reserved_keys)
.collect())
}
fn graph_reached(
&self,
embedding: &[f32],
filter: Option<&Metadata>,
hops: usize,
) -> Result<Vec<Candidate>, MemoryError> {
let seeds = self.search(embedding, 1, filter)?;
let Some((seed_id, _score, seed_content)) = seeds.into_iter().next() else {
return Ok(Vec::new());
};
let explanation = self.traverse(seed_id, seed_content, hops)?;
let nodes: Vec<&MemoryNode> = explanation.nodes.iter().filter(|n| n.hop != 0).collect();
let ids: Vec<u64> = nodes.iter().map(|n| n.id).collect();
let raw_payloads = self.store.get_metadata_batch(&ids)?;
let mut idf_cache: HashMap<u64, f64> = HashMap::new();
let mut reached = Vec::new();
for (node, raw) in nodes.into_iter().zip(raw_payloads) {
if let Some(candidate) =
self.reached_candidate(node, raw, &explanation.edges, filter, &mut idf_cache)?
{
reached.push(candidate);
}
}
Ok(reached)
}
fn reached_candidate(
&self,
node: &MemoryNode,
raw: Option<Metadata>,
edges: &[MemoryEdge],
filter: Option<&Metadata>,
idf_cache: &mut HashMap<u64, f64>,
) -> Result<Option<Candidate>, MemoryError> {
if raw
.as_ref()
.is_some_and(|meta| meta.get(HUB_FIELD) == Some(&Value::Bool(true)))
{
return Ok(None);
}
let metadata = strip_reserved_keys(raw);
if !matches_filter(metadata.as_ref(), filter) {
return Ok(None);
}
let weight = self.reach_weight(node.id, edges, idf_cache)?;
Ok(Some(Candidate {
recollection: Recollection {
id: node.id,
score: 0.0,
content: node.content.clone(),
metadata,
},
vector_score: 0.0,
graph_weight: weight,
}))
}
fn reach_weight(
&self,
fact_id: u64,
edges: &[MemoryEdge],
idf_cache: &mut HashMap<u64, f64>,
) -> Result<f64, MemoryError> {
let mut weight: Option<f64> = None;
for edge in edges {
if edge.to == fact_id && edge.relation == MENTIONS_RELATION {
let idf = self.cached_entity_idf(edge.from, idf_cache)?;
weight = Some(weight.map_or(idf, |w: f64| w.max(idf)));
}
}
Ok(weight.unwrap_or(1.0))
}
fn cached_entity_idf(
&self,
hub_id: u64,
cache: &mut HashMap<u64, f64>,
) -> Result<f64, MemoryError> {
if let Some(&idf) = cache.get(&hub_id) {
return Ok(idf);
}
let idf = self.entity_idf(hub_id)?;
cache.insert(hub_id, idf);
Ok(idf)
}
fn entity_idf(&self, hub_id: u64) -> Result<f64, MemoryError> {
let degree = self.store.relations(hub_id)?.len();
let n = self.store.count();
if degree == 0 || n <= 1 {
return Ok(0.0);
}
#[allow(clippy::cast_precision_loss)] let (n, d) = (n as f64, degree as f64);
Ok((n / d).ln() / n.ln())
}
}
fn matches_filter(metadata: Option<&Metadata>, filter: Option<&Metadata>) -> bool {
let Some(filter) = filter else {
return true;
};
if filter.is_empty() {
return true;
}
let Some(metadata) = metadata else {
return false;
};
filter.iter().all(|(k, v)| metadata.get(k) == Some(v))
}
fn pool_depth(k: usize, opts: FusionOptions) -> usize {
let depth = opts.pool.map_or_else(|| fusion::pool_size(k), |p| p.max(1));
crate::limits::clamp_recall_limit(depth)
}