Skip to main content

butterflow_state/
mock_adapter.rs

1use std::collections::HashMap;
2
3use uuid::Uuid;
4
5use butterflow_models::{Result, Task, WorkflowRun};
6
7use crate::StateAdapter;
8
9// Mock state adapter for testing
10pub struct MockStateAdapter {
11    workflow_runs: HashMap<Uuid, WorkflowRun>,
12    tasks: HashMap<Uuid, Task>,
13    state: HashMap<String, serde_json::Value>,
14}
15
16impl Default for MockStateAdapter {
17    fn default() -> Self {
18        Self::new()
19    }
20}
21
22impl MockStateAdapter {
23    pub fn new() -> Self {
24        Self {
25            workflow_runs: HashMap::new(),
26            tasks: HashMap::new(),
27            state: HashMap::new(),
28        }
29    }
30}
31
32#[async_trait::async_trait]
33impl StateAdapter for MockStateAdapter {
34    async fn save_workflow_run(&mut self, workflow_run: &WorkflowRun) -> Result<()> {
35        self.workflow_runs
36            .insert(workflow_run.id, workflow_run.clone());
37        Ok(())
38    }
39
40    async fn apply_workflow_run_diff(
41        &mut self,
42        diff: &butterflow_models::WorkflowRunDiff,
43    ) -> Result<()> {
44        let mut workflow_run = self.get_workflow_run(diff.workflow_run_id).await?;
45
46        for (field, field_diff) in &diff.fields {
47            match field_diff.operation {
48                butterflow_models::DiffOperation::Add
49                | butterflow_models::DiffOperation::Update
50                | butterflow_models::DiffOperation::Append => {
51                    if let Some(value) = &field_diff.value {
52                        let mut workflow_run_value = serde_json::to_value(&workflow_run)?;
53                        if let serde_json::Value::Object(obj) = &mut workflow_run_value {
54                            obj.insert(field.clone(), value.clone());
55                        }
56                        workflow_run = serde_json::from_value(workflow_run_value)?;
57                    }
58                }
59                butterflow_models::DiffOperation::Remove => {
60                    let mut workflow_run_value = serde_json::to_value(&workflow_run)?;
61                    if let serde_json::Value::Object(obj) = &mut workflow_run_value {
62                        obj.remove(field);
63                    }
64                    workflow_run = serde_json::from_value(workflow_run_value)?;
65                }
66            }
67        }
68
69        self.save_workflow_run(&workflow_run).await
70    }
71
72    async fn get_workflow_run(&self, workflow_run_id: Uuid) -> Result<WorkflowRun> {
73        self.workflow_runs
74            .get(&workflow_run_id)
75            .cloned()
76            .ok_or_else(|| {
77                butterflow_models::Error::Other(format!(
78                    "Workflow run {} not found",
79                    workflow_run_id
80                ))
81            })
82    }
83
84    async fn list_workflow_runs(&self, limit: usize) -> Result<Vec<WorkflowRun>> {
85        let mut runs: Vec<WorkflowRun> = self.workflow_runs.values().cloned().collect();
86        runs.sort_by(|a, b| b.started_at.cmp(&a.started_at));
87        Ok(runs.into_iter().take(limit).collect())
88    }
89
90    async fn save_task(&mut self, task: &Task) -> Result<()> {
91        self.tasks.insert(task.id, task.clone());
92        Ok(())
93    }
94
95    async fn apply_task_diff(&mut self, diff: &butterflow_models::TaskDiff) -> Result<()> {
96        let mut task = self.get_task(diff.task_id).await?;
97
98        for (field, field_diff) in &diff.fields {
99            match field_diff.operation {
100                butterflow_models::DiffOperation::Add
101                | butterflow_models::DiffOperation::Update
102                | butterflow_models::DiffOperation::Append => {
103                    if let Some(value) = &field_diff.value {
104                        let mut task_value = serde_json::to_value(&task)?;
105                        if let serde_json::Value::Object(obj) = &mut task_value {
106                            obj.insert(field.clone(), value.clone());
107                        }
108                        task = serde_json::from_value(task_value)?;
109                    }
110                }
111                butterflow_models::DiffOperation::Remove => {
112                    let mut task_value = serde_json::to_value(&task)?;
113                    if let serde_json::Value::Object(obj) = &mut task_value {
114                        obj.remove(field);
115                    }
116                    task = serde_json::from_value(task_value)?;
117                }
118            }
119        }
120
121        self.save_task(&task).await
122    }
123
124    async fn get_task(&self, task_id: Uuid) -> Result<Task> {
125        self.tasks
126            .get(&task_id)
127            .cloned()
128            .ok_or_else(|| butterflow_models::Error::Other(format!("Task {} not found", task_id)))
129    }
130
131    async fn get_tasks(&self, workflow_run_id: Uuid) -> Result<Vec<Task>> {
132        Ok(self
133            .tasks
134            .values()
135            .filter(|t| t.workflow_run_id == workflow_run_id)
136            .cloned()
137            .collect())
138    }
139
140    async fn update_state(
141        &mut self,
142        _workflow_run_id: Uuid,
143        state: HashMap<String, serde_json::Value>,
144    ) -> Result<()> {
145        self.state = state;
146        Ok(())
147    }
148
149    async fn apply_state_diff(&mut self, diff: &butterflow_models::StateDiff) -> Result<()> {
150        let mut state = self.get_state(diff.workflow_run_id).await?;
151
152        for (field, field_diff) in &diff.fields {
153            match field_diff.operation {
154                butterflow_models::DiffOperation::Add
155                | butterflow_models::DiffOperation::Update => {
156                    if let Some(value) = &field_diff.value {
157                        state.insert(field.clone(), value.clone());
158                    }
159                }
160                butterflow_models::DiffOperation::Remove => {
161                    state.remove(field);
162                }
163                butterflow_models::DiffOperation::Append => {
164                    if let Some(new_value) = &field_diff.value {
165                        if let Some(existing) = state.get_mut(field) {
166                            // If the existing value is an array, append to it
167                            if let serde_json::Value::Array(arr) = existing {
168                                arr.push(new_value.clone());
169                            } else {
170                                // If the existing value is not an array, replace it with a new array containing both values
171                                let old_value = existing.clone();
172                                *existing =
173                                    serde_json::Value::Array(vec![old_value, new_value.clone()]);
174                            }
175                        } else {
176                            // Field doesn't exist yet, create a new array with just this value
177                            state.insert(
178                                field.clone(),
179                                serde_json::Value::Array(vec![new_value.clone()]),
180                            );
181                        }
182                    }
183                }
184            }
185        }
186
187        self.update_state(diff.workflow_run_id, state).await
188    }
189
190    async fn get_state(
191        &self,
192        _workflow_run_id: Uuid,
193    ) -> Result<HashMap<String, serde_json::Value>> {
194        Ok(self.state.clone())
195    }
196}