Skip to main content

stasis/application/runtime/
memory_recall_job_handler.rs

1use std::sync::Arc;
2use std::time::Instant;
3
4use async_trait::async_trait;
5use serde_json::json;
6
7use crate::application::orchestration::runtime_job_payloads::{
8    MemoryPolicyPayload, MemoryRecallJobPayload,
9};
10use crate::application::runtime::in_memory_runtime::{JobExecutionOutcome, JobHandler};
11use crate::application::runtime::memory_recall_request_builder::build_memory_recall_request;
12use crate::application::runtime::runtime_diagnostics_helpers::memory_nodes_json;
13use crate::application::telemetry::operation::OperationTelemetry;
14use crate::domain::errors::Result;
15use crate::domain::runtime::job::Job;
16use crate::ports::outbound::memory::memory_context_reader::MemoryContextReader;
17use crate::ports::outbound::memory::memory_models::MemoryRecallRequest;
18
19pub struct MemoryRecallJobHandler {
20    reader: Arc<dyn MemoryContextReader>,
21    telemetry: Option<OperationTelemetry>,
22}
23
24impl MemoryRecallJobHandler {
25    pub fn new(reader: Arc<dyn MemoryContextReader>) -> Self {
26        Self {
27            reader,
28            telemetry: None,
29        }
30    }
31
32    pub fn with_operation_telemetry(mut self, telemetry: Option<OperationTelemetry>) -> Self {
33        self.telemetry = telemetry;
34        self
35    }
36
37    fn parse_payload(raw: &str) -> std::result::Result<MemoryRecallJobPayload, String> {
38        serde_json::from_str(raw)
39            .map_err(|err| format!("policy violation: invalid memory-recall payload json: {err}"))
40    }
41
42    fn build_request(
43        correlation_id: &str,
44        policy: Option<&MemoryPolicyPayload>,
45    ) -> MemoryRecallRequest {
46        build_memory_recall_request(correlation_id, None, policy)
47    }
48}
49
50#[async_trait]
51impl JobHandler for MemoryRecallJobHandler {
52    fn job_type(&self) -> &'static str {
53        "workflow.stasis.memory.recall"
54    }
55
56    async fn execute(&self, job: &Job) -> Result<JobExecutionOutcome> {
57        let recall_started = Instant::now();
58        let _recall_span = self.telemetry.as_ref().map(|telemetry| {
59            telemetry.record_recall_started();
60            telemetry.recall_span(&job.correlation_id)
61        });
62
63        let payload = match Self::parse_payload(&job.payload_ref) {
64            Ok(payload) => payload,
65            Err(message) => {
66                let diagnostics = json!({
67                    "provider": "stasis-memory-recall",
68                    "status": "failure",
69                    "guardrail_code": "POLICY_VIOLATION",
70                    "policy_reason": message,
71                })
72                .to_string();
73                return Ok(JobExecutionOutcome::FatalFailure {
74                    message: "invalid memory recall payload".to_string(),
75                    execution_id: None,
76                    diagnostics: Some(diagnostics),
77                });
78            }
79        };
80
81        let recall_request =
82            Self::build_request(&job.correlation_id, payload.memory_policy.as_ref());
83        match self.reader.recall(&recall_request).await {
84            Ok(response) => {
85                if let Some(telemetry) = &self.telemetry {
86                    telemetry.record_recall_success(recall_started);
87                }
88                Ok(JobExecutionOutcome::Success {
89                    sttp_output_node_id: format!("sttp:memory-recall:{}", job.id),
90                    execution_id: None,
91                    diagnostics: Some(
92                        json!({
93                            "provider": "stasis-memory-recall",
94                            "status": "success",
95                            "retrieved": response.retrieved,
96                            "retrieval_path": response.retrieval_path,
97                            "fallback_triggered": response.fallback_triggered,
98                            "fallback_reason": response.fallback_reason,
99                            "has_more": response.has_more,
100                            "node_sync_keys": response.node_sync_keys,
101                            "nodes": memory_nodes_json(&response.nodes),
102                        })
103                        .to_string(),
104                    ),
105                })
106            }
107            Err(err) => {
108                if let Some(telemetry) = &self.telemetry {
109                    telemetry.record_recall_error(recall_started);
110                }
111                Ok(JobExecutionOutcome::FatalFailure {
112                    message: err.to_string(),
113                    execution_id: None,
114                    diagnostics: Some(
115                        json!({
116                            "provider": "stasis-memory-recall",
117                            "status": "failure",
118                            "error": err.to_string(),
119                        })
120                        .to_string(),
121                    ),
122                })
123            }
124        }
125    }
126}