use std::collections::HashMap;
use std::sync::Arc;
use futures_util::stream::{self, FuturesUnordered, Stream, StreamExt, TryStreamExt};
use serde::{Deserialize, Serialize};
use taquba::object_store::{ObjectStore, path::Path};
use taquba::{Queue, SettlementEffects};
use tracing::warn;
use crate::blob::ObjectPrefix;
use crate::durable::{self, DurableMember, DurableTermination};
use crate::error::{Error, Result};
use crate::keys::{
HEADER_GROUP, HEADER_GROUP_KEY, RunId, group_member_kv_key, group_members_kv_prefix,
outcome_kv_key,
};
use crate::memo::MemoStore;
use crate::runtime::{RunOptions, RunSpec, RunTermination, RuntimeCore};
use crate::sweep::Clearable;
use crate::terminal::{RunOutcome, TerminalStatus};
const SUBMIT_CONCURRENCY: usize = 32;
const MEMBER_PAGE_SIZE: usize = 1000;
pub(crate) fn member_run_id(group_id: &RunId, key: &str) -> RunId {
RunId::digest(&[group_id.as_bytes(), b"/", key.as_bytes()])
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Membership {
pub(crate) group_id: RunId,
pub(crate) key: String,
}
impl Membership {
pub(crate) fn from_headers(headers: &HashMap<String, String>) -> Result<Option<Self>> {
let (Some(group_id), Some(key)) =
(headers.get(HEADER_GROUP), headers.get(HEADER_GROUP_KEY))
else {
return Ok(None);
};
Ok(Some(Self {
group_id: RunId::new(group_id.as_str())?,
key: key.clone(),
}))
}
pub(crate) fn reserved_headers(&self) -> Vec<(&'static str, String)> {
vec![
(HEADER_GROUP, self.group_id.to_string()),
(HEADER_GROUP_KEY, self.key.clone()),
]
}
pub(crate) fn kv_key(&self) -> Vec<u8> {
group_member_kv_key(&self.group_id, &self.key)
}
pub(crate) fn run_id(&self) -> RunId {
member_run_id(&self.group_id, &self.key)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct GroupMember {
pub key: String,
pub input: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub(crate) struct Manifest {
pub(crate) group_id: RunId,
pub(crate) members: Vec<GroupMember>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GroupStatus {
pub group_id: RunId,
pub total: usize,
pub pending: usize,
pub succeeded: usize,
pub failed: usize,
pub cancelled: usize,
}
#[derive(Debug, Clone)]
pub struct MemberResult {
pub key: String,
pub run_id: RunId,
pub termination: RunTermination,
pub outcome: Option<RunOutcome>,
}
pub(crate) struct MemberState {
pub(crate) key: String,
pub(crate) record: DurableMember,
}
impl MemberState {
pub(crate) fn status(&self) -> Option<TerminalStatus> {
self.record
.terminated
.as_ref()
.map(|termination| termination.status.into())
}
}
#[derive(Clone)]
pub(crate) struct GroupStore {
objects: ObjectPrefix,
memo_store: MemoStore,
queue: Arc<Queue>,
}
impl GroupStore {
pub(crate) fn new(
store: Arc<dyn ObjectStore>,
prefix: impl Into<String>,
memo_store: MemoStore,
queue: Arc<Queue>,
) -> Self {
Self {
objects: ObjectPrefix::new(store, prefix),
memo_store,
queue,
}
}
fn manifest_path(&self, group_id: &RunId) -> Path {
self.objects.path(&format!("groups/{group_id}/manifest"))
}
pub(crate) async fn read_manifest(&self, group_id: &RunId) -> Result<Option<Manifest>> {
match self.objects.get(&self.manifest_path(group_id)).await? {
Some(bytes) => durable::decode(&bytes).map(Some),
None => Ok(None),
}
}
async fn write_manifest(&self, manifest: &Manifest) -> Result<()> {
self.objects
.put(
&self.manifest_path(&manifest.group_id),
&durable::encode(manifest),
)
.await
}
pub(crate) async fn member(
&self,
group_id: &RunId,
key: &str,
) -> Result<Option<DurableMember>> {
durable::kv_record(self.queue.view(), &group_member_kv_key(group_id, key)).await
}
pub(crate) async fn members(&self, group_id: &RunId) -> Result<Vec<MemberState>> {
let prefix = group_members_kv_prefix(group_id);
let mut members = Vec::new();
let mut entries =
std::pin::pin!(self.queue.view().kv_entries(&prefix, .., MEMBER_PAGE_SIZE));
while let Some((kv_key, value)) = entries.try_next().await? {
let key = String::from_utf8_lossy(&kv_key[prefix.len()..]).into_owned();
if let Some(record) = durable::decode_or_absent(
&value,
"group member record",
&format_args!("{group_id}/{key}"),
) {
members.push(MemberState { key, record });
}
}
Ok(members)
}
pub(crate) async fn forget(&self, group_id: &RunId) -> Result<()> {
let mut keys = Vec::new();
if let Some(manifest) = self.read_manifest(group_id).await? {
for member in &manifest.members {
let run_id = member_run_id(group_id, &member.key);
self.memo_store.clear_memos_for_run(&run_id).await?;
keys.push(outcome_kv_key(&run_id));
if keys.len() == MEMBER_PAGE_SIZE {
self.delete_keys(std::mem::take(&mut keys)).await?;
}
}
}
let prefix = group_members_kv_prefix(group_id);
let mut entries =
std::pin::pin!(self.queue.view().kv_entries(&prefix, .., MEMBER_PAGE_SIZE));
while let Some((key, _)) = entries.try_next().await? {
keys.push(key);
if keys.len() == MEMBER_PAGE_SIZE {
self.delete_keys(std::mem::take(&mut keys)).await?;
}
}
if !keys.is_empty() {
self.delete_keys(keys).await?;
}
self.objects.delete(&self.manifest_path(group_id)).await?;
Ok(())
}
async fn delete_keys(&self, keys: Vec<Vec<u8>>) -> Result<()> {
self.queue
.commit_effects(SettlementEffects::default().kv_deletes(keys))
.await?;
Ok(())
}
}
impl Clearable for GroupStore {
type Error = Error;
async fn clear(&self, group_id: &RunId) -> Result<Vec<Vec<u8>>> {
self.forget(group_id).await.map(|()| Vec::new())
}
}
#[derive(Clone)]
pub struct RunGroup {
runtime: Arc<RuntimeCore>,
id: RunId,
}
impl RunGroup {
pub(crate) fn new(runtime: Arc<RuntimeCore>, id: RunId) -> Self {
Self { runtime, id }
}
pub fn id(&self) -> &RunId {
&self.id
}
fn core(&self) -> &RuntimeCore {
&self.runtime
}
fn store(&self) -> &GroupStore {
&self.core().group_store
}
pub(crate) async fn manifest(&self) -> Result<Manifest> {
self.store()
.read_manifest(&self.id)
.await?
.ok_or_else(|| Error::GroupNotFound(self.id.clone()))
}
pub(crate) async fn members(&self) -> Result<Vec<MemberState>> {
self.store().members(&self.id).await
}
pub async fn submit(&self, members: Vec<GroupMember>, options: &RunOptions) -> Result<()> {
let mut seen = std::collections::HashSet::new();
for member in &members {
if !seen.insert(member.key.as_str()) {
return Err(Error::DuplicateMemberKey(member.key.clone()));
}
}
let manifest = Manifest {
group_id: self.id.clone(),
members,
};
match self.store().read_manifest(&self.id).await? {
Some(existing) if existing.members != manifest.members => {
return Err(Error::GroupMismatch(self.id.clone()));
}
Some(_) => {}
None => self.store().write_manifest(&manifest).await?,
}
self.submit_members(manifest.members, options).await
}
pub async fn resume(&self, options: &RunOptions) -> Result<()> {
let manifest = self.manifest().await?;
self.submit_members(manifest.members, options).await
}
async fn submit_members(&self, members: Vec<GroupMember>, options: &RunOptions) -> Result<()> {
let succeeded: std::collections::HashSet<String> = self
.members()
.await?
.into_iter()
.filter(|member| member.status() == Some(TerminalStatus::Succeeded))
.map(|member| member.key)
.collect();
let mut submissions = stream::iter(
members
.into_iter()
.filter(|member| !succeeded.contains(&member.key)),
)
.map(|member| async move {
let membership = Membership {
group_id: self.id.clone(),
key: member.key,
};
self.runtime
.submit_member(
&membership,
RunSpec {
run_id: Some(membership.run_id()),
input: member.input,
options: options.clone(),
effects: SettlementEffects::default(),
},
)
.await
.map(|_| ())
})
.buffer_unordered(SUBMIT_CONCURRENCY);
while submissions.try_next().await?.is_some() {}
Ok(())
}
async fn terminations(&self) -> Result<impl Stream<Item = Result<MemberState>> + use<>> {
let manifest = self.manifest().await?;
let waits: FuturesUnordered<_> = manifest
.members
.into_iter()
.map(|member| {
let group = self.clone();
async move {
let record = group.wait_member(&member.key).await?;
Ok(MemberState {
key: member.key,
record,
})
}
})
.collect();
let group = self.clone();
let marker = stream::once(async move {
if let Err(err) = group.mark_terminated().await {
warn!(group_id = %group.id, "group terminal marker write failed: {err}");
}
})
.filter_map(|()| async { None::<Result<MemberState>> });
Ok(waits.chain(marker))
}
pub async fn results(&self) -> Result<impl Stream<Item = Result<MemberResult>> + use<>> {
let terminations = self.terminations().await?;
let group = self.clone();
Ok(terminations.then(move |member| {
let group = group.clone();
async move {
let member = member?;
let termination: RunTermination = member
.record
.terminated
.ok_or_else(|| Error::InconsistentRunState(member.record.run_id.clone()))?
.into();
let outcome = group
.core()
.view
.run_result_of(&member.record.run_id, &termination)
.await?;
Ok(MemberResult {
key: member.key,
run_id: member.record.run_id,
termination,
outcome: outcome.map(|result| result.outcome),
})
}
}))
}
pub async fn status(&self) -> Result<GroupStatus> {
let manifest = self.manifest().await?;
let mut status = GroupStatus {
group_id: self.id.clone(),
total: manifest.members.len(),
pending: 0,
succeeded: 0,
failed: 0,
cancelled: 0,
};
for member in self.members().await? {
match member.status() {
None => status.pending += 1,
Some(TerminalStatus::Succeeded) => status.succeeded += 1,
Some(TerminalStatus::Failed) => status.failed += 1,
Some(TerminalStatus::Cancelled) => status.cancelled += 1,
}
}
Ok(status)
}
pub async fn cancel(&self) -> Result<usize> {
let mut cancellations = stream::iter(
self.members()
.await?
.into_iter()
.filter(|member| member.status().is_none()),
)
.map(|member| async move { self.runtime.cancel(&member.record.run_id).await })
.buffer_unordered(SUBMIT_CONCURRENCY);
let mut cancelled = 0;
while let Some(recorded) = cancellations.try_next().await? {
cancelled += usize::from(recorded);
}
Ok(cancelled)
}
async fn wait_member(&self, key: &str) -> Result<DurableMember> {
let member = |member: Option<DurableMember>| {
member.ok_or_else(|| Error::MemberNotSubmitted {
group_id: self.id.clone(),
key: key.to_string(),
})
};
let record = member(self.store().member(&self.id, key).await?)?;
if record.terminated.is_some() {
return Ok(record);
}
let run_id = member_run_id(&self.id, key);
self.core().wait_run(&run_id).await?;
let record = member(self.store().member(&self.id, key).await?)?;
if record.terminated.is_none() {
return Err(Error::InconsistentRunState(run_id));
}
Ok(record)
}
async fn mark_terminated(&self) -> Result<()> {
let core = self.core();
if let Some(sweep) = &core.group_sweep {
let effects = sweep.mark(SettlementEffects::default(), &self.id, core.clock.now_ms());
core.queue.commit_effects(effects).await?;
}
Ok(())
}
pub async fn forget(&self) -> Result<()> {
self.store().forget(&self.id).await
}
}
pub(crate) fn pending_member(run_id: &RunId) -> DurableMember {
DurableMember {
run_id: run_id.clone(),
terminated: None,
}
}
pub(crate) fn terminated_member(run_id: &RunId, termination: DurableTermination) -> DurableMember {
DurableMember {
run_id: run_id.clone(),
terminated: Some(termination),
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
use crate::runner::{Step, StepError, StepOutcome, StepRunner};
use crate::runtime::{RunSpec, WorkflowRuntime};
use crate::terminal::NoopTerminalHook;
use crate::test_util::{open_queue, open_queue_at, rid};
struct TwoSteps;
impl StepRunner for TwoSteps {
async fn run_step(&self, step: &Step) -> std::result::Result<StepOutcome, StepError> {
if step.step_number == 0 {
Ok(StepOutcome::continue_now(step.payload.clone()))
} else {
Ok(StepOutcome::Succeed {
result: step.step_number.to_string().into_bytes(),
})
}
}
}
fn member(key: &str) -> GroupMember {
GroupMember {
key: key.to_string(),
input: key.as_bytes().to_vec(),
}
}
struct Rejecting;
impl StepRunner for Rejecting {
async fn run_step(&self, _: &Step) -> std::result::Result<StepOutcome, StepError> {
Err(StepError::permanent("rejected"))
}
}
#[tokio::test(start_paused = true)]
async fn a_result_record_of_an_earlier_termination_is_not_reported_for_a_re_run_member() {
let (queue, store, clock) = open_queue_at(10_000).await;
let runtime =
WorkflowRuntime::builder(queue.clone(), store, Rejecting, NoopTerminalHook).build();
let group = runtime.group(rid("g"));
group
.submit(vec![member("a")], &RunOptions::default())
.await
.unwrap();
let worker = runtime.spawn(std::future::pending::<()>());
let first: Vec<MemberResult> = tokio::time::timeout(Duration::from_secs(10), async {
group.results().await.unwrap().try_collect().await
})
.await
.expect("results finished in time")
.unwrap();
assert_eq!(first[0].termination.error.as_deref(), Some("rejected"));
assert_eq!(
first[0].termination.error_kind,
Some(crate::StepErrorKind::Permanent)
);
assert!(
first[0].outcome.is_some(),
"the worker recorded the failure"
);
worker.shutdown().await.unwrap();
clock.advance(Duration::from_secs(1));
group
.submit(vec![member("a")], &RunOptions::default())
.await
.unwrap();
let claim = queue
.claim("workflow-steps", Duration::from_secs(60))
.await
.unwrap()
.unwrap();
queue.dead_letter(&claim, "hung").await.unwrap();
assert_eq!(runtime.inner.core.reconcile_dead_steps().await.unwrap(), 1);
let second: Vec<MemberResult> = group.results().await.unwrap().try_collect().await.unwrap();
assert_eq!(second[0].termination.error.as_deref(), Some("hung"));
assert_eq!(second[0].termination.error_kind, None);
assert_eq!(second[0].termination.terminated_at_ms, 11_000);
assert!(
second[0].outcome.is_none(),
"the first termination's record does not belong to the second",
);
}
#[tokio::test(start_paused = true)]
async fn membership_holds_across_steps_and_terminations_wait_for_the_last_one() {
let (queue, store) = open_queue().await;
let runtime = WorkflowRuntime::builder(queue.clone(), store, TwoSteps, NoopTerminalHook)
.poll_interval(Duration::from_millis(10))
.build();
let group = runtime.group(rid("g"));
group
.submit(vec![member("a"), member("b")], &RunOptions::default())
.await
.unwrap();
let pending = group.members().await.unwrap();
assert_eq!(pending.len(), 2);
assert!(pending.iter().all(|m| m.status().is_none()));
assert_eq!(pending[0].record.run_id, member_run_id(&rid("g"), "a"));
assert!(matches!(
group.submit(vec![member("a"), member("a")], &RunOptions::default()).await,
Err(Error::DuplicateMemberKey(key)) if key == "a"
));
let worker = runtime.spawn(std::future::pending::<()>());
let terminated: Vec<MemberResult> = tokio::time::timeout(Duration::from_secs(10), async {
group.results().await.unwrap().try_collect().await
})
.await
.expect("results finished in time")
.unwrap();
assert_eq!(terminated.len(), 2);
for m in &terminated {
assert_eq!(m.termination.status, TerminalStatus::Succeeded);
let outcome = m.outcome.as_ref().expect("the worker recorded the outcome");
assert_eq!(outcome.run_id, m.run_id);
assert_eq!(
outcome.final_step, 1,
"the member terminated at its second step"
);
assert_eq!(outcome.result.as_deref(), Some(b"1".as_slice()));
}
let status = group.status().await.unwrap();
assert_eq!((status.total, status.succeeded, status.pending), (2, 2, 0));
let again: Vec<MemberResult> = group.results().await.unwrap().try_collect().await.unwrap();
assert_eq!(again.len(), 2);
worker.shutdown().await.unwrap();
}
#[tokio::test(start_paused = true)]
async fn the_run_options_apply_to_every_member() {
let (queue, store) = open_queue().await;
let runtime =
WorkflowRuntime::builder(queue.clone(), store, TwoSteps, NoopTerminalHook).build();
let group = runtime.group(rid("g"));
let options = RunOptions {
priority: Some(3),
max_attempts_per_step: Some(5),
headers: HashMap::from([("tenant".to_string(), "acme".to_string())]),
..Default::default()
};
group
.submit(vec![member("a"), member("b")], &options)
.await
.unwrap();
for _ in 0..2 {
let job = queue
.claim("workflow-steps", Duration::from_secs(30))
.await
.unwrap()
.expect("a member's step job");
assert_eq!(job.priority, 3);
assert_eq!(job.max_attempts, 5);
assert_eq!(job.headers.get("tenant").map(String::as_str), Some("acme"));
}
}
#[tokio::test(start_paused = true)]
async fn a_group_cancellation_records_the_member_cancelled() {
let (queue, store, _clock) = open_queue_at(10_000).await;
let runtime =
WorkflowRuntime::builder(queue.clone(), store, TwoSteps, NoopTerminalHook).build();
let group = runtime.group(rid("g"));
group
.submit(vec![member("a")], &RunOptions::default())
.await
.unwrap();
let run_id = member_run_id(&rid("g"), "a");
assert_eq!(group.cancel().await.unwrap(), 1);
assert_eq!(group.cancel().await.unwrap(), 0, "no member is active");
let results: Vec<MemberResult> =
group.results().await.unwrap().try_collect().await.unwrap();
assert_eq!(results[0].run_id, run_id);
assert_eq!(
results[0].termination,
RunTermination {
status: TerminalStatus::Cancelled,
error: None,
error_kind: None,
final_step: 0,
terminated_at_ms: 10_000,
}
);
assert!(
results[0].outcome.is_none(),
"a pending step is cancelled without a worker, so no result is recorded",
);
group
.submit(vec![member("a")], &RunOptions::default())
.await
.unwrap();
assert!(group.members().await.unwrap()[0].status().is_none());
let plain = runtime
.submit(RunSpec {
run_id: Some(run_id.clone()),
input: b"a".to_vec(),
..Default::default()
})
.await
.unwrap();
assert!(!plain.newly_submitted);
group.forget().await.unwrap();
assert!(group.members().await.unwrap().is_empty());
assert!(matches!(
group.manifest().await,
Err(Error::GroupNotFound(_))
));
}
#[tokio::test(start_paused = true)]
async fn the_group_sweep_removes_the_members_records_with_the_group() {
let (queue, store, clock) = open_queue_at(10_000).await;
let runtime =
WorkflowRuntime::builder(queue.clone(), store.clone(), TwoSteps, NoopTerminalHook)
.group_retention(Duration::from_secs(1))
.build();
let group = runtime.group(rid("g"));
group
.submit(vec![member("a")], &RunOptions::default())
.await
.unwrap();
let run_id = member_run_id(&rid("g"), "a");
assert_eq!(group.cancel().await.unwrap(), 1);
let memos = crate::memo::MemoStore::new(store, "workflow-steps-memo");
memos.new_run_memo(&run_id).put("k", b"v").await.unwrap();
let results: Vec<MemberResult> =
group.results().await.unwrap().try_collect().await.unwrap();
assert_eq!(results.len(), 1);
let sweep = runtime.inner.core.group_sweep.as_ref().unwrap();
let marker = sweep.marker_key(&rid("g"), 10_000);
assert!(
queue.view().kv_get(&marker).await.unwrap().is_some(),
"the marker is written when the last termination is observed"
);
let terminal_record = outcome_kv_key(&run_id);
assert!(
queue
.view()
.kv_get(&terminal_record)
.await
.unwrap()
.is_some()
);
clock.advance(Duration::from_millis(999));
assert_eq!(
runtime.inner.core.sweep_once().await.unwrap(),
0,
"the marker is not yet expired"
);
clock.advance(Duration::from_millis(1));
assert_eq!(runtime.inner.core.sweep_once().await.unwrap(), 1);
assert!(queue.view().kv_get(&marker).await.unwrap().is_none());
assert!(group.members().await.unwrap().is_empty());
assert!(matches!(
group.manifest().await,
Err(Error::GroupNotFound(_))
));
assert!(
queue
.view()
.kv_get(&terminal_record)
.await
.unwrap()
.is_none(),
"the member's terminal record is removed with the group"
);
assert!(
memos
.new_run_memo(&run_id)
.get("k")
.await
.unwrap()
.is_none()
);
assert!(
runtime.status(&run_id).await.unwrap().is_none(),
"nothing of the member remains"
);
}
}