use std::borrow::Cow;
use std::collections::HashMap;
use std::sync::atomic::{AtomicU32, Ordering};
use parking_lot::Mutex;
use crate::error::CanoError;
use crate::observer::WorkflowObserver;
use crate::recovery::{CheckpointRow, CheckpointStore};
use crate::resource::{Resource, Resources};
use crate::store::MemoryStore;
use crate::task::{Task, TaskConfig, TaskResult};
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RecordedEvent {
StateEntered(String),
TaskStarted {
state: String,
},
TaskSucceeded {
state: String,
},
TaskFailed {
state: String,
error: String,
},
Retry {
state: String,
attempt: u32,
},
CircuitOpen {
state: String,
},
Checkpoint {
workflow_id: String,
sequence: u64,
},
Resume {
workflow_id: String,
sequence: u64,
},
Cancelled {
state: String,
},
}
#[derive(Default, Debug)]
pub struct RecordingObserver {
events: Mutex<Vec<RecordedEvent>>,
}
impl RecordingObserver {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn events(&self) -> Vec<RecordedEvent> {
self.events.lock().clone()
}
#[must_use]
pub fn states_entered(&self) -> Vec<String> {
self.events
.lock()
.iter()
.filter_map(|e| match e {
RecordedEvent::StateEntered(s) => Some(s.clone()),
_ => None,
})
.collect()
}
pub fn clear(&self) {
self.events.lock().clear();
}
pub fn take(&self) -> Vec<RecordedEvent> {
std::mem::take(&mut self.events.lock())
}
pub fn assert_completed_with(&self, final_state: &str) {
let last = self.states_entered().last().cloned();
assert_eq!(last.as_deref(), Some(final_state));
}
pub fn assert_path(&self, expected: &[&str]) {
let actual = self.states_entered();
let actual_refs: Vec<&str> = actual.iter().map(String::as_str).collect();
assert_eq!(actual_refs.as_slice(), expected);
}
pub fn assert_all_states_entered<TState: std::fmt::Debug>(
&self,
expected: &[TState],
) -> Result<(), Vec<String>> {
let entered: std::collections::HashSet<String> =
self.states_entered().into_iter().collect();
let mut missing: Vec<String> = Vec::new();
let mut seen_in_input: std::collections::HashSet<String> = std::collections::HashSet::new();
for state in expected {
let label = format!("{state:?}");
if seen_in_input.insert(label.clone()) && !entered.contains(&label) {
missing.push(label);
}
}
if missing.is_empty() {
Ok(())
} else {
Err(missing)
}
}
pub fn assert_registered_states_entered<TState, TResourceKey>(
&self,
workflow: &crate::workflow::Workflow<TState, TResourceKey>,
) -> Result<(), Vec<String>>
where
TState: Clone + std::fmt::Debug + std::hash::Hash + Eq + Send + Sync + 'static,
TResourceKey: std::hash::Hash + Eq + Send + Sync + 'static,
{
let registered: Vec<TState> = workflow.registered_states().cloned().collect();
self.assert_all_states_entered(®istered)
}
}
impl WorkflowObserver for RecordingObserver {
fn on_state_enter(&self, state: &str) {
self.events
.lock()
.push(RecordedEvent::StateEntered(state.into()));
}
fn on_task_start(&self, task_id: &str) {
self.events.lock().push(RecordedEvent::TaskStarted {
state: task_id.into(),
});
}
fn on_task_success(&self, task_id: &str) {
self.events.lock().push(RecordedEvent::TaskSucceeded {
state: task_id.into(),
});
}
fn on_task_failure(&self, task_id: &str, err: &CanoError) {
self.events.lock().push(RecordedEvent::TaskFailed {
state: task_id.into(),
error: err.to_string(),
});
}
fn on_retry(&self, task_id: &str, attempt: u32) {
self.events.lock().push(RecordedEvent::Retry {
state: task_id.into(),
attempt,
});
}
fn on_circuit_open(&self, task_id: &str) {
self.events.lock().push(RecordedEvent::CircuitOpen {
state: task_id.into(),
});
}
fn on_checkpoint(&self, workflow_id: &str, sequence: u64) {
self.events.lock().push(RecordedEvent::Checkpoint {
workflow_id: workflow_id.into(),
sequence,
});
}
fn on_resume(&self, workflow_id: &str, sequence: u64) {
self.events.lock().push(RecordedEvent::Resume {
workflow_id: workflow_id.into(),
sequence,
});
}
fn on_cancelled(&self, state: &str) {
self.events.lock().push(RecordedEvent::Cancelled {
state: state.into(),
});
}
}
#[derive(Default, Debug)]
pub struct InMemoryCheckpointStore {
inner: Mutex<HashMap<String, Vec<CheckpointRow>>>,
}
impl InMemoryCheckpointStore {
#[must_use]
pub fn new() -> Self {
Self::default()
}
}
#[crate::checkpoint_store]
impl CheckpointStore for InMemoryCheckpointStore {
async fn append(&self, workflow_id: &str, row: CheckpointRow) -> Result<(), CanoError> {
let mut runs = self.inner.lock();
let rows = runs.entry(workflow_id.to_string()).or_default();
if rows.iter().any(|r| r.sequence == row.sequence) {
return Err(CanoError::checkpoint_store(format!(
"checkpoint conflict: {workflow_id:?} already has sequence {}",
row.sequence
)));
}
rows.push(row);
Ok(())
}
async fn load_run(&self, workflow_id: &str) -> Result<Vec<CheckpointRow>, CanoError> {
let mut rows = self
.inner
.lock()
.get(workflow_id)
.cloned()
.unwrap_or_default();
rows.sort_by_key(|r| r.sequence);
Ok(rows)
}
async fn clear(&self, workflow_id: &str) -> Result<(), CanoError> {
self.inner.lock().remove(workflow_id);
Ok(())
}
}
#[derive(Default)]
pub struct TestResources {
inner: Resources,
}
impl TestResources {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_store(mut self, key: &'static str) -> Self {
self.inner = std::mem::take(&mut self.inner).insert(key, MemoryStore::new());
self
}
#[must_use]
pub fn with_resource<R: Resource + 'static>(mut self, key: &'static str, resource: R) -> Self {
self.inner = std::mem::take(&mut self.inner).insert(key, resource);
self
}
#[must_use]
pub fn build(self) -> Resources {
self.inner
}
}
pub struct PanicOnAttempt<TState> {
panics: u32,
next_state: TState,
attempts: AtomicU32,
config: TaskConfig,
}
impl<TState> PanicOnAttempt<TState> {
#[must_use]
pub fn with_config(mut self, config: TaskConfig) -> Self {
self.config = config;
self
}
#[must_use]
pub fn attempts(&self) -> u32 {
self.attempts.load(Ordering::SeqCst)
}
}
pub fn panic_on_attempt<TState>(panics: u32, next_state: TState) -> PanicOnAttempt<TState> {
PanicOnAttempt {
panics,
next_state,
attempts: AtomicU32::new(0),
config: TaskConfig::minimal(),
}
}
#[crate::task]
impl<TState> Task<TState> for PanicOnAttempt<TState>
where
TState: Clone + std::fmt::Debug + Send + Sync + 'static,
{
fn config(&self) -> TaskConfig {
self.config.clone()
}
fn name(&self) -> Cow<'static, str> {
"PanicOnAttempt".into()
}
async fn run_bare(&self) -> Result<TaskResult<TState>, CanoError> {
let attempt = self.attempts.fetch_add(1, Ordering::SeqCst) + 1;
if attempt <= self.panics {
panic!("panic_on_attempt: forced panic on attempt {attempt}");
}
Ok(TaskResult::Single(self.next_state.clone()))
}
}
pub fn assert_compensation_ran<S: AsRef<str>>(actual: &[S], expected: &[&str]) {
let actual_refs: Vec<&str> = actual.iter().map(AsRef::as_ref).collect();
assert_eq!(
actual_refs.as_slice(),
expected,
"compensation order mismatch:\n expected: {expected:?}\n actual: {actual_refs:?}"
);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
use std::sync::Arc;
use std::time::Duration;
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
enum S {
Start,
Done,
A,
B,
C,
}
struct OkTask;
#[task]
impl Task<S> for OkTask {
async fn run_bare(&self) -> Result<TaskResult<S>, CanoError> {
Ok(TaskResult::Single(S::Done))
}
fn name(&self) -> Cow<'static, str> {
"OkTask".into()
}
}
#[derive(Clone)]
struct Go(S);
#[task]
impl Task<S> for Go {
async fn run_bare(&self) -> Result<TaskResult<S>, CanoError> {
Ok(TaskResult::Single(self.0.clone()))
}
}
#[tokio::test]
async fn recording_observer_captures_success_path() {
let observer = Arc::new(RecordingObserver::default());
let wf = Workflow::bare()
.register(S::Start, OkTask)
.add_exit_state(S::Done)
.with_observer(observer.clone());
assert_eq!(
wf.orchestrate(S::Start, CancellationToken::disabled())
.await
.unwrap(),
S::Done
);
observer.assert_path(&["Start", "Done"]);
observer.assert_completed_with("Done");
assert!(observer.events().contains(&RecordedEvent::TaskSucceeded {
state: "OkTask".into()
}));
}
#[tokio::test]
async fn recording_observer_clear_resets_events() {
let observer = Arc::new(RecordingObserver::new());
let wf = Workflow::bare()
.register(S::Start, OkTask)
.add_exit_state(S::Done)
.with_observer(observer.clone());
wf.orchestrate(S::Start, CancellationToken::disabled())
.await
.unwrap();
assert!(!observer.events().is_empty());
observer.clear();
assert!(observer.events().is_empty());
}
#[tokio::test]
async fn in_memory_checkpoint_store_roundtrip() {
let store = InMemoryCheckpointStore::new();
store
.append("run", CheckpointRow::new(0, "A", "t0"))
.await
.unwrap();
store
.append("run", CheckpointRow::new(1, "B", "t1"))
.await
.unwrap();
let rows = store.load_run("run").await.unwrap();
assert_eq!(
rows.iter().map(|r| r.sequence).collect::<Vec<_>>(),
vec![0, 1]
);
}
#[tokio::test]
async fn in_memory_checkpoint_store_sorts_on_load() {
let store = InMemoryCheckpointStore::new();
store
.append("run", CheckpointRow::new(2, "C", "t2"))
.await
.unwrap();
store
.append("run", CheckpointRow::new(0, "A", "t0"))
.await
.unwrap();
let rows = store.load_run("run").await.unwrap();
assert_eq!(
rows.iter().map(|r| r.sequence).collect::<Vec<_>>(),
vec![0, 2]
);
}
#[tokio::test]
async fn in_memory_checkpoint_store_rejects_duplicate_sequence() {
let store = InMemoryCheckpointStore::new();
store
.append("run", CheckpointRow::new(0, "A", "t0"))
.await
.unwrap();
let err = store
.append("run", CheckpointRow::new(0, "A2", "t0"))
.await
.expect_err("duplicate sequence must be rejected");
assert_eq!(err.category(), "checkpoint_store");
}
#[tokio::test]
async fn in_memory_checkpoint_store_clear_is_per_id() {
let store = InMemoryCheckpointStore::new();
store
.append("a", CheckpointRow::new(0, "A", "t0"))
.await
.unwrap();
store
.append("b", CheckpointRow::new(0, "B", "t0"))
.await
.unwrap();
store.clear("a").await.unwrap();
assert!(store.load_run("a").await.unwrap().is_empty());
assert_eq!(store.load_run("b").await.unwrap().len(), 1);
}
#[test]
fn test_resources_with_store_inserts_memory_store() {
let resources = TestResources::new().with_store("store").build();
assert!(resources.get::<MemoryStore, _>("store").is_ok());
}
#[test]
fn test_resources_with_resource_inserts_custom() {
#[derive(Clone)]
struct Cfg {
limit: usize,
}
#[resource]
impl Resource for Cfg {}
let resources = TestResources::new()
.with_store("store")
.with_resource("cfg", Cfg { limit: 7 })
.build();
assert_eq!(resources.get::<Cfg, _>("cfg").unwrap().limit, 7);
assert!(resources.get::<MemoryStore, _>("store").is_ok());
}
#[tokio::test]
async fn panic_on_attempt_first_attempt_surfaces_as_error() {
let task = panic_on_attempt(1, S::Done);
let err = crate::workflow::catch_panic_to_error(
<PanicOnAttempt<S> as Task<S>>::run_bare(&task),
"Single task",
)
.await
.expect_err("expected the forced panic to surface as an error");
assert!(err.to_string().contains("panic"), "{err}");
}
#[tokio::test]
async fn panic_on_attempt_fails_fast_even_with_retries_configured() {
let observer = Arc::new(RecordingObserver::new());
let task = panic_on_attempt(2, S::Done)
.with_config(TaskConfig::new().with_fixed_retry(5, Duration::from_millis(1)));
let wf = Workflow::bare()
.register(S::Start, task)
.add_exit_state(S::Done)
.with_observer(observer.clone());
let err = wf
.orchestrate(S::Start, CancellationToken::disabled())
.await
.unwrap_err();
assert!(err.to_string().contains("panic"), "{err}");
let retries = observer
.events()
.into_iter()
.filter(|e| matches!(e, RecordedEvent::Retry { .. }))
.count();
assert_eq!(retries, 0, "panics are not retried");
assert!(
observer
.events()
.iter()
.any(|e| matches!(e, RecordedEvent::TaskFailed { .. }))
);
}
#[tokio::test]
async fn panic_on_attempt_zero_panics_transitions_immediately() {
let wf = Workflow::bare()
.register(S::Start, panic_on_attempt(0, S::Done))
.add_exit_state(S::Done);
assert_eq!(
wf.orchestrate(S::Start, CancellationToken::disabled())
.await
.unwrap(),
S::Done
);
}
#[test]
fn assert_compensation_ran_matches() {
let ran = vec!["charge".to_string(), "reserve".to_string()];
assert_compensation_ran(&ran, &["charge", "reserve"]);
}
#[test]
#[should_panic(expected = "compensation order mismatch")]
fn assert_compensation_ran_mismatch_panics() {
let ran = vec!["reserve".to_string()];
assert_compensation_ran(&ran, &["charge", "reserve"]);
}
#[tokio::test]
async fn assert_all_states_entered_passes_when_all_visited() {
let observer = Arc::new(RecordingObserver::new());
let wf = Workflow::bare()
.register(S::A, Go(S::B))
.register(S::B, Go(S::Done))
.add_exit_state(S::Done)
.with_observer(observer.clone());
wf.orchestrate(S::A, CancellationToken::disabled())
.await
.unwrap();
observer
.assert_all_states_entered(&[S::A, S::B, S::Done])
.expect("all states visited");
}
#[tokio::test]
async fn assert_all_states_entered_returns_missing_in_input_order() {
let observer = Arc::new(RecordingObserver::new());
let wf = Workflow::bare()
.register(S::A, Go(S::Done))
.add_exit_state(S::Done)
.with_observer(observer.clone());
wf.orchestrate(S::A, CancellationToken::disabled())
.await
.unwrap();
let missing = observer
.assert_all_states_entered(&[S::A, S::B, S::C, S::Done])
.unwrap_err();
assert_eq!(missing, vec!["B".to_string(), "C".to_string()]);
}
#[tokio::test]
async fn assert_all_states_entered_handles_duplicates_in_input() {
let observer = Arc::new(RecordingObserver::new());
let wf = Workflow::bare()
.register(S::A, Go(S::Done))
.add_exit_state(S::Done)
.with_observer(observer.clone());
wf.orchestrate(S::A, CancellationToken::disabled())
.await
.unwrap();
let missing = observer
.assert_all_states_entered(&[S::A, S::A, S::B])
.unwrap_err();
assert_eq!(missing, vec!["B".to_string()]);
}
#[tokio::test]
async fn assert_registered_states_entered_passes_for_full_path() {
let observer = Arc::new(RecordingObserver::new());
let wf = Workflow::bare()
.register(S::A, Go(S::B))
.register(S::B, Go(S::Done))
.add_exit_state(S::Done)
.with_observer(observer.clone());
wf.orchestrate(S::A, CancellationToken::disabled())
.await
.unwrap();
observer
.assert_registered_states_entered(&wf)
.expect("all registered states visited");
}
#[tokio::test]
async fn assert_registered_states_entered_reports_dead_state() {
let observer = Arc::new(RecordingObserver::new());
let wf = Workflow::bare()
.register(S::A, Go(S::Done))
.register(S::C, Go(S::Done)) .add_exit_state(S::Done)
.with_observer(observer.clone());
wf.orchestrate(S::A, CancellationToken::disabled())
.await
.unwrap();
let missing = observer.assert_registered_states_entered(&wf).unwrap_err();
assert!(missing.contains(&"C".to_string()), "missing={missing:?}");
}
}