use crate::graph::checkpoint::{Checkpoint, Checkpointer};
use crate::graph::orchestration::{
OrchestrationTaskFilter, OrchestrationTaskKind, OrchestrationTaskResult,
OrchestrationTaskStatus, TaskStore,
};
use crate::harness::ids::{NodeId, TaskId};
fn contract_checkpoint(
thread: &str,
id: &str,
parent: Option<&str>,
step: usize,
) -> Checkpoint<i32> {
Checkpoint {
thread_id: thread.to_string(),
checkpoint_id: id.to_string(),
run_id: None,
parent_checkpoint_id: parent.map(str::to_string),
namespace: vec![],
state: step as i32,
next_nodes: vec![NodeId::from("n")],
completed_tasks: vec![],
pending_writes: vec![],
interrupts: vec![],
metadata: serde_json::json!({ "source": "loop", "step": step }),
}
}
pub async fn checkpointer_contract<C>(cp: C)
where
C: Checkpointer<i32>,
{
cp.put(contract_checkpoint("t1", "c1", None, 1))
.await
.expect("put c1");
cp.put(contract_checkpoint("t1", "c2", Some("c1"), 2))
.await
.expect("put c2");
let latest = cp.get("t1", None).await.expect("get latest").expect("some");
assert_eq!(latest.checkpoint_id, "c2", "latest checkpoint id");
assert_eq!(latest.state, 2, "latest state");
let first = cp
.get("t1", Some("c1"))
.await
.expect("get specific")
.expect("some");
assert_eq!(first.checkpoint_id, "c1", "specific checkpoint id");
assert!(cp.get("nope", None).await.expect("get miss").is_none());
assert!(
cp.get("t1", Some("missing"))
.await
.expect("get miss")
.is_none()
);
let list = cp.list("t1").await.expect("list");
assert_eq!(list.len(), 2, "listed count");
assert_eq!(list[0].checkpoint_id, "c1", "list order[0]");
assert_eq!(
list[1].parent_checkpoint_id.as_deref(),
Some("c1"),
"list lineage"
);
let threads = cp.list_threads().await.expect("list_threads");
assert!(
threads.iter().any(|t| t == "t1"),
"list_threads contains t1"
);
let second_thread = "t2";
for i in 1..=3 {
let parent = (i > 1).then(|| format!("c{}", i - 1));
cp.put(contract_checkpoint(
second_thread,
&format!("c{i}"),
parent.as_deref(),
i,
))
.await
.expect("put for prune");
}
cp.prune(second_thread, 1).await.expect("prune");
let pruned = cp.list(second_thread).await.expect("list pruned");
assert!(!pruned.is_empty(), "prune keeps at least the window");
cp.delete_thread("t1").await.expect("delete_thread");
assert!(
cp.get("t1", None)
.await
.expect("get after delete")
.is_none()
);
}
pub fn taskstore_contract<S>(store: S)
where
S: TaskStore,
{
let spec = |id: &str| {
crate::graph::orchestration::OrchestrationTaskSpec::new(
id,
OrchestrationTaskKind::SubAgent {
agent: "worker".into(),
},
)
};
let happy = TaskId::new("happy");
let record = store.insert(spec("happy")).expect("insert");
assert_eq!(record.status, OrchestrationTaskStatus::Pending);
assert_eq!(
store.mark_running(&happy).expect("running").status,
OrchestrationTaskStatus::Running
);
let done = store
.complete(&happy, OrchestrationTaskResult::text("ok"))
.expect("complete");
assert_eq!(done.status, OrchestrationTaskStatus::Completed);
assert!(done.ended_at.is_some(), "completed sets ended_at");
assert!(
store.request_cancel(&happy).is_err(),
"terminal task rejects cancel"
);
let bad = TaskId::new("bad");
store.insert(spec("bad")).expect("insert bad");
assert_eq!(
store.fail(&bad, "boom".into()).expect("fail").status,
OrchestrationTaskStatus::Failed
);
let slow = TaskId::new("slow");
store.insert(spec("slow")).expect("insert slow");
assert_eq!(
store
.timeout(&slow, "deadline".into())
.expect("timeout")
.status,
OrchestrationTaskStatus::TimedOut
);
let cancelable = TaskId::new("cancelable");
store.insert(spec("cancelable")).expect("insert cancelable");
store.request_cancel(&cancelable).expect("request_cancel");
assert_eq!(
store.get(&cancelable).expect("get").status,
OrchestrationTaskStatus::CancelRequested
);
assert_eq!(
store.mark_cancelled(&cancelable).expect("mark").status,
OrchestrationTaskStatus::Cancelled
);
let doomed = TaskId::new("doomed");
store.insert(spec("doomed")).expect("insert doomed");
store.kill(&doomed).expect("kill");
assert_eq!(
store.get(&doomed).expect("get").status,
OrchestrationTaskStatus::Abandoned
);
let timed = TaskId::new("timed");
store.insert(spec("timed")).expect("insert timed");
let updated = store.set_timeout_ms(&timed, 1234).expect("set_timeout_ms");
assert_eq!(updated.spec.timeout_ms, Some(1234));
let completed = store.list(OrchestrationTaskFilter {
status: Some(OrchestrationTaskStatus::Completed),
..OrchestrationTaskFilter::default()
});
assert_eq!(completed.len(), 1, "one completed task");
assert_eq!(completed[0].spec.task_id.as_str(), "happy");
}
pub async fn checkpointer_concurrent_contract<C>(cp: std::sync::Arc<C>)
where
C: Checkpointer<i32> + 'static,
{
const WRITERS: usize = 8;
let mut handles = Vec::new();
for w in 0..WRITERS {
let cp = cp.clone();
handles.push(tokio::spawn(async move {
cp.put(contract_checkpoint("shared", &format!("c{w}"), None, w))
.await
.expect("concurrent put");
}));
}
for h in handles {
h.await.expect("writer task joins");
}
for w in 0..WRITERS {
let got = cp.get("shared", Some(&format!("c{w}"))).await.expect("get");
assert!(got.is_some(), "checkpoint c{w} survived concurrent writes");
}
}
pub fn taskstore_concurrent_contract<S>(store: std::sync::Arc<S>)
where
S: TaskStore + 'static,
{
const WRITERS: usize = 8;
std::thread::scope(|scope| {
for w in 0..WRITERS {
let store = store.clone();
scope.spawn(move || {
let id = TaskId::new(format!("task-{w}"));
store
.insert(crate::graph::orchestration::OrchestrationTaskSpec::new(
id.as_str(),
OrchestrationTaskKind::SubAgent {
agent: "worker".into(),
},
))
.expect("concurrent insert");
store.mark_running(&id).expect("concurrent mark_running");
});
}
});
let all = store.list(OrchestrationTaskFilter::default());
assert_eq!(all.len(), WRITERS, "every concurrent insert must land once");
assert!(
all.iter()
.all(|r| r.status == OrchestrationTaskStatus::Running),
"every concurrently-advanced task reached Running"
);
}
pub fn taskstore_replay_contract<S, F>(reopen: F)
where
S: TaskStore,
F: Fn() -> S,
{
let id = TaskId::new("survivor");
let spec = || {
crate::graph::orchestration::OrchestrationTaskSpec::new(
"survivor",
OrchestrationTaskKind::SubAgent {
agent: "worker".into(),
},
)
};
{
let store = reopen();
store.insert(spec()).expect("insert");
store.mark_running(&id).expect("running");
store
.complete(&id, OrchestrationTaskResult::text("done"))
.expect("complete");
}
let reopened = reopen();
let record = reopened.get(&id).expect("task survives reopen");
assert_eq!(record.status, OrchestrationTaskStatus::Completed);
assert_eq!(
reopened.history(&id).len(),
3,
"pending → running → completed history replays"
);
}