use std::future::Future;
use std::time::Duration;
use dapr_durabletask::api::ExternalEventResult;
use dapr_durabletask::worker::{ActivityResult, OrchestratorResult, Registry};
use serde::Serialize;
use serde::de::DeserializeOwned;
use super::options::{ActivityOptions, SubWorkflowOptions};
pub type WorkflowContext = dapr_durabletask::task::OrchestrationContext;
pub type ActivityContext = dapr_durabletask::task::ActivityContext;
pub trait WorkflowContextExt {
fn get_input_typed<T: DeserializeOwned>(&self) -> dapr_durabletask::api::Result<T>;
fn call_activity_typed<T, I>(
&self,
name: &str,
input: I,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static,
I: Serialize + Send + 'static;
fn call_activity_with_options_typed<T>(
&self,
name: &str,
options: ActivityOptions,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static;
fn call_sub_workflow_typed<T, I>(
&self,
name: &str,
input: I,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static,
I: Serialize + Send + 'static;
fn call_sub_workflow_with_options_typed<T>(
&self,
name: &str,
options: SubWorkflowOptions,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static;
fn propagated_history(
&self,
) -> Option<std::sync::Arc<dapr_durabletask::api::PropagatedHistory>>;
fn wait_for_external_event_typed<T>(
&self,
name: &str,
timeout: Option<Duration>,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static;
}
impl WorkflowContextExt for WorkflowContext {
fn get_input_typed<T: DeserializeOwned>(&self) -> dapr_durabletask::api::Result<T> {
self.input()
}
fn call_activity_typed<T, I>(
&self,
name: &str,
input: I,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static,
I: Serialize + Send + 'static,
{
let task = self.call_activity(name, input);
async move { deserialize_task_output(task.await?) }
}
fn call_activity_with_options_typed<T>(
&self,
name: &str,
options: ActivityOptions,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static,
{
let (input, task_options) = options.into_parts();
let task = self.call_activity_with_options(name, input, task_options);
async move { deserialize_task_output(task.await?) }
}
fn call_sub_workflow_typed<T, I>(
&self,
name: &str,
input: I,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static,
I: Serialize + Send + 'static,
{
let task = self.call_sub_orchestrator_with_options(
name,
input,
dapr_durabletask::task::SubOrchestratorOptions::new(),
);
async move { deserialize_task_output(task.await?) }
}
fn call_sub_workflow_with_options_typed<T>(
&self,
name: &str,
options: SubWorkflowOptions,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static,
{
let (input, task_options) = options.into_parts();
let task = self.call_sub_orchestrator_with_options(name, input, task_options);
async move { deserialize_task_output(task.await?) }
}
fn propagated_history(
&self,
) -> Option<std::sync::Arc<dapr_durabletask::api::PropagatedHistory>> {
dapr_durabletask::task::OrchestrationContext::propagated_history(self)
}
fn wait_for_external_event_typed<T>(
&self,
name: &str,
timeout: Option<Duration>,
) -> impl Future<Output = dapr_durabletask::api::Result<T>> + Send + 'static
where
T: DeserializeOwned + Send + 'static,
{
let ctx = self.clone();
let name = name.to_string();
async move {
let output = match timeout {
Some(duration) => match ctx
.wait_for_external_event_with_timeout(&name, duration)
.await?
{
ExternalEventResult::Received(output) => output,
ExternalEventResult::TimedOut => {
return Err(dapr_durabletask::api::DurableTaskError::Timeout);
}
},
None => ctx.wait_for_external_event(&name).await?,
};
deserialize_task_output(output)
}
}
}
pub trait ActivityContextExt {
fn get_input<T: DeserializeOwned>(
&self,
input: Option<&str>,
) -> dapr_durabletask::api::Result<T>;
fn propagated_history(&self) -> Option<&dapr_durabletask::api::PropagatedHistory>;
}
impl ActivityContextExt for ActivityContext {
fn get_input<T: DeserializeOwned>(
&self,
input: Option<&str>,
) -> dapr_durabletask::api::Result<T> {
deserialize_task_output(input.map(ToOwned::to_owned))
}
fn propagated_history(&self) -> Option<&dapr_durabletask::api::PropagatedHistory> {
dapr_durabletask::task::ActivityContext::propagated_history(self)
}
}
pub trait RegistryExt {
fn add_workflow<F, Fut>(&mut self, name: &str, f: F)
where
F: Fn(WorkflowContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = OrchestratorResult> + Send + 'static;
fn add_activity<F, Fut>(&mut self, name: &str, f: F)
where
F: Fn(ActivityContext, Option<String>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = ActivityResult> + Send + 'static;
}
impl RegistryExt for Registry {
fn add_workflow<F, Fut>(&mut self, name: &str, f: F)
where
F: Fn(WorkflowContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = OrchestratorResult> + Send + 'static,
{
self.add_named_orchestrator(name, f);
}
fn add_activity<F, Fut>(&mut self, name: &str, f: F)
where
F: Fn(ActivityContext, Option<String>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = ActivityResult> + Send + 'static,
{
self.add_named_activity(name, f);
}
}
fn deserialize_task_output<T: DeserializeOwned>(
output: Option<String>,
) -> dapr_durabletask::api::Result<T> {
match output {
Some(value) => serde_json::from_str(&value).map_err(Into::into),
None => serde_json::from_value(serde_json::Value::Null).map_err(Into::into),
}
}