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;
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    MemoryFallbackPolicy, MemoryFilter, MemoryFindRequest, MemoryFindResponse,
18    MemoryMetricRange, 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 mut response = MemoryRecallResponse {
77            retrieved: recall_result.retrieved,
78            next_cursor: recall_result.next_cursor,
79            has_more: recall_result.has_more,
80            retrieval_path: Some(format!("{:?}", recall_result.retrieval_path)),
81            node_sync_keys: recall_result
82                .nodes
83                .iter()
84                .map(|node| node.sync_key.clone())
85                .collect(),
86            ..Default::default()
87        };
88
89        if request.include_explain {
90            let explain_result = self
91                .explain
92                .execute(&MemoryExplainRequest {
93                    recall: locus_request,
94                })
95                .await
96                .map_err(|e| StasisError::PortFailure(format!("locus explain failed: {e}")))?;
97
98            response.fallback_triggered = explain_result.fallback_triggered;
99            response.fallback_reason = explain_result.fallback_reason;
100        }
101
102        Ok(response)
103    }
104
105    async fn find(&self, request: &MemoryFindRequest) -> Result<MemoryFindResponse> {
106        let locus_request = LocusFindRequest {
107            scope: locus_sdk::prelude::MemoryScope {
108                session_ids: request.scope.session_ids.clone(),
109                tiers: request.scope.tiers.clone(),
110                from_utc: request.scope.from_utc,
111                to_utc: request.scope.to_utc,
112                ..Default::default()
113            },
114            filter: map_filter(&request.filter),
115            page: locus_sdk::prelude::MemoryPage {
116                limit: request.limit,
117                cursor: request.cursor.clone(),
118            },
119            sort: MemorySort {
120                field: map_sort_field(request.sort_field),
121                direction: map_sort_direction(request.sort_direction),
122            },
123        };
124
125        let find_result = self
126            .find
127            .execute(&locus_request)
128            .await
129            .map_err(|e| StasisError::PortFailure(format!("locus find failed: {e}")))?;
130
131        Ok(MemoryFindResponse {
132            retrieved: find_result.retrieved,
133            has_more: find_result.has_more,
134            next_cursor: find_result.next_cursor,
135            node_sync_keys: find_result
136                .nodes
137                .iter()
138                .map(|node| node.sync_key.clone())
139                .collect(),
140        })
141    }
142}
143
144fn map_fallback(value: MemoryFallbackPolicy) -> FallbackPolicy {
145    match value {
146        MemoryFallbackPolicy::Never => FallbackPolicy::Never,
147        MemoryFallbackPolicy::OnEmpty => FallbackPolicy::OnEmpty,
148        MemoryFallbackPolicy::Always => FallbackPolicy::Always,
149    }
150}
151
152fn map_strictness(value: MemoryStrictnessMode) -> StrictnessMode {
153    match value {
154        MemoryStrictnessMode::Precision => StrictnessMode::Precision,
155        MemoryStrictnessMode::Balanced => StrictnessMode::Balanced,
156        MemoryStrictnessMode::Recall => StrictnessMode::Recall,
157    }
158}
159
160fn map_metric_range(value: &MemoryMetricRange) -> LocusMetricRange {
161    LocusMetricRange {
162        min: value.min,
163        max: value.max,
164    }
165}
166
167fn map_filter(value: &MemoryFilter) -> LocusFilter {
168    LocusFilter {
169        has_embedding: value.has_embedding,
170        embedding_model: value.embedding_model.clone(),
171        psi: value.psi.as_ref().map(map_metric_range),
172        rho: value.rho.as_ref().map(map_metric_range),
173        kappa: value.kappa.as_ref().map(map_metric_range),
174        text_contains: value.text_contains.clone(),
175    }
176}
177
178fn map_sort_field(value: MemorySortField) -> LocusSortField {
179    match value {
180        MemorySortField::Timestamp => LocusSortField::Timestamp,
181        MemorySortField::UpdatedAt => LocusSortField::UpdatedAt,
182        MemorySortField::Psi => LocusSortField::Psi,
183        MemorySortField::Rho => LocusSortField::Rho,
184        MemorySortField::Kappa => LocusSortField::Kappa,
185    }
186}
187
188fn map_sort_direction(value: MemorySortDirection) -> LocusSortDirection {
189    match value {
190        MemorySortDirection::Asc => LocusSortDirection::Asc,
191        MemorySortDirection::Desc => LocusSortDirection::Desc,
192    }
193}