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