use super::super::error::CortexAdapterError;
use super::adapter::WorkflowAdapter;
use super::state::WorkflowState;
use super::types::{TaskId, TaskStatus};
pub fn derive_shard_ids(parent: TaskId, count: usize) -> Vec<TaskId> {
(0..count)
.map(|k| {
let mut z = parent.wrapping_add(
(k as u64)
.wrapping_add(1)
.wrapping_mul(0x9E37_79B9_7F4A_7C15),
);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
})
.collect()
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ShardGroup {
shards: Vec<TaskId>,
reduce: TaskId,
parent: Option<TaskId>,
}
impl ShardGroup {
pub fn new(shards: Vec<TaskId>, reduce: TaskId) -> Self {
Self {
shards,
reduce,
parent: None,
}
}
pub fn derived(parent: TaskId, shard_count: usize, reduce: TaskId) -> Self {
Self {
shards: derive_shard_ids(parent, shard_count),
reduce,
parent: Some(parent),
}
}
pub fn shards(&self) -> &[TaskId] {
&self.shards
}
pub fn reduce(&self) -> TaskId {
self.reduce
}
pub fn parent(&self) -> Option<TaskId> {
self.parent
}
pub fn shard_count(&self) -> usize {
self.shards.len()
}
pub fn join_ready(&self, state: &WorkflowState) -> bool {
self.shards.iter().all(|s| {
state
.get(*s)
.map(|t| t.status == TaskStatus::Done)
.unwrap_or(false)
})
}
pub fn pending(&self, state: &WorkflowState) -> Vec<TaskId> {
self.shards
.iter()
.copied()
.filter(|s| {
state
.get(*s)
.map(|t| t.status != TaskStatus::Done)
.unwrap_or(true)
})
.collect()
}
pub fn failed(&self, state: &WorkflowState) -> Vec<TaskId> {
self.shards
.iter()
.copied()
.filter(|s| {
state
.get(*s)
.map(|t| t.status == TaskStatus::Failed)
.unwrap_or(false)
})
.collect()
}
pub fn done_count(&self, state: &WorkflowState) -> usize {
self.shards
.iter()
.filter(|s| {
state
.get(**s)
.map(|t| t.status == TaskStatus::Done)
.unwrap_or(false)
})
.count()
}
pub fn join_status(&self, state: &WorkflowState, policy: JoinPolicy) -> JoinStatus {
let failed = self.failed(state);
let done = self.done_count(state);
let total = self.shards.len();
match policy {
JoinPolicy::AllOrNothing => {
if !failed.is_empty() {
JoinStatus::Failed(failed)
} else if done == total {
JoinStatus::Ready
} else {
JoinStatus::Pending
}
}
JoinPolicy::BestEffort => {
let terminal = self
.shards
.iter()
.filter(|s| {
state
.get(**s)
.map(|t| t.status.is_terminal())
.unwrap_or(false)
})
.count();
if terminal == total {
JoinStatus::Ready
} else {
JoinStatus::Pending
}
}
JoinPolicy::Threshold(n) => {
debug_assert!(
n <= total,
"Threshold({n}) exceeds shard count {total}: unsatisfiable by construction"
);
if done >= n {
JoinStatus::Ready
} else if total.saturating_sub(failed.len()) < n {
JoinStatus::Failed(failed)
} else {
JoinStatus::Pending
}
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum JoinPolicy {
#[default]
AllOrNothing,
BestEffort,
Threshold(usize),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum JoinStatus {
Ready,
Failed(Vec<TaskId>),
Pending,
}
pub fn fan_out(wf: &WorkflowAdapter, group: &ShardGroup) -> Result<u64, CortexAdapterError> {
let mut last = 0;
for &shard in group.shards() {
last = wf.submit(shard)?;
if let Some(parent) = group.parent() {
last = wf.link(parent, shard)?;
}
}
Ok(last)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Join {
Submitted(u64),
AlreadySubmitted,
Pending,
Failed(Vec<TaskId>),
}
pub fn try_join(wf: &WorkflowAdapter, group: &ShardGroup) -> Result<Join, CortexAdapterError> {
try_join_with(wf, group, JoinPolicy::AllOrNothing)
}
pub fn try_join_with(
wf: &WorkflowAdapter,
group: &ShardGroup,
policy: JoinPolicy,
) -> Result<Join, CortexAdapterError> {
let (already, status) = {
let state = wf.state();
let guard = state.read();
(
guard.contains(group.reduce()),
group.join_status(&guard, policy),
)
};
if already {
return Ok(Join::AlreadySubmitted);
}
match status {
JoinStatus::Ready => {
let mut seq = wf.submit(group.reduce())?;
if let Some(parent) = group.parent() {
seq = wf.link(parent, group.reduce())?;
}
Ok(Join::Submitted(seq))
}
JoinStatus::Failed(f) => Ok(Join::Failed(f)),
JoinStatus::Pending => Ok(Join::Pending),
}
}
pub fn propagate_failure(
wf: &WorkflowAdapter,
group: &ShardGroup,
) -> Result<u64, CortexAdapterError> {
let mut last = 0;
let pending = group.pending(&wf.state().read());
for shard in pending {
last = wf.request_cancel(shard)?;
}
if let Some(parent) = group.parent() {
last = wf.fail(parent)?;
}
Ok(last)
}
pub fn block_on_failure(
wf: &WorkflowAdapter,
group: &ShardGroup,
) -> Result<u64, CortexAdapterError> {
match group.parent() {
Some(parent) => wf.block(parent),
None => Ok(0),
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::super::types::TaskState;
use super::*;
use crate::adapter::net::redex::Redex;
fn done_state(ids: &[TaskId]) -> WorkflowState {
let mut s = WorkflowState::new();
for id in ids {
s.tasks.insert(
*id,
TaskState {
step: 0,
status: TaskStatus::Done,
attempts: 0,
},
);
}
s
}
fn state_with(pairs: &[(TaskId, TaskStatus)]) -> WorkflowState {
let mut s = WorkflowState::new();
for (id, status) in pairs {
s.tasks.insert(
*id,
TaskState {
step: 0,
status: *status,
attempts: 0,
},
);
}
s
}
#[test]
fn derived_shard_ids_are_deterministic_and_distinct() {
let a = derive_shard_ids(0xABCD, 4);
let b = derive_shard_ids(0xABCD, 4);
assert_eq!(a, b, "same parent → same shard ids (replay-stable)");
assert_eq!(a.len(), 4);
let mut sorted = a.clone();
sorted.sort_unstable();
sorted.dedup();
assert_eq!(sorted.len(), 4, "shard ids are distinct");
assert!(!a.contains(&0xABCD));
}
#[test]
fn join_ready_only_when_all_shards_done() {
let group = ShardGroup::new(vec![1, 2, 3], 9);
assert!(!group.join_ready(&WorkflowState::new()));
assert_eq!(group.pending(&WorkflowState::new()), vec![1, 2, 3]);
let partial = done_state(&[1, 2]);
assert!(!group.join_ready(&partial));
assert_eq!(group.pending(&partial), vec![3]);
let all = done_state(&[1, 2, 3]);
assert!(group.join_ready(&all));
assert!(group.pending(&all).is_empty());
}
#[test]
fn empty_group_joins_immediately() {
let group = ShardGroup::new(vec![], 9);
assert!(group.join_ready(&WorkflowState::new()));
}
#[test]
#[should_panic(expected = "exceeds shard count")]
fn threshold_above_shard_count_is_caught() {
let group = ShardGroup::new(vec![1, 2, 3], 9);
let running = state_with(&[
(1, TaskStatus::Running),
(2, TaskStatus::Running),
(3, TaskStatus::Running),
]);
let _ = group.join_status(&running, JoinPolicy::Threshold(5));
}
#[tokio::test]
async fn map_reduce_join_fires_reduce_only_after_all_shards_done() {
let redex = Redex::new();
let wf = WorkflowAdapter::open(&redex, 0x0F10_00C1).await.unwrap();
let group = ShardGroup::new(vec![10, 11, 12], 99);
let seq = fan_out(&wf, &group).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(group.pending(&wf.state().read()).len(), 3);
wf.start(10).unwrap();
wf.complete(10).unwrap();
wf.start(11).unwrap();
wf.retry(11).unwrap(); let seq = wf.complete(11).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(try_join(&wf, &group).unwrap(), Join::Pending);
assert!(
!wf.state().read().contains(99),
"reduce not submitted until every shard is Done",
);
wf.start(12).unwrap();
let seq = wf.complete(12).unwrap();
wf.wait_for_seq(seq).await.unwrap();
let reduce_seq = match try_join(&wf, &group).unwrap() {
Join::Submitted(seq) => seq,
other => panic!("expected Submitted, got {other:?}"),
};
wf.wait_for_seq(reduce_seq).await.unwrap();
assert!(
wf.state().read().contains(99),
"reduce submitted on all-done"
);
assert_eq!(try_join(&wf, &group).unwrap(), Join::AlreadySubmitted);
assert_eq!(wf.get(11).unwrap().attempts, 1);
}
#[test]
fn failed_shard_surfaces_as_join_failed_not_pending() {
let group = ShardGroup::new(vec![1, 2, 3], 9);
let st = state_with(&[
(1, TaskStatus::Done),
(2, TaskStatus::Failed),
(3, TaskStatus::Running),
]);
assert!(!group.join_ready(&st));
assert_eq!(group.failed(&st), vec![2]);
assert_eq!(
group.join_status(&st, JoinPolicy::AllOrNothing),
JoinStatus::Failed(vec![2]),
);
}
#[test]
fn best_effort_joins_when_every_shard_is_terminal() {
let group = ShardGroup::new(vec![1, 2, 3], 9);
let mixed = state_with(&[
(1, TaskStatus::Done),
(2, TaskStatus::Failed),
(3, TaskStatus::Running),
]);
assert_eq!(
group.join_status(&mixed, JoinPolicy::BestEffort),
JoinStatus::Pending
);
let terminal = state_with(&[
(1, TaskStatus::Done),
(2, TaskStatus::Failed),
(3, TaskStatus::Done),
]);
assert_eq!(
group.join_status(&terminal, JoinPolicy::BestEffort),
JoinStatus::Ready
);
}
#[test]
fn threshold_joins_at_n_done_and_fails_once_unreachable() {
let group = ShardGroup::new(vec![1, 2, 3], 9);
let one = state_with(&[
(1, TaskStatus::Done),
(2, TaskStatus::Running),
(3, TaskStatus::Running),
]);
assert_eq!(
group.join_status(&one, JoinPolicy::Threshold(2)),
JoinStatus::Pending
);
let two = state_with(&[
(1, TaskStatus::Done),
(2, TaskStatus::Done),
(3, TaskStatus::Running),
]);
assert_eq!(
group.join_status(&two, JoinPolicy::Threshold(2)),
JoinStatus::Ready
);
let lost = state_with(&[
(1, TaskStatus::Done),
(2, TaskStatus::Failed),
(3, TaskStatus::Failed),
]);
assert_eq!(
group.join_status(&lost, JoinPolicy::Threshold(2)),
JoinStatus::Failed(vec![2, 3]),
);
}
#[tokio::test]
async fn propagate_failure_cancels_pending_and_fails_parent() {
let redex = Redex::new();
let wf = WorkflowAdapter::open(&redex, 0x0F10_00C2).await.unwrap();
let group = ShardGroup::derived(7, 3, 99);
let shards = group.shards().to_vec();
let seq = wf.submit(7).unwrap(); wf.wait_for_seq(seq).await.unwrap();
let seq = fan_out(&wf, &group).unwrap();
wf.wait_for_seq(seq).await.unwrap();
wf.complete(shards[0]).unwrap();
let seq = wf.fail(shards[1]).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(
try_join(&wf, &group).unwrap(),
Join::Failed(vec![shards[1]]),
);
assert!(
!wf.state().read().contains(99),
"reduce never submitted on failure"
);
let seq = propagate_failure(&wf, &group).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert!(
wf.is_cancel_requested(shards[2]),
"running sibling cancelled"
);
assert!(
!wf.is_cancel_requested(shards[0]),
"the Done shard isn't cancelled"
);
assert_eq!(
wf.get(7).unwrap().status,
TaskStatus::Failed,
"parent failed"
);
}
#[tokio::test]
async fn block_on_failure_marks_parent_blocked_and_spares_siblings() {
let redex = Redex::new();
let wf = WorkflowAdapter::open(&redex, 0x0F10_00C3).await.unwrap();
let group = ShardGroup::derived(7, 3, 99);
let shards = group.shards().to_vec();
let seq = wf.submit(7).unwrap();
wf.wait_for_seq(seq).await.unwrap();
fan_out(&wf, &group).unwrap();
wf.start(shards[2]).unwrap();
let seq = wf.fail(shards[1]).unwrap();
wf.wait_for_seq(seq).await.unwrap();
let seq = block_on_failure(&wf, &group).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(
wf.get(7).unwrap().status,
TaskStatus::Blocked,
"parent parked Blocked (external state), distinct from Waiting",
);
assert!(!wf.is_cancel_requested(shards[2]));
}
#[tokio::test]
async fn delete_parent_reclaims_all_shards() {
let redex = Redex::new();
let wf = WorkflowAdapter::open(&redex, 0x0F10_00C4).await.unwrap();
let group = ShardGroup::derived(7, 3, 99);
let shards = group.shards().to_vec();
let seq = wf.submit(7).unwrap();
wf.wait_for_seq(seq).await.unwrap();
let seq = fan_out(&wf, &group).unwrap();
wf.wait_for_seq(seq).await.unwrap();
for s in &shards {
assert!(wf.get(*s).is_some());
}
let seq = wf.delete(7).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert!(wf.get(7).is_none());
for s in &shards {
assert!(wf.get(*s).is_none(), "shard {s} reclaimed with the parent");
}
}
#[test]
fn shards_have_independent_leases() {
use super::super::TaskLease;
use super::super::TaskLeaseOutcome;
use crate::adapter::net::behavior::fold::{Fold, ReservationFold};
use crate::adapter::net::current_timestamp_micros;
use crate::adapter::net::identity::EntityKeypair;
let fold = Fold::<ReservationFold>::with_sweep_interval(Duration::ZERO);
let a = EntityKeypair::generate();
let b = EntityKeypair::generate();
let (na, nb) = (a.entity_id().node_id(), b.entity_id().node_id());
let mut la = TaskLease::new(&fold, &a, na);
let mut lb = TaskLease::new(&fold, &b, nb);
let until = current_timestamp_micros() + 60_000_000;
assert_eq!(la.acquire(10, until).unwrap(), TaskLeaseOutcome::Acquired);
assert_eq!(lb.acquire(11, until).unwrap(), TaskLeaseOutcome::Acquired);
assert_eq!(la.current_holder(10), Some(na));
assert_eq!(lb.current_holder(11), Some(nb));
assert_eq!(la.acquire(11, until).unwrap(), TaskLeaseOutcome::Contended);
}
}