Skip to main content

locus_sdk/application/
memory_aggregate.rs

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    /// Create an aggregate service over a shared node store.
20    pub fn new(store: Arc<dyn NodeStore>) -> Self {
21        Self { store }
22    }
23
24    /// Compute grouped aggregate statistics over filtered memory nodes.
25    ///
26    /// Returns group-level coverage and AVEC/metric summaries,
27    /// capped by request node/group limits.
28    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}