Skip to main content

stasis/application/runtime/
memory_find_job_handler.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use serde_json::json;
5
6use crate::application::orchestration::runtime_job_payloads::MemoryFindJobPayload;
7use crate::application::runtime::in_memory_runtime::{JobExecutionOutcome, JobHandler};
8use crate::application::runtime::memory_operation_job_outcome_helpers::{
9    operation_failure, operation_success, policy_violation_failure,
10};
11use crate::application::runtime::runtime_diagnostics_helpers::memory_nodes_json;
12use crate::domain::errors::Result;
13use crate::domain::runtime::job::Job;
14use crate::ports::outbound::memory::memory_context_reader::MemoryContextReader;
15use crate::ports::outbound::memory::memory_models::{
16    MemoryFilter, MemoryFindRequest, MemoryScope, MemorySortDirection, MemorySortField,
17};
18
19pub struct MemoryFindJobHandler {
20    reader: Arc<dyn MemoryContextReader>,
21}
22
23impl MemoryFindJobHandler {
24    pub fn new(reader: Arc<dyn MemoryContextReader>) -> Self {
25        Self { reader }
26    }
27
28    fn parse_payload(raw: &str) -> std::result::Result<MemoryFindJobPayload, String> {
29        serde_json::from_str(raw)
30            .map_err(|err| format!("policy violation: invalid memory-find payload json: {err}"))
31    }
32
33    fn map_sort_field(value: Option<&str>) -> MemorySortField {
34        match value {
35            Some("updated_at") => MemorySortField::UpdatedAt,
36            Some("psi") => MemorySortField::Psi,
37            Some("rho") => MemorySortField::Rho,
38            Some("kappa") => MemorySortField::Kappa,
39            _ => MemorySortField::Timestamp,
40        }
41    }
42
43    fn map_sort_direction(value: Option<&str>) -> MemorySortDirection {
44        match value {
45            Some("asc") => MemorySortDirection::Asc,
46            _ => MemorySortDirection::Desc,
47        }
48    }
49}
50
51#[async_trait]
52impl JobHandler for MemoryFindJobHandler {
53    fn job_type(&self) -> &'static str {
54        "workflow.stasis.memory.find"
55    }
56
57    async fn execute(&self, job: &Job) -> Result<JobExecutionOutcome> {
58        let payload = match Self::parse_payload(&job.payload_ref) {
59            Ok(payload) => payload,
60            Err(message) => return Ok(policy_violation_failure("stasis-memory-find", message)),
61        };
62
63        let request = MemoryFindRequest {
64            scope: MemoryScope {
65                session_ids: payload.session_ids,
66                tiers: payload.tiers,
67                from_utc: payload.from_utc,
68                to_utc: payload.to_utc,
69            },
70            filter: MemoryFilter {
71                text_contains: payload.text_contains,
72                ..Default::default()
73            },
74            limit: payload.limit.unwrap_or(50),
75            cursor: payload.cursor,
76            sort_field: Self::map_sort_field(payload.sort_field.as_deref()),
77            sort_direction: Self::map_sort_direction(payload.sort_direction.as_deref()),
78        };
79
80        match self.reader.find(&request).await {
81            Ok(result) => Ok(operation_success(
82                "stasis-memory-find",
83                "memory-find",
84                &job.id,
85                json!({
86                    "retrieved": result.retrieved,
87                    "has_more": result.has_more,
88                    "next_cursor": result.next_cursor,
89                    "node_sync_keys": result.node_sync_keys,
90                    "nodes": memory_nodes_json(&result.nodes),
91                }),
92            )),
93            Err(err) => Ok(operation_failure("stasis-memory-find", err.to_string())),
94        }
95    }
96}