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