Skip to main content

stasis/infrastructure/memory/
locus_memory_operations.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use locus_core_rs::NodeStore;
5use locus_sdk::prelude::{
6    AiProviderRegistry, MemoryAggregateRequest as LocusAggregateRequest, MemoryAggregateService,
7    MemoryCompositionService, MemoryDailyRollupRequest, MemoryGroupBy, MemorySchemaService,
8    MemoryTransformOperation as LocusTransformOperation,
9    MemoryTransformRequest as LocusTransformRequest, MemoryTransformService,
10};
11
12use crate::domain::errors::{Result, StasisError};
13use crate::ports::outbound::memory::memory_models::{
14    MemoryAggregateRequest, MemoryAggregateResponse, MemoryRollupRequest, MemoryRollupResponse,
15    MemorySchemaResponse, MemoryTransformOperation, MemoryTransformRequest,
16    MemoryTransformResponse,
17};
18use crate::ports::outbound::memory::memory_operations::MemoryOperations;
19
20pub struct LocusMemoryOperations {
21    aggregate: MemoryAggregateService,
22    composition: MemoryCompositionService,
23    transform_store: Arc<dyn NodeStore>,
24    schema: MemorySchemaService,
25    providers: Option<Arc<dyn AiProviderRegistry>>,
26}
27
28impl LocusMemoryOperations {
29    pub fn new(store: Arc<dyn NodeStore>, providers: Option<Arc<dyn AiProviderRegistry>>) -> Self {
30        Self {
31            aggregate: MemoryAggregateService::new(store.clone()),
32            composition: MemoryCompositionService::new(store.clone()),
33            transform_store: store,
34            schema: MemorySchemaService::new(),
35            providers,
36        }
37    }
38}
39
40#[async_trait]
41impl MemoryOperations for LocusMemoryOperations {
42    async fn aggregate(&self, request: &MemoryAggregateRequest) -> Result<MemoryAggregateResponse> {
43        let result = self
44            .aggregate
45            .execute(&LocusAggregateRequest {
46                scope: locus_sdk::prelude::MemoryScope {
47                    session_ids: request.scope.session_ids.clone(),
48                    tiers: request.scope.tiers.clone(),
49                    from_utc: request.scope.from_utc,
50                    to_utc: request.scope.to_utc,
51                    ..Default::default()
52                },
53                group_by: MemoryGroupBy::DateDay,
54                max_groups: request.max_groups,
55                max_nodes: request.max_nodes,
56                ..Default::default()
57            })
58            .await
59            .map_err(|e| StasisError::PortFailure(format!("locus aggregate failed: {e}")))?;
60
61        Ok(MemoryAggregateResponse {
62            total_groups: result.total_groups,
63            scanned_nodes: result.scanned_nodes,
64        })
65    }
66
67    async fn transform(&self, request: &MemoryTransformRequest) -> Result<MemoryTransformResponse> {
68        let providers = self.providers.clone().ok_or_else(|| {
69            StasisError::PortFailure("locus transform requires ai provider registry".to_string())
70        })?;
71
72        let service = MemoryTransformService::new(self.transform_store.clone(), providers);
73        let result = service
74            .execute(&LocusTransformRequest {
75                scope: locus_sdk::prelude::MemoryScope {
76                    session_ids: request.scope.session_ids.clone(),
77                    tiers: request.scope.tiers.clone(),
78                    from_utc: request.scope.from_utc,
79                    to_utc: request.scope.to_utc,
80                    ..Default::default()
81                },
82                operation: map_transform_operation(request.operation),
83                dry_run: request.dry_run,
84                batch_size: request.batch_size,
85                max_nodes: request.max_nodes,
86                provider_id: request.provider_id.clone(),
87                model: request.model.clone(),
88                ..Default::default()
89            })
90            .await
91            .map_err(|e| StasisError::PortFailure(format!("locus transform failed: {e}")))?;
92
93        Ok(MemoryTransformResponse {
94            scanned: result.scanned,
95            selected: result.selected,
96            updated: result.updated,
97            skipped: result.skipped,
98            failed: result.failed,
99            duplicate: result.duplicate,
100            failures: result.failures,
101        })
102    }
103
104    async fn rollup(&self, request: &MemoryRollupRequest) -> Result<MemoryRollupResponse> {
105        let result = self
106            .composition
107            .daily_rollup(&MemoryDailyRollupRequest {
108                scope: locus_sdk::prelude::MemoryScope {
109                    session_ids: request.scope.session_ids.clone(),
110                    tiers: request.scope.tiers.clone(),
111                    from_utc: request.scope.from_utc,
112                    to_utc: request.scope.to_utc,
113                    ..Default::default()
114                },
115                max_days: request.max_days,
116                max_nodes: request.max_nodes,
117                ..Default::default()
118            })
119            .await
120            .map_err(|e| StasisError::PortFailure(format!("locus daily rollup failed: {e}")))?;
121
122        Ok(MemoryRollupResponse {
123            total_groups: result.total_groups,
124            scanned_nodes: result.scanned_nodes,
125        })
126    }
127
128    async fn schema(&self) -> Result<MemorySchemaResponse> {
129        let schema = self.schema.execute();
130        Ok(MemorySchemaResponse {
131            schema_version: schema.schema_version,
132            sort_fields: schema.sort_fields,
133            filter_fields: schema.filter_fields,
134            group_by_fields: schema.group_by_fields,
135            fallback_policies: schema.fallback_policies,
136            strictness_modes: schema.strictness_modes,
137            transform_operations: schema.transform_operations,
138        })
139    }
140}
141
142fn map_transform_operation(value: MemoryTransformOperation) -> LocusTransformOperation {
143    match value {
144        MemoryTransformOperation::EmbedBackfill => LocusTransformOperation::EmbedBackfill,
145        MemoryTransformOperation::ReindexEmbeddings => LocusTransformOperation::ReindexEmbeddings,
146    }
147}