use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use aion::{ActivityDispatch, ActivityDispatcher};
use aion_core::WorkflowId;
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum MockedActivity {
Succeeds {
result_json: String,
},
Fails {
message: String,
},
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
struct MockKey {
workflow_id: WorkflowId,
activity_name: String,
}
#[derive(Clone, Default)]
pub struct ActivityMockRegistry {
mocks: Arc<Mutex<HashMap<MockKey, MockedActivity>>>,
}
impl std::fmt::Debug for ActivityMockRegistry {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ActivityMockRegistry")
.finish_non_exhaustive()
}
}
impl ActivityMockRegistry {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn register(
&self,
workflow_id: WorkflowId,
activity_name: impl Into<String>,
mock: MockedActivity,
) -> Result<(), String> {
let key = MockKey {
workflow_id,
activity_name: activity_name.into(),
};
self.mocks
.lock()
.map_err(|_| "activity mock registry mutex poisoned".to_owned())?
.insert(key, mock);
Ok(())
}
fn lookup(
&self,
workflow_id: &WorkflowId,
activity_name: &str,
) -> Result<Option<MockedActivity>, String> {
let guard = self
.mocks
.lock()
.map_err(|_| "activity mock registry mutex poisoned".to_owned())?;
Ok(guard
.get(&MockKey {
workflow_id: workflow_id.clone(),
activity_name: activity_name.to_owned(),
})
.cloned())
}
pub fn has_any_for(&self, workflow_id: &WorkflowId) -> Result<bool, String> {
let guard = self
.mocks
.lock()
.map_err(|_| "activity mock registry mutex poisoned".to_owned())?;
Ok(guard.keys().any(|key| &key.workflow_id == workflow_id))
}
}
pub struct DevMockingDispatcher {
inner: Arc<dyn ActivityDispatcher>,
registry: ActivityMockRegistry,
}
impl std::fmt::Debug for DevMockingDispatcher {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("DevMockingDispatcher")
.field("registry", &self.registry)
.finish_non_exhaustive()
}
}
impl DevMockingDispatcher {
#[must_use]
pub fn new(inner: Arc<dyn ActivityDispatcher>, registry: ActivityMockRegistry) -> Self {
Self { inner, registry }
}
}
impl ActivityDispatcher for DevMockingDispatcher {
fn dispatch(&self, request: ActivityDispatch) -> Result<String, String> {
match self.registry.lookup(&request.workflow_id, &request.name)? {
Some(MockedActivity::Succeeds { result_json }) => {
tracing::info!(
operation = "dev.activity_mock",
workflow_id = %request.workflow_id,
activity_id = %request.activity_id,
activity_type = %request.name,
outcome = "succeeded",
"dev activity mock returned a canned result"
);
Ok(result_json)
}
Some(MockedActivity::Fails { message }) => {
tracing::info!(
operation = "dev.activity_mock",
workflow_id = %request.workflow_id,
activity_id = %request.activity_id,
activity_type = %request.name,
outcome = "failed",
"dev activity mock returned a canned failure"
);
Err(message)
}
None => self.inner.dispatch(request),
}
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use std::sync::Arc;
use aion::{ActivityDispatch, ActivityDispatcher};
use aion_core::{ActivityId, WorkflowId};
use super::{ActivityMockRegistry, DevMockingDispatcher, MockedActivity};
#[derive(Default)]
struct RecordingInner {
calls: std::sync::Mutex<Vec<String>>,
}
impl ActivityDispatcher for RecordingInner {
fn dispatch(&self, request: ActivityDispatch) -> Result<String, String> {
self.calls
.lock()
.map_err(|_| "poisoned".to_owned())?
.push(request.name.clone());
Ok(format!("real:{}", request.input))
}
}
fn dispatch(workflow_id: WorkflowId, name: &str) -> ActivityDispatch {
ActivityDispatch {
namespace: "default".to_owned(),
task_queue: "default".to_owned(),
node: None,
workflow_id,
run_id: aion_core::RunId::new_v4(),
activity_id: ActivityId::from_sequence_position(0),
name: name.to_owned(),
input: "{}".to_owned(),
config: "{}".to_owned(),
attempt: 1,
labels: BTreeMap::new(),
advisory: false,
}
}
#[test]
fn mocked_activity_returns_canned_result_without_delegating() -> Result<(), String> {
let inner = Arc::new(RecordingInner::default());
let registry = ActivityMockRegistry::new();
let workflow_id = WorkflowId::new_v4();
registry.register(
workflow_id.clone(),
"charge-card",
MockedActivity::Succeeds {
result_json: r#"{"charged":true}"#.to_owned(),
},
)?;
let dispatcher = DevMockingDispatcher::new(inner.clone(), registry);
let result = dispatcher.dispatch(dispatch(workflow_id, "charge-card"));
assert_eq!(result, Ok(r#"{"charged":true}"#.to_owned()));
assert!(
inner
.calls
.lock()
.map_err(|_| "poisoned".to_owned())?
.is_empty(),
"a mocked activity must not reach the real dispatcher"
);
Ok(())
}
#[test]
fn mocked_failure_short_circuits_with_the_canned_message() -> Result<(), String> {
let inner = Arc::new(RecordingInner::default());
let registry = ActivityMockRegistry::new();
let workflow_id = WorkflowId::new_v4();
registry.register(
workflow_id.clone(),
"charge-card",
MockedActivity::Fails {
message: "card declined".to_owned(),
},
)?;
let dispatcher = DevMockingDispatcher::new(inner, registry);
assert_eq!(
dispatcher.dispatch(dispatch(workflow_id, "charge-card")),
Err("card declined".to_owned())
);
Ok(())
}
#[test]
fn unmocked_activity_delegates_to_the_real_dispatcher() -> Result<(), String> {
let inner = Arc::new(RecordingInner::default());
let registry = ActivityMockRegistry::new();
let dispatcher = DevMockingDispatcher::new(inner.clone(), registry);
let result = dispatcher.dispatch(dispatch(WorkflowId::new_v4(), "ship-order"));
assert_eq!(result, Ok("real:{}".to_owned()));
assert_eq!(
inner
.calls
.lock()
.map_err(|_| "poisoned".to_owned())?
.as_slice(),
["ship-order"]
);
Ok(())
}
#[test]
fn mock_is_scoped_to_its_workflow_run() -> Result<(), String> {
let inner = Arc::new(RecordingInner::default());
let registry = ActivityMockRegistry::new();
let mocked = WorkflowId::new_v4();
let other = WorkflowId::new_v4();
registry.register(
mocked.clone(),
"charge-card",
MockedActivity::Succeeds {
result_json: r#"{"charged":true}"#.to_owned(),
},
)?;
let dispatcher = DevMockingDispatcher::new(inner.clone(), registry.clone());
assert_eq!(
dispatcher.dispatch(dispatch(other.clone(), "charge-card")),
Ok("real:{}".to_owned())
);
assert!(registry.has_any_for(&mocked)?);
assert!(!registry.has_any_for(&other)?);
Ok(())
}
}