Skip to main content

stasis/infrastructure/memory/
locus_context_reader.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use locus_sdk::application::memory_graph::MemoryGraphService;
5use locus_sdk::prelude::{
6    MemoryExplainRequest, MemoryExplainService, MemoryFindRequest as LocusFindRequest,
7    MemoryFindService, MemoryRecallRequest as LocusRecallRequest, MemoryRecallService,
8    MemorySort, MemoryPage,
9};
10use locus_sdk::domain::graph::MemoryGraphRequest as LocusGraphRequest;
11
12use crate::domain::errors::{Result, StasisError};
13use crate::infrastructure::memory::locus_memory_mapping::{
14    map_filter, map_node, map_scope, map_scoring, map_sort_direction, map_sort_field,
15};
16use crate::infrastructure::memory::locus_node_store_factory::LocusMemoryStore;
17use crate::ports::outbound::memory::memory_context_reader::MemoryContextReader;
18use crate::ports::outbound::memory::memory_models::{
19    MemoryFindRequest, MemoryFindResponse, MemoryGraphRequest, MemoryGraphResponse,
20    MemoryRecallRequest, MemoryRecallResponse,
21};
22
23pub struct LocusContextReader {
24    recall: MemoryRecallService,
25    find: MemoryFindService,
26    explain: MemoryExplainService,
27    graph: MemoryGraphService,
28}
29
30impl LocusContextReader {
31    pub fn new(memory: Arc<LocusMemoryStore>) -> Self {
32        let node_store = memory.node_store.clone();
33        let semantic_index = memory.semantic_index.clone();
34        Self {
35            recall: MemoryRecallService::new(node_store.clone())
36                .with_semantic_index(semantic_index.clone()),
37            find: MemoryFindService::new(node_store.clone())
38                .with_semantic_index(semantic_index.clone()),
39            explain: MemoryExplainService::new(node_store.clone()),
40            graph: MemoryGraphService::new(node_store).with_semantic_index(semantic_index),
41        }
42    }
43}
44
45#[async_trait]
46impl MemoryContextReader for LocusContextReader {
47    async fn recall(&self, request: &MemoryRecallRequest) -> Result<MemoryRecallResponse> {
48        let locus_request = LocusRecallRequest {
49            scope: map_scope(&request.scope),
50            filter: map_filter(&request.filter),
51            scoring: map_scoring(
52                request.alpha,
53                request.beta,
54                request.gamma,
55                request.fallback_policy,
56                request.strictness,
57            ),
58            page: MemoryPage {
59                limit: request.limit,
60                cursor: None,
61            },
62            current_avec: request.current_avec.map(|avec| locus_core_rs::domain::models::AvecState {
63                stability: avec.stability,
64                friction: avec.friction,
65                logic: avec.logic,
66                autonomy: avec.autonomy,
67            }),
68            query_text: request.query_text.clone(),
69            ..Default::default()
70        };
71
72        let recall_result = self
73            .recall
74            .execute(&locus_request)
75            .await
76            .map_err(|e| StasisError::PortFailure(format!("locus recall failed: {e}")))?;
77
78        let nodes: Vec<_> = recall_result.nodes.iter().map(map_node).collect();
79        let node_sync_keys: Vec<String> = nodes.iter().map(|node| node.sync_key.clone()).collect();
80
81        let mut response = MemoryRecallResponse {
82            retrieved: recall_result.retrieved,
83            next_cursor: recall_result.next_cursor,
84            has_more: recall_result.has_more,
85            retrieval_path: Some(format!("{:?}", recall_result.retrieval_path)),
86            nodes,
87            node_sync_keys,
88            ..Default::default()
89        };
90
91        if request.include_explain {
92            let explain_result = self
93                .explain
94                .execute(&MemoryExplainRequest {
95                    recall: locus_request,
96                })
97                .await
98                .map_err(|e| StasisError::PortFailure(format!("locus explain failed: {e}")))?;
99
100            response.fallback_triggered = explain_result.fallback_triggered;
101            response.fallback_reason = explain_result.fallback_reason;
102        }
103
104        Ok(response)
105    }
106
107    async fn find(&self, request: &MemoryFindRequest) -> Result<MemoryFindResponse> {
108        let locus_request = LocusFindRequest {
109            scope: map_scope(&request.scope),
110            filter: map_filter(&request.filter),
111            page: MemoryPage {
112                limit: request.limit,
113                cursor: request.cursor.clone(),
114            },
115            sort: MemorySort {
116                field: map_sort_field(request.sort_field),
117                direction: map_sort_direction(request.sort_direction),
118            },
119        };
120
121        let find_result = self
122            .find
123            .execute(&locus_request)
124            .await
125            .map_err(|e| StasisError::PortFailure(format!("locus find failed: {e}")))?;
126
127        let nodes: Vec<_> = find_result.nodes.iter().map(map_node).collect();
128        let node_sync_keys: Vec<String> = nodes.iter().map(|node| node.sync_key.clone()).collect();
129
130        Ok(MemoryFindResponse {
131            retrieved: find_result.retrieved,
132            has_more: find_result.has_more,
133            next_cursor: find_result.next_cursor,
134            nodes,
135            node_sync_keys,
136        })
137    }
138
139    async fn graph(&self, request: &MemoryGraphRequest) -> Result<MemoryGraphResponse> {
140        let locus_request = LocusGraphRequest {
141            scope: map_scope(&request.scope),
142            filter: map_filter(&request.filter),
143            include_lineage: request.include_lineage,
144            include_semantic: request.include_semantic,
145            include_session_topology: request.include_session_topology,
146            rel: request.rel.clone(),
147            target_prefix: request.target_prefix.clone(),
148            limit: request.limit,
149        };
150
151        let result = self
152            .graph
153            .execute(&locus_request)
154            .await
155            .map_err(|e| StasisError::PortFailure(format!("locus graph failed: {e}")))?;
156
157        Ok(MemoryGraphResponse {
158            sessions: result.sessions,
159            nodes: result.nodes,
160            edges: result.edges,
161            retrieved: result.retrieved,
162        })
163    }
164}