Skip to main content

stasis/application/runtime/
memory_transform_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::{
7    MemoryTransformJobPayload, MemoryTransformOperationPayload,
8};
9use crate::application::runtime::in_memory_runtime::{JobExecutionOutcome, JobHandler};
10use crate::application::runtime::memory_job_request_helpers::memory_scope_from_fields;
11use crate::application::runtime::memory_operation_job_outcome_helpers::{
12    operation_failure, operation_success, policy_violation_failure,
13};
14use crate::domain::errors::Result;
15use crate::domain::runtime::job::Job;
16use crate::ports::outbound::memory::memory_models::{
17    MemoryTransformOperation, MemoryTransformRequest,
18};
19use crate::ports::outbound::memory::memory_operations::MemoryOperations;
20
21pub struct MemoryTransformJobHandler {
22    operations: Arc<dyn MemoryOperations>,
23}
24
25impl MemoryTransformJobHandler {
26    pub fn new(operations: Arc<dyn MemoryOperations>) -> Self {
27        Self { operations }
28    }
29
30    fn parse_payload(raw: &str) -> std::result::Result<MemoryTransformJobPayload, String> {
31        serde_json::from_str(raw).map_err(|err| {
32            format!("policy violation: invalid memory-transform payload json: {err}")
33        })
34    }
35
36    fn map_operation(value: Option<MemoryTransformOperationPayload>) -> MemoryTransformOperation {
37        match value {
38            Some(MemoryTransformOperationPayload::ReindexEmbeddings) => {
39                MemoryTransformOperation::ReindexEmbeddings
40            }
41            Some(MemoryTransformOperationPayload::EmbedTagBackfill) => {
42                MemoryTransformOperation::EmbedTagBackfill
43            }
44            Some(MemoryTransformOperationPayload::ReindexTagEmbeddings) => {
45                MemoryTransformOperation::ReindexTagEmbeddings
46            }
47            _ => MemoryTransformOperation::EmbedBackfill,
48        }
49    }
50}
51
52#[async_trait]
53impl JobHandler for MemoryTransformJobHandler {
54    fn job_type(&self) -> &'static str {
55        "workflow.stasis.memory.transform"
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-transform", message)),
62        };
63
64        let request = MemoryTransformRequest {
65            scope: memory_scope_from_fields(
66                None,
67                payload.session_ids,
68                payload.tiers,
69                payload.from_utc,
70                payload.to_utc,
71            ),
72            filter: Default::default(),
73            operation: Self::map_operation(payload.operation),
74            dry_run: payload.dry_run.unwrap_or(true),
75            batch_size: payload.batch_size.unwrap_or(100),
76            max_nodes: payload.max_nodes.unwrap_or(5000),
77            provider_id: payload.provider_id,
78            model: payload.model,
79        };
80
81        match self.operations.transform(&request).await {
82            Ok(result) => Ok(operation_success(
83                "stasis-memory-transform",
84                "memory-transform",
85                &job.id,
86                json!({
87                    "scanned": result.scanned,
88                    "selected": result.selected,
89                    "updated": result.updated,
90                    "failed": result.failed,
91                }),
92            )),
93            Err(err) => Ok(operation_failure("stasis-memory-transform", err.to_string())),
94        }
95    }
96}