use std::sync::Arc;
use async_trait::async_trait;
use serde::{Serialize, de::DeserializeOwned};
use serde_json::Value;
use tokio::sync::watch;
use tokio::task::{AbortHandle, JoinHandle};
use crate::error::{Error, JobFailure, Result};
#[async_trait]
pub(crate) trait JobHandle: Send + Sync {
fn id(&self) -> Option<&str> {
None
}
async fn status(&self) -> Result<String>;
async fn wait(&self) -> Result<TerminalResult>;
async fn cancel(&self) -> Result<()>;
}
#[derive(Clone)]
pub(crate) struct TerminalResult {
value: Option<Value>,
request_id: Option<String>,
}
impl TerminalResult {
fn local(value: Value) -> Self {
Self {
value: Some(value),
request_id: None,
}
}
pub(crate) fn remote(value: Option<Value>, request_id: String) -> Self {
Self {
value,
request_id: Some(request_id),
}
}
pub(crate) fn value(&self) -> Option<&Value> {
self.value.as_ref()
}
fn decode<T: DeserializeOwned>(self) -> Result<T> {
let value = self.value.ok_or_else(|| match &self.request_id {
Some(request_id) => Error::Http {
source: "successful typed job response did not contain a result".into(),
request_id: request_id.clone(),
status_code: None,
},
None => Error::Runtime {
message: "successful typed job did not contain a result".to_string(),
},
})?;
serde_json::from_value(value).map_err(|error| match self.request_id {
Some(request_id) => Error::Http {
source: format!("failed to parse typed job result: {error}").into(),
request_id,
status_code: None,
},
None => Error::Runtime {
message: format!("failed to parse typed job result: {error}"),
},
})
}
}
type ResultDecoder<T> = Arc<dyn Fn(TerminalResult) -> Result<T> + Send + Sync>;
enum JobInner<T> {
Handle {
handle: Box<dyn JobHandle>,
decode: ResultDecoder<T>,
},
Completed(T),
}
pub struct Job<T = ()>
where
T: Clone + Send + Sync + 'static,
{
inner: JobInner<T>,
}
impl<T> std::fmt::Debug for Job<T>
where
T: Clone + Send + Sync + 'static,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Job")
.field("id", &self.id())
.field("done", &matches!(self.inner, JobInner::Completed(_)))
.finish()
}
}
impl Job<()> {
pub(crate) fn new_done() -> Self {
Self {
inner: JobInner::Completed(()),
}
}
pub(crate) fn new(handle: Box<dyn JobHandle>) -> Self {
Self {
inner: JobInner::Handle {
handle,
decode: Arc::new(|_| Ok(())),
},
}
}
}
impl<T> Job<T>
where
T: Clone + DeserializeOwned + Send + Sync + 'static,
{
pub(crate) fn new_typed(handle: Box<dyn JobHandle>) -> Self {
Self {
inner: JobInner::Handle {
handle,
decode: Arc::new(TerminalResult::decode::<T>),
},
}
}
}
impl<T> Job<T>
where
T: Clone + Serialize + DeserializeOwned + Send + Sync + 'static,
{
pub(crate) fn spawned(task: JoinHandle<Result<T>>) -> Self {
Self::new_typed(Box::new(SpawnedJob::new(task)))
}
}
impl<T> Job<T>
where
T: Clone + Send + Sync + 'static,
{
pub fn id(&self) -> Option<&str> {
match &self.inner {
JobInner::Handle { handle, .. } => handle.id(),
JobInner::Completed(_) => None,
}
}
pub async fn status(&self) -> Result<String> {
match &self.inner {
JobInner::Handle { handle, .. } => handle.status().await,
JobInner::Completed(_) => Ok("finished".to_string()),
}
}
pub async fn wait(&self) -> Result<T> {
match &self.inner {
JobInner::Handle { handle, decode } => (decode)(handle.wait().await?),
JobInner::Completed(result) => Ok(result.clone()),
}
}
pub async fn cancel(&self) -> Result<()> {
match &self.inner {
JobInner::Handle { handle, .. } => handle.cancel().await,
JobInner::Completed(_) => Ok(()),
}
}
pub fn map<U, F>(self, map: F) -> Job<U>
where
U: Clone + Send + Sync + 'static,
F: Fn(T) -> U + Send + Sync + 'static,
{
match self.inner {
JobInner::Handle { handle, decode } => Job {
inner: JobInner::Handle {
handle,
decode: Arc::new(move |result| Ok(map((decode)(result)?))),
},
},
JobInner::Completed(result) => Job {
inner: JobInner::Completed(map(result)),
},
}
}
}
#[derive(Clone)]
enum Outcome {
Succeeded(TerminalResult),
Failed(Arc<Error>),
Cancelled,
}
impl Outcome {
fn into_result(self) -> Result<TerminalResult> {
match self {
Self::Succeeded(result) => Ok(result),
Self::Failed(source) => Err(Error::JobFailed {
job_id: None,
failure: JobFailure::from_source(source),
}),
Self::Cancelled => Err(Error::JobCancelled { job_id: None }),
}
}
}
struct SpawnedJob {
outcome: watch::Receiver<Option<Outcome>>,
abort: AbortHandle,
}
impl SpawnedJob {
fn new<T>(task: JoinHandle<Result<T>>) -> Self
where
T: Serialize + Send + 'static,
{
let abort = task.abort_handle();
let (tx, outcome) = watch::channel(None);
tokio::spawn(async move {
let outcome = match task.await {
Ok(Ok(result)) => match serde_json::to_value(result) {
Ok(value) => Outcome::Succeeded(TerminalResult::local(value)),
Err(err) => Outcome::Failed(Arc::new(Error::Runtime {
message: format!("failed to serialize job result: {err}"),
})),
},
Ok(Err(err)) => Outcome::Failed(Arc::new(err)),
Err(err) if err.is_cancelled() => Outcome::Cancelled,
Err(err) => Outcome::Failed(Arc::new(Error::Runtime {
message: format!("job task failed: {err}"),
})),
};
let _ = tx.send(Some(outcome));
});
Self { outcome, abort }
}
}
#[async_trait]
impl JobHandle for SpawnedJob {
async fn status(&self) -> Result<String> {
let label = match &*self.outcome.borrow() {
None => "running",
Some(Outcome::Succeeded(_)) => "finished",
Some(Outcome::Failed(_)) => "failed",
Some(Outcome::Cancelled) => "cancelled",
};
Ok(label.to_string())
}
async fn wait(&self) -> Result<TerminalResult> {
let mut outcome = self.outcome.clone();
let settled = outcome
.wait_for(|outcome| outcome.is_some())
.await
.map_err(|_| Error::Runtime {
message: "job outcome was dropped before it completed".to_string(),
})?
.clone()
.expect("wait_for returns once an outcome is set");
settled.into_result()
}
async fn cancel(&self) -> Result<()> {
self.abort.abort();
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::future::pending;
use super::*;
#[tokio::test]
async fn mapped_spawned_job_reuses_outcome() {
let job = Job::spawned(tokio::spawn(async { Ok(41_u64) })).map(|value| value + 1);
assert_eq!(job.wait().await.unwrap(), 42);
assert_eq!(job.wait().await.unwrap(), 42);
assert_eq!(job.status().await.unwrap(), "finished");
}
#[tokio::test]
async fn mapped_spawned_job_preserves_cancellation() {
let job = Job::spawned(tokio::spawn(async { pending::<Result<u64>>().await }))
.map(|value| value.to_string());
job.cancel().await.unwrap();
assert!(matches!(job.wait().await, Err(Error::JobCancelled { .. })));
assert_eq!(job.status().await.unwrap(), "cancelled");
}
}