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;
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}