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