1use std::collections::HashMap;
2use std::sync::Arc;
3
4use anyhow::Result;
5use locus_core_rs::domain::contracts::NodeStore;
6use locus_core_rs::domain::models::{AvecState, NodeQuery, SttpNode};
7
8use crate::application::memory_filters::{build_session_filter, node_matches_common_filters};
9use crate::domain::memory::{
10 MemoryAggregateGroup, MemoryAggregateRequest, MemoryAggregateResult, MemoryGroupBy, NumericStats,
11 clamp_groups, clamp_nodes,
12};
13
14pub struct MemoryAggregateService {
15 store: Arc<dyn NodeStore>,
16}
17
18impl MemoryAggregateService {
19 pub fn new(store: Arc<dyn NodeStore>) -> Self {
21 Self { store }
22 }
23
24 pub async fn execute(&self, request: &MemoryAggregateRequest) -> Result<MemoryAggregateResult> {
29 let max_nodes = clamp_nodes(if request.max_nodes == 0 {
30 5000
31 } else {
32 request.max_nodes
33 });
34 let max_groups = clamp_groups(if request.max_groups == 0 {
35 500
36 } else {
37 request.max_groups
38 });
39
40 let single_session = request
41 .scope
42 .session_ids
43 .as_deref()
44 .filter(|sessions| sessions.len() == 1)
45 .and_then(|sessions| sessions.first().cloned());
46
47 let nodes = self
48 .store
49 .query_nodes_async(NodeQuery {
50 limit: max_nodes,
51 session_id: single_session,
52 from_utc: request.scope.from_utc,
53 to_utc: request.scope.to_utc,
54 tiers: request.scope.tiers.clone(),
55 })
56 .await?;
57
58 let session_filter = build_session_filter(&request.scope);
59
60 let filtered = nodes
61 .into_iter()
62 .filter(|node| {
63 node_matches_common_filters(node, &request.scope, &request.filter, session_filter.as_ref())
64 })
65 .collect::<Vec<_>>();
66
67 let scanned_nodes = filtered.len();
68
69 let mut grouped: HashMap<String, Vec<SttpNode>> = HashMap::new();
70 for node in filtered {
71 if request.group_by == MemoryGroupBy::SemanticTag {
72 let tags = node.semantic_tags.as_deref().unwrap_or_default();
73 if tags.is_empty() {
74 grouped
75 .entry("untagged".to_string())
76 .or_default()
77 .push(node);
78 continue;
79 }
80
81 for tag in tags {
82 let key = tag.trim().to_ascii_lowercase();
83 if key.is_empty() {
84 continue;
85 }
86 grouped.entry(key).or_default().push(node.clone());
87 }
88 continue;
89 }
90
91 let key = group_key(&node, request.group_by);
92 grouped.entry(key).or_default().push(node);
93 }
94
95 let mut groups = grouped
96 .into_iter()
97 .map(|(key, nodes)| to_group(key, &nodes))
98 .collect::<Vec<_>>();
99
100 groups.sort_by(|left, right| {
101 right
102 .node_count
103 .cmp(&left.node_count)
104 .then_with(|| left.key.cmp(&right.key))
105 });
106
107 let total_groups = groups.len();
108 groups.truncate(max_groups);
109
110 Ok(MemoryAggregateResult {
111 groups,
112 total_groups,
113 scanned_nodes,
114 })
115 }
116}
117
118fn group_key(node: &SttpNode, group_by: MemoryGroupBy) -> String {
119 match group_by {
120 MemoryGroupBy::SessionId => node.session_id.clone(),
121 MemoryGroupBy::Tier => node.tier.clone(),
122 MemoryGroupBy::EmbeddingModel => node
123 .embedding_model
124 .clone()
125 .unwrap_or_else(|| "none".to_string()),
126 MemoryGroupBy::DateDay => node.timestamp.date_naive().to_string(),
127 MemoryGroupBy::SemanticTag => "semantic_tag".to_string(),
128 }
129}
130
131fn to_group(key: String, nodes: &[SttpNode]) -> MemoryAggregateGroup {
132 let node_count = nodes.len();
133
134 let embedding_count = nodes
135 .iter()
136 .filter(|node| node.embedding.as_ref().is_some_and(|values| !values.is_empty()))
137 .count();
138
139 let embedding_coverage = if node_count == 0 {
140 0.0
141 } else {
142 embedding_count as f32 / node_count as f32
143 };
144
145 let avg_user_avec = average_avec(nodes.iter().map(|node| node.user_avec).collect::<Vec<_>>().as_slice());
146 let avg_model_avec =
147 average_avec(nodes.iter().map(|node| node.model_avec).collect::<Vec<_>>().as_slice());
148
149 let compression_states = nodes
150 .iter()
151 .filter_map(|node| node.compression_avec)
152 .collect::<Vec<_>>();
153
154 let avg_compression_avec = if compression_states.is_empty() {
155 None
156 } else {
157 Some(average_avec(compression_states.as_slice()))
158 };
159
160 let psi_stats = average_metric(nodes.iter().map(|node| node.psi).collect::<Vec<_>>().as_slice());
161 let rho_stats = average_metric(nodes.iter().map(|node| node.rho).collect::<Vec<_>>().as_slice());
162 let kappa_stats =
163 average_metric(nodes.iter().map(|node| node.kappa).collect::<Vec<_>>().as_slice());
164
165 MemoryAggregateGroup {
166 key,
167 node_count,
168 embedding_coverage,
169 avg_user_avec,
170 avg_model_avec,
171 avg_compression_avec,
172 psi_stats,
173 rho_stats,
174 kappa_stats,
175 }
176}
177
178fn average_avec(values: &[AvecState]) -> AvecState {
179 if values.is_empty() {
180 return AvecState::zero();
181 }
182
183 let mut stability = 0.0_f32;
184 let mut friction = 0.0_f32;
185 let mut logic = 0.0_f32;
186 let mut autonomy = 0.0_f32;
187
188 for value in values {
189 stability += value.stability;
190 friction += value.friction;
191 logic += value.logic;
192 autonomy += value.autonomy;
193 }
194
195 let count = values.len() as f32;
196
197 AvecState {
198 stability: stability / count,
199 friction: friction / count,
200 logic: logic / count,
201 autonomy: autonomy / count,
202 }
203}
204
205fn average_metric(values: &[f32]) -> NumericStats {
206 if values.is_empty() {
207 return NumericStats::default();
208 }
209
210 let (min, max, sum) = values.iter().fold(
211 (f32::MAX, f32::MIN, 0.0_f32),
212 |(min, max, sum), value| (min.min(*value), max.max(*value), sum + *value),
213 );
214
215 NumericStats {
216 min,
217 max,
218 average: sum / values.len() as f32,
219 }
220}
221
222#[cfg(test)]
223mod tests {
224 use std::sync::Arc;
225
226 use chrono::{Duration, Utc};
227 use locus_core_rs::{InMemoryNodeStore, NodeStore};
228 use locus_core_rs::domain::models::{AvecState, SttpNode};
229
230 use super::MemoryAggregateService;
231 use crate::domain::memory::{MemoryAggregateRequest, MemoryGroupBy};
232
233 #[tokio::test]
234 async fn aggregates_nodes_by_session_with_coverage() {
235 let store: Arc<dyn NodeStore> = Arc::new(InMemoryNodeStore::new());
236 let now = Utc::now();
237
238 store
239 .upsert_node_async(test_node("s-1", "raw", now - Duration::minutes(2), Some(vec![0.1, 0.2])))
240 .await
241 .expect("upsert should succeed");
242 store
243 .upsert_node_async(test_node("s-1", "raw", now - Duration::minutes(1), None))
244 .await
245 .expect("upsert should succeed");
246 store
247 .upsert_node_async(test_node("s-2", "raw", now, Some(vec![0.3, 0.4])))
248 .await
249 .expect("upsert should succeed");
250
251 let service = MemoryAggregateService::new(store);
252 let request = MemoryAggregateRequest {
253 group_by: MemoryGroupBy::SessionId,
254 max_groups: 10,
255 max_nodes: 100,
256 ..Default::default()
257 };
258
259 let result = service.execute(&request).await.expect("aggregate should succeed");
260
261 assert_eq!(result.total_groups, 2);
262 let s1 = result
263 .groups
264 .iter()
265 .find(|group| group.key == "s-1")
266 .expect("s-1 group should exist");
267 assert_eq!(s1.node_count, 2);
268 assert!((s1.embedding_coverage - 0.5).abs() < f32::EPSILON);
269 }
270
271 fn test_node(session_id: &str, tier: &str, timestamp: chrono::DateTime<Utc>, embedding: Option<Vec<f32>>) -> SttpNode {
272 let user = AvecState {
273 stability: 0.6,
274 friction: 0.4,
275 logic: 0.8,
276 autonomy: 0.7,
277 };
278 let model = AvecState {
279 stability: 0.5,
280 friction: 0.3,
281 logic: 0.9,
282 autonomy: 0.6,
283 };
284
285 SttpNode {
286 raw: format!("raw:{session_id}:{tier}:{timestamp}"),
287 session_id: session_id.to_string(),
288 tier: tier.to_string(),
289 timestamp,
290 compression_depth: 1,
291 parent_node_id: None,
292 sync_key: format!("{}:{}:{}", session_id, tier, timestamp.timestamp_nanos_opt().unwrap_or_default()),
293 updated_at: timestamp,
294 source_metadata: None,
295 context_summary: Some("summary".to_string()),
296 semantic_tags: None,
297 semantic_links: None,
298 embedding_dimensions: embedding.as_ref().map(|v| v.len()),
299 embedding_model: embedding.as_ref().map(|_| "test-model".to_string()),
300 embedding,
301 embedded_at: None,
302 user_avec: user,
303 model_avec: model,
304 compression_avec: Some(model),
305 rho: 0.9,
306 kappa: 0.8,
307 psi: 2.5,
308 }
309 }
310}