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::{AvecState, PsiRange, SemanticTagQueryFilter, SttpNode};
8use locus_core_rs::storage::derive_tenant_id_from_session;
9
10use crate::application::memory_filters::{
11    build_session_filter, node_matches_common_filters, resolve_indexed_sync_keys,
12};
13use crate::domain::memory::{
14    FallbackPolicy, MemoryRecallRequest, MemoryRecallResult, RetrievalPath, clamp_limit,
15};
16
17pub struct MemoryRecallService {
18    context_query: ContextQueryService,
19    semantic_index: Option<Arc<dyn SemanticIndexStore>>,
20}
21
22impl MemoryRecallService {
23    /// Create a recall service backed by the core resonance query pipeline.
24    pub fn new(store: Arc<dyn NodeStore>) -> Self {
25        Self {
26            context_query: ContextQueryService::new(store),
27            semantic_index: None,
28        }
29    }
30
31    pub fn with_semantic_index(
32        mut self,
33        semantic_index: Arc<dyn SemanticIndexStore>,
34    ) -> Self {
35        self.semantic_index = Some(semantic_index);
36        self
37    }
38
39    /// Retrieve context nodes using resonance or hybrid scoring,
40    /// with optional lexical fallback when configured.
41    pub async fn execute(&self, request: &MemoryRecallRequest) -> Result<MemoryRecallResult> {
42        let limit = clamp_limit(request.page.limit);
43        let expanded_limit = (limit.saturating_mul(5)).clamp(1, 200);
44
45        let current = request.current_avec.unwrap_or_else(AvecState::zero);
46        let session_scope = request
47            .scope
48            .session_ids
49            .as_deref()
50            .filter(|sessions| sessions.len() == 1)
51            .and_then(|sessions| sessions.first().map(String::as_str));
52        let session_filter = build_session_filter(&request.scope);
53        let tenant_id = request
54            .scope
55            .tenant_id
56            .clone()
57            .or_else(|| session_scope.map(derive_tenant_id_from_session))
58            .unwrap_or_else(|| "default".to_string());
59
60        let indexed_sync_keys = if let Some(index) = self.semantic_index.as_ref() {
61            resolve_indexed_sync_keys(
62                index.as_ref(),
63                &tenant_id,
64                &request.filter,
65                session_scope,
66                expanded_limit,
67            )
68            .await?
69        } else {
70            None
71        };
72
73        let mut path = if request.query_embedding.is_some() {
74            RetrievalPath::Hybrid
75        } else {
76            RetrievalPath::ResonanceOnly
77        };
78
79        let primary = if let Some(query_embedding) = request.query_embedding.as_deref() {
80            self.context_query
81                .get_context_hybrid_scoped_filtered_async(
82                    session_scope,
83                    current.stability,
84                    current.friction,
85                    current.logic,
86                    current.autonomy,
87                    request.scope.from_utc,
88                    request.scope.to_utc,
89                    request.scope.tiers.as_deref(),
90                    Some(query_embedding),
91                    request.scoring.alpha,
92                    request.scoring.beta,
93                    expanded_limit,
94                )
95                .await
96        } else {
97            self.context_query
98                .get_context_scoped_filtered_async(
99                    session_scope,
100                    current.stability,
101                    current.friction,
102                    current.logic,
103                    current.autonomy,
104                    request.scope.from_utc,
105                    request.scope.to_utc,
106                    request.scope.tiers.as_deref(),
107                    expanded_limit,
108                )
109                .await
110        };
111
112        let mut nodes = filter_nodes(
113            primary.nodes,
114            request,
115            session_filter.as_ref(),
116            indexed_sync_keys.as_ref(),
117        );
118
119        if let Some(query_text) = request.query_text.as_deref() {
120            let need_fallback = match request.scoring.fallback_policy {
121                FallbackPolicy::Never => false,
122                FallbackPolicy::OnEmpty => nodes.is_empty(),
123                FallbackPolicy::Always => true,
124            };
125
126            if need_fallback {
127                let fallback_result = self
128                    .context_query
129                    .get_context_scoped_filtered_async(
130                        session_scope,
131                        current.stability,
132                        current.friction,
133                        current.logic,
134                        current.autonomy,
135                        request.scope.from_utc,
136                        request.scope.to_utc,
137                        request.scope.tiers.as_deref(),
138                        expanded_limit,
139                    )
140                    .await;
141
142                let lexical = lexical_filter(
143                    filter_nodes(
144                        fallback_result.nodes,
145                        request,
146                        session_filter.as_ref(),
147                        indexed_sync_keys.as_ref(),
148                    ),
149                    query_text,
150                );
151
152                if request.scoring.fallback_policy == FallbackPolicy::Always && !nodes.is_empty() {
153                    nodes = merge_unique(nodes, lexical);
154                } else {
155                    nodes = lexical;
156                }
157
158                path = RetrievalPath::LexicalFallback;
159            }
160        }
161
162        if request.scoring.gamma > 0.0
163            && let Some(query_tag_embedding) = request.query_tag_embedding.as_deref()
164            && let Some(index) = self.semantic_index.as_ref()
165        {
166            rerank_by_tag_similarity(
167                &mut nodes,
168                index.as_ref(),
169                &tenant_id,
170                query_tag_embedding,
171                request.scoring.gamma,
172            )
173            .await?;
174        }
175
176        let has_more = nodes.len() > limit;
177        nodes.truncate(limit);
178
179        let next_cursor = nodes
180            .last()
181            .map(|node| format!("{}|{}", node.updated_at.to_rfc3339(), node.sync_key));
182
183        let psi_range = psi_range_from_nodes(&nodes);
184
185        Ok(MemoryRecallResult {
186            retrieved: nodes.len(),
187            nodes,
188            psi_range,
189            retrieval_path: path,
190            has_more,
191            next_cursor,
192        })
193    }
194}
195
196async fn rerank_by_tag_similarity(
197    nodes: &mut Vec<SttpNode>,
198    index: &dyn SemanticIndexStore,
199    tenant_id: &str,
200    query_embedding: &[f32],
201    gamma: f32,
202) -> Result<()> {
203    if nodes.is_empty() {
204        return Ok(());
205    }
206
207    let sync_keys: Vec<String> = nodes.iter().map(|node| node.sync_key.clone()).collect();
208    let records = index
209        .query_tag_records_async(SemanticTagQueryFilter {
210            tenant_id: Some(tenant_id.to_string()),
211            tags: None,
212            tag_prefix: None,
213            has_embedding: Some(true),
214            missing_embedding_only: false,
215            limit: sync_keys.len().saturating_mul(16).max(64),
216            session_id: None,
217        })
218        .await?;
219
220    let mut scores: Vec<(usize, f32)> = nodes
221        .iter()
222        .enumerate()
223        .map(|(index, node)| {
224            let tag_score = records
225                .iter()
226                .filter(|record| record.sync_key == node.sync_key)
227                .filter_map(|record| record.embedding.as_deref())
228                .filter_map(|embedding| cosine_similarity(query_embedding, embedding))
229                .fold(0.0_f32, f32::max);
230            (index, tag_score)
231        })
232        .collect();
233
234    scores.sort_by(|left, right| right.1.partial_cmp(&left.1).unwrap_or(std::cmp::Ordering::Equal));
235
236    let mut reranked = Vec::with_capacity(nodes.len());
237    let mut used = HashSet::new();
238    for (index, _) in scores {
239        if used.insert(index) {
240            reranked.push(nodes[index].clone());
241        }
242    }
243
244    if gamma >= 1.0 {
245        *nodes = reranked;
246    } else {
247        let blend_count = ((nodes.len() as f32) * gamma).ceil() as usize;
248        for (slot, node) in reranked.into_iter().take(blend_count).enumerate() {
249            nodes[slot] = node;
250        }
251    }
252
253    Ok(())
254}
255
256fn cosine_similarity(left: &[f32], right: &[f32]) -> Option<f32> {
257    if left.len() != right.len() || left.is_empty() {
258        return None;
259    }
260
261    let mut dot = 0.0_f32;
262    let mut left_norm = 0.0_f32;
263    let mut right_norm = 0.0_f32;
264
265    for (left_value, right_value) in left.iter().zip(right.iter()) {
266        dot += left_value * right_value;
267        left_norm += left_value * left_value;
268        right_norm += right_value * right_value;
269    }
270
271    if left_norm == 0.0 || right_norm == 0.0 {
272        return None;
273    }
274
275    Some(dot / (left_norm.sqrt() * right_norm.sqrt()))
276}
277
278fn filter_nodes(
279    nodes: Vec<SttpNode>,
280    request: &MemoryRecallRequest,
281    session_filter: Option<&HashSet<String>>,
282    indexed_sync_keys: Option<&HashSet<String>>,
283) -> Vec<SttpNode> {
284    nodes.into_iter()
285        .filter(|node| {
286            if let Some(keys) = indexed_sync_keys
287                && !keys.contains(&node.sync_key)
288            {
289                return false;
290            }
291
292            node_matches_common_filters(node, &request.scope, &request.filter, session_filter)
293        })
294        .collect()
295}
296
297fn lexical_filter(nodes: Vec<SttpNode>, query_text: &str) -> Vec<SttpNode> {
298    let needle = query_text.trim().to_ascii_lowercase();
299    if needle.is_empty() {
300        return nodes;
301    }
302
303    let mut scored = nodes
304        .into_iter()
305        .filter_map(|node| {
306            let summary = node
307                .context_summary
308                .as_deref()
309                .unwrap_or_default()
310                .to_ascii_lowercase();
311            let session = node.session_id.to_ascii_lowercase();
312            let raw = node.raw.to_ascii_lowercase();
313
314            let mut score = 0usize;
315            if summary.contains(&needle) {
316                score += 3;
317            }
318            if session.contains(&needle) {
319                score += 2;
320            }
321            if raw.contains(&needle) {
322                score += 1;
323            }
324
325            if score > 0 {
326                Some((score, node.timestamp, node))
327            } else {
328                None
329            }
330        })
331        .collect::<Vec<_>>();
332
333    scored.sort_by(|left, right| right.0.cmp(&left.0).then_with(|| right.1.cmp(&left.1)));
334
335    scored.into_iter().map(|(_, _, node)| node).collect()
336}
337
338fn merge_unique(primary: Vec<SttpNode>, secondary: Vec<SttpNode>) -> Vec<SttpNode> {
339    let mut merged = Vec::with_capacity(primary.len() + secondary.len());
340    let mut seen = HashSet::new();
341
342    for node in primary.into_iter().chain(secondary.into_iter()) {
343        if seen.insert(node.sync_key.clone()) {
344            merged.push(node);
345        }
346    }
347
348    merged
349}
350
351fn psi_range_from_nodes(nodes: &[SttpNode]) -> PsiRange {
352    if nodes.is_empty() {
353        return PsiRange::default();
354    }
355
356    let (min, max, sum) = nodes
357        .iter()
358        .fold((f32::MAX, f32::MIN, 0.0_f32), |(min, max, sum), node| {
359            (min.min(node.psi), max.max(node.psi), sum + node.psi)
360        });
361
362    PsiRange {
363        min,
364        max,
365        average: sum / nodes.len() as f32,
366    }
367}