stasis/infrastructure/memory/
locus_context_reader.rs1use 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}