Skip to main content

locus_sdk/application/
memory_recall.rs

1use std::collections::HashSet;
2use std::sync::Arc;
3
4use anyhow::Result;
5use locus_core_rs::ContextQueryService;
6use locus_core_rs::domain::contracts::{NodeStore, SemanticIndexStore};
7use locus_core_rs::domain::models::{
8    AvecState, NodeQuery, PsiRange, SemanticTagQueryFilter, SttpNode,
9};
10use locus_core_rs::storage::derive_tenant_id_from_session;
11
12use crate::application::memory_filters::{
13    build_session_filter, node_matches_common_filters, resolve_indexed_sync_keys,
14};
15use crate::application::memory_lexical::{
16    self, LEXICAL_SCAN_LIMIT, LexicalActivation, LexicalFields,
17};
18use crate::domain::memory::{
19    FallbackPolicy, MemoryRecallRequest, MemoryRecallResult, RetrievalPath, clamp_limit,
20};
21
22pub struct MemoryRecallService {
23    store: Arc<dyn NodeStore>,
24    context_query: ContextQueryService,
25    semantic_index: Option<Arc<dyn SemanticIndexStore>>,
26}
27
28impl MemoryRecallService {
29    /// Create a recall service backed by the core resonance query pipeline.
30    pub fn new(store: Arc<dyn NodeStore>) -> Self {
31        Self {
32            context_query: ContextQueryService::new(store.clone()),
33            store,
34            semantic_index: None,
35        }
36    }
37
38    pub fn with_semantic_index(mut self, semantic_index: Arc<dyn SemanticIndexStore>) -> Self {
39        self.semantic_index = Some(semantic_index);
40        self
41    }
42
43    /// Retrieve context nodes using resonance or hybrid scoring,
44    /// with optional lexical fallback when configured.
45    pub async fn execute(&self, request: &MemoryRecallRequest) -> Result<MemoryRecallResult> {
46        let limit = clamp_limit(request.page.limit);
47        let expanded_limit = (limit.saturating_mul(5)).clamp(1, 200);
48
49        let current = request.current_avec.unwrap_or_else(AvecState::zero);
50        let session_scope = request
51            .scope
52            .session_ids
53            .as_deref()
54            .filter(|sessions| sessions.len() == 1)
55            .and_then(|sessions| sessions.first().map(String::as_str));
56        let session_filter = build_session_filter(&request.scope);
57        let tenant_id = request
58            .scope
59            .tenant_id
60            .clone()
61            .or_else(|| session_scope.map(derive_tenant_id_from_session))
62            .unwrap_or_else(|| "default".to_string());
63
64        let indexed_sync_keys = if let Some(index) = self.semantic_index.as_ref() {
65            resolve_indexed_sync_keys(
66                index.as_ref(),
67                &tenant_id,
68                &request.filter,
69                session_scope,
70                expanded_limit,
71            )
72            .await?
73        } else {
74            None
75        };
76
77        let mut path = if request.query_embedding.is_some() {
78            RetrievalPath::Hybrid
79        } else {
80            RetrievalPath::ResonanceOnly
81        };
82
83        let primary = if let Some(query_embedding) = request.query_embedding.as_deref() {
84            self.context_query
85                .get_context_hybrid_scoped_filtered_async(
86                    session_scope,
87                    current.stability,
88                    current.friction,
89                    current.logic,
90                    current.autonomy,
91                    request.scope.from_utc,
92                    request.scope.to_utc,
93                    request.scope.tiers.as_deref(),
94                    Some(query_embedding),
95                    request.scoring.alpha,
96                    request.scoring.beta,
97                    expanded_limit,
98                )
99                .await
100        } else {
101            self.context_query
102                .get_context_scoped_filtered_async(
103                    session_scope,
104                    current.stability,
105                    current.friction,
106                    current.logic,
107                    current.autonomy,
108                    request.scope.from_utc,
109                    request.scope.to_utc,
110                    request.scope.tiers.as_deref(),
111                    expanded_limit,
112                )
113                .await
114        };
115
116        let mut nodes = filter_nodes(
117            primary.nodes,
118            request,
119            session_filter.as_ref(),
120            indexed_sync_keys.as_ref(),
121        );
122
123        if let Some(query_text) = request.query_text.as_deref() {
124            let primary_empty = nodes.is_empty();
125            match memory_lexical::activation(
126                request.scoring.fallback_policy,
127                query_text,
128                primary_empty,
129            ) {
130                LexicalActivation::Skip => {}
131                LexicalActivation::Legacy => {
132                    let fallback_result = self
133                        .context_query
134                        .get_context_scoped_filtered_async(
135                            session_scope,
136                            current.stability,
137                            current.friction,
138                            current.logic,
139                            current.autonomy,
140                            request.scope.from_utc,
141                            request.scope.to_utc,
142                            request.scope.tiers.as_deref(),
143                            expanded_limit,
144                        )
145                        .await;
146
147                    let lexical = memory_lexical::legacy_phrase_filter(
148                        filter_nodes(
149                            fallback_result.nodes,
150                            request,
151                            session_filter.as_ref(),
152                            indexed_sync_keys.as_ref(),
153                        ),
154                        query_text,
155                    );
156
157                    if request.scoring.fallback_policy == FallbackPolicy::Always
158                        && !nodes.is_empty()
159                    {
160                        nodes = memory_lexical::merge_unique(nodes, lexical);
161                    } else {
162                        nodes = lexical;
163                    }
164
165                    path = RetrievalPath::LexicalFallback;
166                }
167                LexicalActivation::NaturalLanguage => {
168                    let scanned = self
169                        .store
170                        .query_nodes_async(NodeQuery {
171                            limit: LEXICAL_SCAN_LIMIT,
172                            session_id: session_scope.map(str::to_string),
173                            from_utc: request.scope.from_utc,
174                            to_utc: request.scope.to_utc,
175                            tiers: request.scope.tiers.clone(),
176                        })
177                        .await?;
178                    let lexical = memory_lexical::select_lexical_matches(
179                        filter_nodes(
180                            scanned,
181                            request,
182                            session_filter.as_ref(),
183                            indexed_sync_keys.as_ref(),
184                        ),
185                        &memory_lexical::parse_lexical_query(query_text),
186                        request.scoring.strictness,
187                        LexicalFields::RECALL,
188                    );
189                    let (ranked, applied) = memory_lexical::apply_natural_language(
190                        nodes,
191                        lexical,
192                        request.query_embedding.is_some(),
193                    );
194                    nodes = ranked;
195                    if request.query_embedding.is_none() && (applied || primary_empty) {
196                        path = RetrievalPath::LexicalFallback;
197                    }
198                }
199            }
200        }
201
202        if request.scoring.gamma > 0.0
203            && let Some(query_tag_embedding) = request.query_tag_embedding.as_deref()
204            && let Some(index) = self.semantic_index.as_ref()
205        {
206            rerank_by_tag_similarity(
207                &mut nodes,
208                index.as_ref(),
209                &tenant_id,
210                query_tag_embedding,
211                request.scoring.gamma,
212            )
213            .await?;
214        }
215
216        let has_more = nodes.len() > limit;
217        nodes.truncate(limit);
218
219        let next_cursor = nodes
220            .last()
221            .map(|node| format!("{}|{}", node.updated_at.to_rfc3339(), node.sync_key));
222
223        let psi_range = psi_range_from_nodes(&nodes);
224
225        Ok(MemoryRecallResult {
226            retrieved: nodes.len(),
227            nodes,
228            psi_range,
229            retrieval_path: path,
230            has_more,
231            next_cursor,
232        })
233    }
234}
235
236async fn rerank_by_tag_similarity(
237    nodes: &mut Vec<SttpNode>,
238    index: &dyn SemanticIndexStore,
239    tenant_id: &str,
240    query_embedding: &[f32],
241    gamma: f32,
242) -> Result<()> {
243    if nodes.is_empty() {
244        return Ok(());
245    }
246
247    let sync_keys: Vec<String> = nodes.iter().map(|node| node.sync_key.clone()).collect();
248    let records = index
249        .query_tag_records_async(SemanticTagQueryFilter {
250            tenant_id: Some(tenant_id.to_string()),
251            tags: None,
252            tag_prefix: None,
253            has_embedding: Some(true),
254            missing_embedding_only: false,
255            limit: sync_keys.len().saturating_mul(16).max(64),
256            session_id: None,
257        })
258        .await?;
259
260    let mut scores: Vec<(usize, f32)> = nodes
261        .iter()
262        .enumerate()
263        .map(|(index, node)| {
264            let tag_score = records
265                .iter()
266                .filter(|record| record.sync_key == node.sync_key)
267                .filter_map(|record| record.embedding.as_deref())
268                .filter_map(|embedding| cosine_similarity(query_embedding, embedding))
269                .fold(0.0_f32, f32::max);
270            (index, tag_score)
271        })
272        .collect();
273
274    scores.sort_by(|left, right| {
275        right
276            .1
277            .partial_cmp(&left.1)
278            .unwrap_or(std::cmp::Ordering::Equal)
279    });
280
281    let mut reranked = Vec::with_capacity(nodes.len());
282    let mut used = HashSet::new();
283    for (index, _) in scores {
284        if used.insert(index) {
285            reranked.push(nodes[index].clone());
286        }
287    }
288
289    if gamma >= 1.0 {
290        *nodes = reranked;
291    } else {
292        let blend_count = ((nodes.len() as f32) * gamma).ceil() as usize;
293        for (slot, node) in reranked.into_iter().take(blend_count).enumerate() {
294            nodes[slot] = node;
295        }
296    }
297
298    Ok(())
299}
300
301fn cosine_similarity(left: &[f32], right: &[f32]) -> Option<f32> {
302    if left.len() != right.len() || left.is_empty() {
303        return None;
304    }
305
306    let mut dot = 0.0_f32;
307    let mut left_norm = 0.0_f32;
308    let mut right_norm = 0.0_f32;
309
310    for (left_value, right_value) in left.iter().zip(right.iter()) {
311        dot += left_value * right_value;
312        left_norm += left_value * left_value;
313        right_norm += right_value * right_value;
314    }
315
316    if left_norm == 0.0 || right_norm == 0.0 {
317        return None;
318    }
319
320    Some(dot / (left_norm.sqrt() * right_norm.sqrt()))
321}
322
323fn filter_nodes(
324    nodes: Vec<SttpNode>,
325    request: &MemoryRecallRequest,
326    session_filter: Option<&HashSet<String>>,
327    indexed_sync_keys: Option<&HashSet<String>>,
328) -> Vec<SttpNode> {
329    nodes
330        .into_iter()
331        .filter(|node| {
332            if let Some(keys) = indexed_sync_keys
333                && !keys.contains(&node.sync_key)
334            {
335                return false;
336            }
337
338            node_matches_common_filters(node, &request.scope, &request.filter, session_filter)
339        })
340        .collect()
341}
342
343fn psi_range_from_nodes(nodes: &[SttpNode]) -> PsiRange {
344    if nodes.is_empty() {
345        return PsiRange::default();
346    }
347
348    let (min, max, sum) = nodes
349        .iter()
350        .fold((f32::MAX, f32::MIN, 0.0_f32), |(min, max, sum), node| {
351            (min.min(node.psi), max.max(node.psi), sum + node.psi)
352        });
353
354    PsiRange {
355        min,
356        max,
357        average: sum / nodes.len() as f32,
358    }
359}
360
361#[cfg(test)]
362mod tests {
363    use std::sync::Arc;
364
365    use chrono::Utc;
366    use locus_core_rs::domain::models::{AvecState, SttpNode};
367    use locus_core_rs::{InMemoryNodeStore, NodeStore};
368
369    use super::MemoryRecallService;
370    use crate::domain::memory::{
371        FallbackPolicy, MemoryPage, MemoryRecallRequest, MemoryScoring, RetrievalPath,
372    };
373
374    #[tokio::test]
375    async fn natural_language_question_returns_the_matching_memory() {
376        let store: Arc<dyn NodeStore> = Arc::new(InMemoryNodeStore::new());
377        store
378            .upsert_node_async(sample(
379                "near",
380                AvecState::zero(),
381                "weekend hiking plan",
382                "unrelated notes",
383            ))
384            .await
385            .expect("upsert near node");
386        store
387            .upsert_node_async(sample(
388                "far",
389                AvecState {
390                    stability: 1.0,
391                    friction: 1.0,
392                    logic: 0.0,
393                    autonomy: 0.0,
394                },
395                "decided to harden the parser grammar",
396                "decision notes",
397            ))
398            .await
399            .expect("upsert far node");
400
401        let service = MemoryRecallService::new(store);
402        let result = service
403            .execute(&MemoryRecallRequest {
404                page: MemoryPage {
405                    limit: 1,
406                    cursor: None,
407                },
408                scoring: MemoryScoring {
409                    fallback_policy: FallbackPolicy::OnEmpty,
410                    ..Default::default()
411                },
412                current_avec: Some(AvecState::zero()),
413                query_text: Some("what did we decide about the parser grammar?".to_string()),
414                ..Default::default()
415            })
416            .await
417            .expect("recall should succeed");
418
419        assert_eq!(result.retrieval_path, RetrievalPath::LexicalFallback);
420        assert_eq!(result.nodes.len(), 1);
421        assert_eq!(result.nodes[0].sync_key, "far");
422    }
423
424    #[tokio::test]
425    async fn single_token_does_not_override_resonance_when_primary_is_non_empty() {
426        let store: Arc<dyn NodeStore> = Arc::new(InMemoryNodeStore::new());
427        store
428            .upsert_node_async(sample(
429                "near",
430                AvecState::zero(),
431                "weekend plans",
432                "alpha notes",
433            ))
434            .await
435            .expect("upsert near node");
436        store
437            .upsert_node_async(sample(
438                "far",
439                AvecState {
440                    stability: 1.0,
441                    friction: 1.0,
442                    logic: 0.0,
443                    autonomy: 0.0,
444                },
445                "hiking notes",
446                "bring boots",
447            ))
448            .await
449            .expect("upsert far node");
450
451        let service = MemoryRecallService::new(store);
452        let result = service
453            .execute(&MemoryRecallRequest {
454                page: MemoryPage {
455                    limit: 1,
456                    cursor: None,
457                },
458                scoring: MemoryScoring {
459                    fallback_policy: FallbackPolicy::OnEmpty,
460                    ..Default::default()
461                },
462                current_avec: Some(AvecState::zero()),
463                query_text: Some("hiking".to_string()),
464                ..Default::default()
465            })
466            .await
467            .expect("recall should succeed");
468
469        assert_eq!(result.retrieval_path, RetrievalPath::ResonanceOnly);
470        assert_eq!(result.nodes[0].sync_key, "near");
471    }
472
473    fn sample(sync_key: &str, avec: AvecState, summary: &str, raw: &str) -> SttpNode {
474        let now = Utc::now();
475        SttpNode {
476            raw: raw.to_string(),
477            session_id: "session".to_string(),
478            tier: "raw".to_string(),
479            timestamp: now,
480            compression_depth: 1,
481            parent_node_id: None,
482            sync_key: sync_key.to_string(),
483            updated_at: now,
484            source_metadata: None,
485            context_summary: Some(summary.to_string()),
486            semantic_tags: None,
487            semantic_links: None,
488            embedding_dimensions: None,
489            embedding_model: None,
490            embedding: None,
491            embedded_at: None,
492            user_avec: avec,
493            model_avec: avec,
494            compression_avec: Some(avec),
495            rho: 0.5,
496            kappa: 0.5,
497            psi: 1.0,
498        }
499    }
500}