use crate::{
Worker, WorkerRunError,
activities::{ActivityContext, ActivityError, ActivityInfo},
};
use anyhow::bail;
use futures_util::future::{BoxFuture, LocalBoxFuture};
use std::{
any::Any,
collections::HashMap,
sync::{Arc, OnceLock},
};
use temporalio_common::{
data_converters::{
GenericPayloadConverter, PayloadConversionError, SerializationContext, TemporalSerializable,
},
protos::{
coresdk::{
workflow_activation::{WorkflowActivation, remove_from_cache::EvictionReason},
workflow_completion::WorkflowActivationCompletion,
},
temporal::api::common::v1::Payload,
},
};
mod activity_execution_value {
use super::*;
pub trait Sealed {
fn to_activity_payload(
&self,
context: &SerializationContext<'_>,
) -> Result<Payload, PayloadConversionError>;
}
impl<T> Sealed for T
where
T: Any + TemporalSerializable + Send + Sync,
{
fn to_activity_payload(
&self,
context: &SerializationContext<'_>,
) -> Result<Payload, PayloadConversionError> {
context.converter.to_payload(context, self)
}
}
}
#[async_trait::async_trait(?Send)]
pub trait WorkerInterceptor: Send + Sync {
fn run_worker<'a>(
&'a self,
input: RunWorkerInput<'a>,
next: Next<'a, RunWorkerInput<'a>, LocalBoxFuture<'a, Result<(), WorkerRunError>>>,
) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> {
next.run(input)
}
fn with_workflow_replay_worker<'a>(
&'a self,
input: WithWorkflowReplayWorkerInput<'a>,
next: Next<
'a,
WithWorkflowReplayWorkerInput<'a>,
LocalBoxFuture<'a, Result<(), WorkerRunError>>,
>,
) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> {
next.run(input)
}
async fn on_workflow_activation_completion(&self, _completion: &WorkflowActivationCompletion) {}
fn on_shutdown(&self, _sdk_worker: &Worker) {}
async fn on_workflow_activation(
&self,
_activation: &WorkflowActivation,
) -> Result<(), anyhow::Error> {
Ok(())
}
}
pub struct Next<'a, I, O> {
inner: Box<dyn FnOnce(I) -> O + Send + 'a>,
}
#[derive(Debug)]
#[non_exhaustive]
pub struct RunWorkerInput<'a> {
pub worker: &'a mut Worker,
}
impl<'a> RunWorkerInput<'a> {
pub(crate) fn new(worker: &'a mut Worker) -> Self {
Self { worker }
}
}
#[derive(Debug)]
#[non_exhaustive]
pub struct WithWorkflowReplayWorkerInput<'a> {
pub worker: &'a mut Worker,
}
impl<'a> WithWorkflowReplayWorkerInput<'a> {
pub(crate) fn new(worker: &'a mut Worker) -> Self {
Self { worker }
}
}
pub(crate) fn call_run_worker<'a>(
interceptors: &'a [Arc<dyn WorkerInterceptor>],
input: RunWorkerInput<'a>,
terminal: Next<'a, RunWorkerInput<'a>, LocalBoxFuture<'a, Result<(), WorkerRunError>>>,
) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> {
if let Some((interceptor, remaining)) = interceptors.split_first() {
let next = Next::new(move |input| call_run_worker(remaining, input, terminal));
interceptor.run_worker(input, next)
} else {
terminal.run(input)
}
}
pub(crate) fn call_with_workflow_replay_worker<'a>(
interceptors: &'a [Arc<dyn WorkerInterceptor>],
input: WithWorkflowReplayWorkerInput<'a>,
terminal: Next<
'a,
WithWorkflowReplayWorkerInput<'a>,
LocalBoxFuture<'a, Result<(), WorkerRunError>>,
>,
) -> LocalBoxFuture<'a, Result<(), WorkerRunError>> {
if let Some((interceptor, remaining)) = interceptors.split_first() {
let next =
Next::new(move |input| call_with_workflow_replay_worker(remaining, input, terminal));
interceptor.with_workflow_replay_worker(input, next)
} else {
terminal.run(input)
}
}
impl<'a, I, O> Next<'a, I, O> {
pub(crate) fn new(f: impl FnOnce(I) -> O + Send + 'a) -> Self {
Self { inner: Box::new(f) }
}
pub fn run(self, input: I) -> O {
(self.inner)(input)
}
}
#[non_exhaustive]
pub struct ExecuteActivityInput {
context: ActivityContext,
args: Box<dyn Any + Send + Sync>,
}
impl ExecuteActivityInput {
pub(crate) fn new(context: ActivityContext, args: Box<dyn Any + Send + Sync>) -> Self {
Self { context, args }
}
pub(crate) fn into_parts(self) -> (ActivityContext, Box<dyn Any + Send + Sync>) {
(self.context, self.args)
}
pub fn activity_info(&self) -> &ActivityInfo {
self.context.info()
}
pub fn headers(&self) -> &HashMap<String, Payload> {
self.context.headers()
}
pub fn headers_mut(&mut self) -> &mut HashMap<String, Payload> {
self.context.headers_mut()
}
pub fn args_ref<T: Any>(&self) -> Option<&T> {
self.args.downcast_ref()
}
pub fn args_mut<T: Any>(&mut self) -> Option<&mut T> {
self.args.downcast_mut()
}
}
pub trait ActivityExecutionValue:
Any + TemporalSerializable + Send + Sync + activity_execution_value::Sealed
{
fn as_any(&self) -> &dyn Any;
}
impl<T> ActivityExecutionValue for T
where
T: Any + TemporalSerializable + Send + Sync,
{
fn as_any(&self) -> &dyn Any {
self
}
}
impl dyn ActivityExecutionValue {
pub fn downcast_ref<T: Any>(&self) -> Option<&T> {
self.as_any().downcast_ref()
}
pub(crate) fn serialize_payload(
&self,
context: &SerializationContext<'_>,
) -> Result<Payload, PayloadConversionError> {
self.to_activity_payload(context)
}
}
pub type ExecuteActivityResult = Result<Box<dyn ActivityExecutionValue>, ActivityError>;
pub type ExecuteActivityOutput<'a> = BoxFuture<'a, ExecuteActivityResult>;
pub trait ActivityInboundInterceptor: Send + Sync + 'static {
fn execute_activity<'a>(
&'a self,
input: ExecuteActivityInput,
next: Next<'a, ExecuteActivityInput, ExecuteActivityOutput<'a>>,
) -> ExecuteActivityOutput<'a> {
next.run(input)
}
}
pub struct FailOnNondeterminismInterceptor {}
#[async_trait::async_trait(?Send)]
impl WorkerInterceptor for FailOnNondeterminismInterceptor {
async fn on_workflow_activation(
&self,
activation: &WorkflowActivation,
) -> Result<(), anyhow::Error> {
if matches!(
activation.eviction_reason(),
Some(EvictionReason::Nondeterminism)
) {
bail!("Workflow is being evicted because of nondeterminism! {activation}");
}
Ok(())
}
}
#[derive(Default)]
pub struct ReturnWorkflowExitValueInterceptor {
result_value: Arc<OnceLock<Payload>>,
}
impl ReturnWorkflowExitValueInterceptor {
pub fn result_handle(&self) -> Arc<OnceLock<Payload>> {
self.result_value.clone()
}
}
#[async_trait::async_trait(?Send)]
impl WorkerInterceptor for ReturnWorkflowExitValueInterceptor {
async fn on_workflow_activation_completion(&self, c: &WorkflowActivationCompletion) {
if let Some(v) = c.complete_workflow_execution_value() {
let _ = self.result_value.set(v.clone());
}
}
}