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 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 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}