1use std::collections::HashMap;
22use std::sync::{Arc, Mutex};
23
24use aion::{ActivityDispatch, ActivityDispatcher};
25use aion_core::WorkflowId;
26
27#[derive(Clone, Debug, Eq, PartialEq)]
29pub enum MockedActivity {
30 Succeeds {
33 result_json: String,
35 },
36 Fails {
39 message: String,
41 },
42}
43
44#[derive(Clone, Debug, Eq, Hash, PartialEq)]
46struct MockKey {
47 workflow_id: WorkflowId,
48 activity_name: String,
49}
50
51#[derive(Clone, Default)]
54pub struct ActivityMockRegistry {
55 mocks: Arc<Mutex<HashMap<MockKey, MockedActivity>>>,
56}
57
58impl std::fmt::Debug for ActivityMockRegistry {
59 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
60 formatter
61 .debug_struct("ActivityMockRegistry")
62 .finish_non_exhaustive()
63 }
64}
65
66impl ActivityMockRegistry {
67 #[must_use]
69 pub fn new() -> Self {
70 Self::default()
71 }
72
73 pub fn register(
81 &self,
82 workflow_id: WorkflowId,
83 activity_name: impl Into<String>,
84 mock: MockedActivity,
85 ) -> Result<(), String> {
86 let key = MockKey {
87 workflow_id,
88 activity_name: activity_name.into(),
89 };
90 self.mocks
91 .lock()
92 .map_err(|_| "activity mock registry mutex poisoned".to_owned())?
93 .insert(key, mock);
94 Ok(())
95 }
96
97 fn lookup(
103 &self,
104 workflow_id: &WorkflowId,
105 activity_name: &str,
106 ) -> Result<Option<MockedActivity>, String> {
107 let guard = self
108 .mocks
109 .lock()
110 .map_err(|_| "activity mock registry mutex poisoned".to_owned())?;
111 Ok(guard
113 .get(&MockKey {
114 workflow_id: workflow_id.clone(),
115 activity_name: activity_name.to_owned(),
116 })
117 .cloned())
118 }
119
120 pub fn has_any_for(&self, workflow_id: &WorkflowId) -> Result<bool, String> {
126 let guard = self
127 .mocks
128 .lock()
129 .map_err(|_| "activity mock registry mutex poisoned".to_owned())?;
130 Ok(guard.keys().any(|key| &key.workflow_id == workflow_id))
131 }
132}
133
134pub struct DevMockingDispatcher {
141 inner: Arc<dyn ActivityDispatcher>,
142 registry: ActivityMockRegistry,
143}
144
145impl std::fmt::Debug for DevMockingDispatcher {
146 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
147 formatter
148 .debug_struct("DevMockingDispatcher")
149 .field("registry", &self.registry)
150 .finish_non_exhaustive()
151 }
152}
153
154impl DevMockingDispatcher {
155 #[must_use]
157 pub fn new(inner: Arc<dyn ActivityDispatcher>, registry: ActivityMockRegistry) -> Self {
158 Self { inner, registry }
159 }
160}
161
162impl ActivityDispatcher for DevMockingDispatcher {
163 fn dispatch(&self, request: ActivityDispatch) -> Result<String, String> {
164 match self.registry.lookup(&request.workflow_id, &request.name)? {
165 Some(MockedActivity::Succeeds { result_json }) => {
166 tracing::info!(
167 operation = "dev.activity_mock",
168 workflow_id = %request.workflow_id,
169 activity_id = %request.activity_id,
170 activity_type = %request.name,
171 outcome = "succeeded",
172 "dev activity mock returned a canned result"
173 );
174 Ok(result_json)
175 }
176 Some(MockedActivity::Fails { message }) => {
177 tracing::info!(
178 operation = "dev.activity_mock",
179 workflow_id = %request.workflow_id,
180 activity_id = %request.activity_id,
181 activity_type = %request.name,
182 outcome = "failed",
183 "dev activity mock returned a canned failure"
184 );
185 Err(message)
186 }
187 None => self.inner.dispatch(request),
188 }
189 }
190}
191
192#[cfg(test)]
193mod tests {
194 use std::collections::BTreeMap;
195 use std::sync::Arc;
196
197 use aion::{ActivityDispatch, ActivityDispatcher};
198 use aion_core::{ActivityId, WorkflowId};
199
200 use super::{ActivityMockRegistry, DevMockingDispatcher, MockedActivity};
201
202 #[derive(Default)]
204 struct RecordingInner {
205 calls: std::sync::Mutex<Vec<String>>,
206 }
207
208 impl ActivityDispatcher for RecordingInner {
209 fn dispatch(&self, request: ActivityDispatch) -> Result<String, String> {
210 self.calls
211 .lock()
212 .map_err(|_| "poisoned".to_owned())?
213 .push(request.name.clone());
214 Ok(format!("real:{}", request.input))
215 }
216 }
217
218 fn dispatch(workflow_id: WorkflowId, name: &str) -> ActivityDispatch {
219 ActivityDispatch {
220 namespace: "default".to_owned(),
221 task_queue: "default".to_owned(),
222 node: None,
223 workflow_id,
224 activity_id: ActivityId::from_sequence_position(0),
225 name: name.to_owned(),
226 input: "{}".to_owned(),
227 config: "{}".to_owned(),
228 attempt: 1,
229 labels: BTreeMap::new(),
230 }
231 }
232
233 #[test]
234 fn mocked_activity_returns_canned_result_without_delegating() -> Result<(), String> {
235 let inner = Arc::new(RecordingInner::default());
236 let registry = ActivityMockRegistry::new();
237 let workflow_id = WorkflowId::new_v4();
238 registry.register(
239 workflow_id.clone(),
240 "charge-card",
241 MockedActivity::Succeeds {
242 result_json: r#"{"charged":true}"#.to_owned(),
243 },
244 )?;
245 let dispatcher = DevMockingDispatcher::new(inner.clone(), registry);
246
247 let result = dispatcher.dispatch(dispatch(workflow_id, "charge-card"));
248
249 assert_eq!(result, Ok(r#"{"charged":true}"#.to_owned()));
250 assert!(
251 inner
252 .calls
253 .lock()
254 .map_err(|_| "poisoned".to_owned())?
255 .is_empty(),
256 "a mocked activity must not reach the real dispatcher"
257 );
258 Ok(())
259 }
260
261 #[test]
262 fn mocked_failure_short_circuits_with_the_canned_message() -> Result<(), String> {
263 let inner = Arc::new(RecordingInner::default());
264 let registry = ActivityMockRegistry::new();
265 let workflow_id = WorkflowId::new_v4();
266 registry.register(
267 workflow_id.clone(),
268 "charge-card",
269 MockedActivity::Fails {
270 message: "card declined".to_owned(),
271 },
272 )?;
273 let dispatcher = DevMockingDispatcher::new(inner, registry);
274
275 assert_eq!(
276 dispatcher.dispatch(dispatch(workflow_id, "charge-card")),
277 Err("card declined".to_owned())
278 );
279 Ok(())
280 }
281
282 #[test]
283 fn unmocked_activity_delegates_to_the_real_dispatcher() -> Result<(), String> {
284 let inner = Arc::new(RecordingInner::default());
285 let registry = ActivityMockRegistry::new();
286 let dispatcher = DevMockingDispatcher::new(inner.clone(), registry);
287
288 let result = dispatcher.dispatch(dispatch(WorkflowId::new_v4(), "ship-order"));
289
290 assert_eq!(result, Ok("real:{}".to_owned()));
291 assert_eq!(
292 inner
293 .calls
294 .lock()
295 .map_err(|_| "poisoned".to_owned())?
296 .as_slice(),
297 ["ship-order"]
298 );
299 Ok(())
300 }
301
302 #[test]
303 fn mock_is_scoped_to_its_workflow_run() -> Result<(), String> {
304 let inner = Arc::new(RecordingInner::default());
305 let registry = ActivityMockRegistry::new();
306 let mocked = WorkflowId::new_v4();
307 let other = WorkflowId::new_v4();
308 registry.register(
309 mocked.clone(),
310 "charge-card",
311 MockedActivity::Succeeds {
312 result_json: r#"{"charged":true}"#.to_owned(),
313 },
314 )?;
315 let dispatcher = DevMockingDispatcher::new(inner.clone(), registry.clone());
316
317 assert_eq!(
319 dispatcher.dispatch(dispatch(other.clone(), "charge-card")),
320 Ok("real:{}".to_owned())
321 );
322 assert!(registry.has_any_for(&mocked)?);
323 assert!(!registry.has_any_for(&other)?);
324 Ok(())
325 }
326}