use std::collections::HashMap;
use std::marker::PhantomData;
use futures_util::{Stream, StreamExt, TryStreamExt};
use crate::group::{GroupMember, GroupStatus, RunGroup};
use crate::jobs::handle::{JobError, decode_end};
use crate::jobs::job::Job;
use crate::jobs::runner::job_payload;
use crate::{Result, RunId, RunOptions};
pub struct JobGroup<J: Job> {
group: RunGroup,
_marker: PhantomData<fn() -> J>,
}
#[derive(Debug)]
pub struct GroupResult<J: Job> {
pub key: String,
pub result: std::result::Result<J::Output, JobError>,
}
impl<J: Job> JobGroup<J> {
pub(crate) fn new(group: RunGroup) -> Self {
Self {
group,
_marker: PhantomData,
}
}
pub fn id(&self) -> &RunId {
self.group.id()
}
pub async fn submit(&self, jobs: impl IntoIterator<Item = J>) -> Result<()> {
self.submit_with(jobs, RunOptions::default()).await
}
pub async fn submit_with(
&self,
jobs: impl IntoIterator<Item = J>,
options: RunOptions,
) -> Result<()> {
let mut members = Vec::new();
for (i, job) in jobs.into_iter().enumerate() {
members.push(GroupMember {
key: job.idempotency_key().unwrap_or_else(|| format!("item-{i}")),
input: job_payload(&job)?,
});
}
self.group.submit(members, &options).await
}
pub async fn resume(&self) -> Result<()> {
self.resume_with(RunOptions::default()).await
}
pub async fn resume_with(&self, options: RunOptions) -> Result<()> {
self.group.resume(&options).await
}
pub async fn results(&self) -> Result<impl Stream<Item = Result<GroupResult<J>>> + use<J>> {
let results = self.group.results().await?;
Ok(results.map(|member| {
let member = member?;
Ok(GroupResult {
key: member.key,
result: decode_end::<J>(member.termination, member.outcome)?,
})
}))
}
pub async fn join(&self) -> Result<Vec<GroupResult<J>>> {
let manifest = self.group.manifest().await?;
let mut results: HashMap<String, std::result::Result<J::Output, JobError>> = HashMap::new();
let mut stream = std::pin::pin!(self.results().await?);
while let Some(result) = stream.try_next().await? {
results.insert(result.key, result.result);
}
Ok(manifest
.members
.into_iter()
.filter_map(|member| {
results.remove(&member.key).map(|result| GroupResult {
key: member.key,
result,
})
})
.collect())
}
pub async fn status(&self) -> Result<GroupStatus> {
self.group.status().await
}
pub async fn cancel(&self) -> Result<usize> {
self.group.cancel().await
}
pub async fn forget(&self) -> Result<()> {
self.group.forget().await
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use futures_util::TryStreamExt;
use serde::{Deserialize, Serialize};
use crate::jobs::{Job, JobContext, JobRunner};
use crate::test_util::{open_queue, rid};
use crate::{Error, StepErrorKind};
#[derive(Debug, thiserror::Error)]
#[error("{0}")]
struct TestError(String);
#[derive(Serialize, Deserialize)]
struct Square {
n: u32,
}
impl Job for Square {
const NAME: &'static str = "test.square";
type Output = u32;
type Error = TestError;
async fn run(&self, ctx: JobContext<'_>) -> std::result::Result<u32, TestError> {
ctx.state::<Arc<AtomicU32>>().fetch_add(1, Ordering::SeqCst);
if self.n == 13 {
return Err(TestError("unlucky".to_string()));
}
Ok(self.n * self.n)
}
fn classify(&self, _error: &TestError) -> StepErrorKind {
StepErrorKind::Permanent
}
fn idempotency_key(&self) -> Option<String> {
Some(format!("square:{}", self.n))
}
}
#[tokio::test(start_paused = true)]
async fn a_group_joins_its_results_in_submission_order_and_reruns_failures() {
let (queue, store) = open_queue().await;
let runs = Arc::new(AtomicU32::new(0));
let runner = JobRunner::builder(queue, store)
.register::<Square>()
.state(runs.clone())
.build();
let worker = runner.spawn(std::future::pending::<()>());
let jobs = || vec![Square { n: 3 }, Square { n: 13 }, Square { n: 2 }];
let group = runner.group::<Square>(rid("squares"));
group.submit(jobs()).await.unwrap();
let results = tokio::time::timeout(std::time::Duration::from_secs(10), group.join())
.await
.expect("join finished in time")
.unwrap();
let keys: Vec<&str> = results.iter().map(|r| r.key.as_str()).collect();
assert_eq!(keys, ["square:3", "square:13", "square:2"]);
assert_eq!(results[0].result.as_ref().unwrap(), &9);
assert_eq!(results[2].result.as_ref().unwrap(), &4);
let failure = results[1].result.as_ref().unwrap_err();
assert_eq!(
(failure.kind, failure.message.as_str()),
(StepErrorKind::Permanent, "unlucky")
);
let status = group.status().await.unwrap();
assert_eq!(
(
status.total,
status.pending,
status.succeeded,
status.failed
),
(3, 0, 2, 1)
);
group.submit(jobs()).await.unwrap();
tokio::time::timeout(std::time::Duration::from_secs(10), group.join())
.await
.expect("join finished in time")
.unwrap();
assert_eq!(runs.load(Ordering::SeqCst), 4);
let err = group.submit(vec![Square { n: 3 }]).await.unwrap_err();
assert!(matches!(err, Error::GroupMismatch(id) if id == "squares"));
group.resume().await.unwrap();
let mut streamed = 0;
let mut results = std::pin::pin!(group.results().await.unwrap());
while let Some(result) =
tokio::time::timeout(std::time::Duration::from_secs(10), results.try_next())
.await
.expect("results finished in time")
.unwrap()
{
streamed += 1;
assert_eq!(result.result.is_ok(), result.key != "square:13");
}
assert_eq!(streamed, 3);
assert_eq!(runs.load(Ordering::SeqCst), 5);
group.forget().await.unwrap();
assert!(matches!(group.status().await, Err(Error::GroupNotFound(_))));
assert!(matches!(group.resume().await, Err(Error::GroupNotFound(_))));
worker.shutdown().await.unwrap();
}
}