Skip to main content

stasis/infrastructure/memory/
locus_context_reader.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use locus_core_rs::NodeStore;
5use locus_core_rs::domain::models::{AvecState, SttpNode};
6use locus_sdk::prelude::{
7    FallbackPolicy, MemoryExplainRequest, MemoryExplainService,
8    MemoryFindRequest as LocusFindRequest, MemoryFindService, MemoryFilter as LocusFilter,
9    MemoryRecallRequest as LocusRecallRequest, MemoryRecallService, MemoryScoring, MemorySort,
10    MemorySortField as LocusSortField, StrictnessMode, SortDirection as LocusSortDirection,
11};
12use locus_sdk::domain::memory::MetricRange as LocusMetricRange;
13
14use crate::domain::errors::{Result, StasisError};
15use crate::ports::outbound::memory::memory_context_reader::MemoryContextReader;
16use crate::ports::outbound::memory::memory_models::{
17    MemoryAvecState, MemoryFallbackPolicy, MemoryFilter, MemoryFindRequest, MemoryFindResponse,
18    MemoryMetricRange, MemoryNode, MemoryRecallRequest, MemoryRecallResponse, MemorySortDirection,
19    MemorySortField, MemoryStrictnessMode,
20};
21
22pub struct LocusContextReader {
23    recall: MemoryRecallService,
24    find: MemoryFindService,
25    explain: MemoryExplainService,
26}
27
28impl LocusContextReader {
29    pub fn new(store: Arc<dyn NodeStore>) -> Self {
30        Self {
31            recall: MemoryRecallService::new(store.clone()),
32            find: MemoryFindService::new(store.clone()),
33            explain: MemoryExplainService::new(store),
34        }
35    }
36}
37
38#[async_trait]
39impl MemoryContextReader for LocusContextReader {
40    async fn recall(&self, request: &MemoryRecallRequest) -> Result<MemoryRecallResponse> {
41        let locus_request = LocusRecallRequest {
42            scope: locus_sdk::prelude::MemoryScope {
43                session_ids: request.scope.session_ids.clone(),
44                tiers: request.scope.tiers.clone(),
45                from_utc: request.scope.from_utc,
46                to_utc: request.scope.to_utc,
47                ..Default::default()
48            },
49            scoring: MemoryScoring {
50                alpha: request.alpha,
51                beta: request.beta,
52                fallback_policy: map_fallback(request.fallback_policy),
53                strictness: map_strictness(request.strictness),
54                ..Default::default()
55            },
56            page: locus_sdk::prelude::MemoryPage {
57                limit: request.limit,
58                cursor: None,
59            },
60            current_avec: request.current_avec.map(|avec| AvecState {
61                stability: avec.stability,
62                friction: avec.friction,
63                logic: avec.logic,
64                autonomy: avec.autonomy,
65            }),
66            query_text: request.query_text.clone(),
67            ..Default::default()
68        };
69
70        let recall_result = self
71            .recall
72            .execute(&locus_request)
73            .await
74            .map_err(|e| StasisError::PortFailure(format!("locus recall failed: {e}")))?;
75
76        let nodes: Vec<MemoryNode> = recall_result
77            .nodes
78            .iter()
79            .map(map_node)
80            .collect();
81        let node_sync_keys: Vec<String> = nodes.iter().map(|node| node.sync_key.clone()).collect();
82
83        let mut response = MemoryRecallResponse {
84            retrieved: recall_result.retrieved,
85            next_cursor: recall_result.next_cursor,
86            has_more: recall_result.has_more,
87            retrieval_path: Some(format!("{:?}", recall_result.retrieval_path)),
88            nodes,
89            node_sync_keys,
90            ..Default::default()
91        };
92
93        if request.include_explain {
94            let explain_result = self
95                .explain
96                .execute(&MemoryExplainRequest {
97                    recall: locus_request,
98                })
99                .await
100                .map_err(|e| StasisError::PortFailure(format!("locus explain failed: {e}")))?;
101
102            response.fallback_triggered = explain_result.fallback_triggered;
103            response.fallback_reason = explain_result.fallback_reason;
104        }
105
106        Ok(response)
107    }
108
109    async fn find(&self, request: &MemoryFindRequest) -> Result<MemoryFindResponse> {
110        let locus_request = LocusFindRequest {
111            scope: locus_sdk::prelude::MemoryScope {
112                session_ids: request.scope.session_ids.clone(),
113                tiers: request.scope.tiers.clone(),
114                from_utc: request.scope.from_utc,
115                to_utc: request.scope.to_utc,
116                ..Default::default()
117            },
118            filter: map_filter(&request.filter),
119            page: locus_sdk::prelude::MemoryPage {
120                limit: request.limit,
121                cursor: request.cursor.clone(),
122            },
123            sort: MemorySort {
124                field: map_sort_field(request.sort_field),
125                direction: map_sort_direction(request.sort_direction),
126            },
127        };
128
129        let find_result = self
130            .find
131            .execute(&locus_request)
132            .await
133            .map_err(|e| StasisError::PortFailure(format!("locus find failed: {e}")))?;
134
135        let nodes: Vec<MemoryNode> = find_result.nodes.iter().map(map_node).collect();
136        let node_sync_keys: Vec<String> = nodes.iter().map(|node| node.sync_key.clone()).collect();
137
138        Ok(MemoryFindResponse {
139            retrieved: find_result.retrieved,
140            has_more: find_result.has_more,
141            next_cursor: find_result.next_cursor,
142            nodes,
143            node_sync_keys,
144        })
145    }
146}
147
148fn map_avec(avec: &AvecState) -> MemoryAvecState {
149    MemoryAvecState {
150        stability: avec.stability,
151        friction: avec.friction,
152        logic: avec.logic,
153        autonomy: avec.autonomy,
154    }
155}
156
157fn map_node(node: &SttpNode) -> MemoryNode {
158    MemoryNode {
159        raw: node.raw.clone(),
160        session_id: node.session_id.clone(),
161        tier: node.tier.clone(),
162        timestamp: node.timestamp,
163        compression_depth: node.compression_depth,
164        parent_node_id: node.parent_node_id.clone(),
165        sync_key: node.sync_key.clone(),
166        context_summary: node.context_summary.clone(),
167        embedding_model: node.embedding_model.clone(),
168        embedding_dimensions: node.embedding_dimensions,
169        embedded_at: node.embedded_at,
170        rho: node.rho,
171        kappa: node.kappa,
172        psi: node.psi,
173        user_avec: map_avec(&node.user_avec),
174        model_avec: map_avec(&node.model_avec),
175        compression_avec: node.compression_avec.as_ref().map(map_avec),
176        updated_at: node.updated_at,
177    }
178}
179
180fn map_fallback(value: MemoryFallbackPolicy) -> FallbackPolicy {
181    match value {
182        MemoryFallbackPolicy::Never => FallbackPolicy::Never,
183        MemoryFallbackPolicy::OnEmpty => FallbackPolicy::OnEmpty,
184        MemoryFallbackPolicy::Always => FallbackPolicy::Always,
185    }
186}
187
188fn map_strictness(value: MemoryStrictnessMode) -> StrictnessMode {
189    match value {
190        MemoryStrictnessMode::Precision => StrictnessMode::Precision,
191        MemoryStrictnessMode::Balanced => StrictnessMode::Balanced,
192        MemoryStrictnessMode::Recall => StrictnessMode::Recall,
193    }
194}
195
196fn map_metric_range(value: &MemoryMetricRange) -> LocusMetricRange {
197    LocusMetricRange {
198        min: value.min,
199        max: value.max,
200    }
201}
202
203fn map_filter(value: &MemoryFilter) -> LocusFilter {
204    LocusFilter {
205        has_embedding: value.has_embedding,
206        embedding_model: value.embedding_model.clone(),
207        psi: value.psi.as_ref().map(map_metric_range),
208        rho: value.rho.as_ref().map(map_metric_range),
209        kappa: value.kappa.as_ref().map(map_metric_range),
210        text_contains: value.text_contains.clone(),
211    }
212}
213
214fn map_sort_field(value: MemorySortField) -> LocusSortField {
215    match value {
216        MemorySortField::Timestamp => LocusSortField::Timestamp,
217        MemorySortField::UpdatedAt => LocusSortField::UpdatedAt,
218        MemorySortField::Psi => LocusSortField::Psi,
219        MemorySortField::Rho => LocusSortField::Rho,
220        MemorySortField::Kappa => LocusSortField::Kappa,
221    }
222}
223
224fn map_sort_direction(value: MemorySortDirection) -> LocusSortDirection {
225    match value {
226        MemorySortDirection::Asc => LocusSortDirection::Asc,
227        MemorySortDirection::Desc => LocusSortDirection::Desc,
228    }
229}