Skip to main content

temporalio_sdk/
workflow_registry.rs

1use std::{collections::HashMap, fmt::Debug, rc::Rc, sync::Arc};
2
3use anyhow::Context;
4use temporalio_common::{
5    WorkflowDefinition,
6    data_converters::{
7        DataConverter, GenericPayloadConverter, PayloadConverter, SerializationContext,
8        SerializationContextData, WorkflowSerializationContext,
9    },
10    protos::{
11        coresdk::workflow_activation::InitializeWorkflow, temporal::api::common::v1::Payload,
12    },
13};
14use temporalio_workflow::{
15    BaseWorkflowContext, InternalPatchActivationCallback as PatchActivationCallback,
16    runtime::{
17        entry::WorkflowImplementation,
18        guest::WorkflowInstance,
19        host::WorkflowHost,
20        instance::{GuestWorkflowInstance, instantiate_workflow},
21        types::{WorkflowDefinitionDescriptor, WorkflowInit},
22    },
23    workflow_interceptors::WorkflowInterceptorConstructor,
24};
25
26/// Host-owned execution inputs used to instantiate a single workflow run.
27pub(crate) struct WorkflowExecutionInput {
28    pub namespace: String,
29    pub task_queue: String,
30    pub run_id: String,
31    pub init_workflow_job: InitializeWorkflow,
32    pub data_converter: DataConverter,
33    pub host: Rc<dyn WorkflowHost>,
34    pub patch_activation_callback: Option<PatchActivationCallback>,
35    pub workflow_interceptor_constructors: Vec<WorkflowInterceptorConstructor>,
36}
37
38/// Creates workflow execution instances from activation input payloads and context.
39pub(crate) type WorkflowExecutionFactory = Arc<
40    dyn Fn(WorkflowExecutionInput) -> Result<Box<dyn WorkflowInstance>, anyhow::Error>
41        + Send
42        + Sync,
43>;
44
45#[derive(Clone)]
46struct RegisteredWorkflow {
47    definition: WorkflowDefinitionDescriptor,
48    factory: WorkflowExecutionFactory,
49}
50
51/// Error returned when a workflow cannot be registered.
52#[derive(Debug, thiserror::Error, Clone, PartialEq, Eq)]
53pub enum WorkflowRegistrationError {
54    /// The workflow type is already registered.
55    #[error("Workflow type {workflow_type} is already registered")]
56    DuplicateWorkflowType {
57        /// The duplicate workflow type.
58        workflow_type: String,
59    },
60
61    /// The workflow type has an `#[init]` method and cannot be registered with a factory.
62    #[error(
63        "Workflow type {workflow_type} must not define an #[init] method when registered with a factory"
64    )]
65    FactoryRegistrationWithInit {
66        /// The workflow type with an `#[init]` method.
67        workflow_type: String,
68    },
69}
70
71/// Contains workflow registrations in a form ready for execution by workers.
72#[derive(Default, Clone)]
73pub struct WorkflowDefinitions {
74    workflows: HashMap<String, RegisteredWorkflow>,
75}
76
77impl WorkflowDefinitions {
78    // Only used by Plugins so feature flagged to avoid dead code.
79    #[cfg(feature = "experimental")]
80    pub(crate) fn extend(&mut self, other: &Self) -> Result<(), WorkflowRegistrationError> {
81        for workflow in other.workflows.values() {
82            self.insert_workflow(workflow.definition.clone(), workflow.factory.clone())?;
83        }
84        Ok(())
85    }
86
87    /// Creates a new empty `WorkflowDefinitions`.
88    pub fn new() -> Self {
89        Self::default()
90    }
91
92    /// Register a workflow implementation.
93    ///
94    /// Returns an error if a workflow with the same type is already registered.
95    pub fn register_workflow<W: WorkflowImplementation>(
96        &mut self,
97    ) -> Result<&mut Self, WorkflowRegistrationError>
98    where
99        <W::Run as WorkflowDefinition>::Input: Send,
100    {
101        let factory = Arc::new(move |input| {
102            let (payloads, payload_converter, base_ctx) = workflow_input_parts(input);
103            instantiate_workflow::<W>(payloads, payload_converter, base_ctx)
104                .context("Failed to instantiate native workflow")
105        });
106        self.insert_workflow(W::definition(), factory)?;
107        Ok(self)
108    }
109
110    /// Register a workflow with a custom factory for instance creation.
111    ///
112    /// Returns an error if a workflow with the same type is already registered, or if the workflow
113    /// type defines an `#[init]` method.
114    pub fn register_workflow_run_with_factory<W, F>(
115        &mut self,
116        user_factory: F,
117    ) -> Result<&mut Self, WorkflowRegistrationError>
118    where
119        W: WorkflowImplementation,
120        <W::Run as WorkflowDefinition>::Input: Send,
121        F: Fn() -> W + Send + Sync + 'static,
122    {
123        if W::HAS_INIT {
124            return Err(WorkflowRegistrationError::FactoryRegistrationWithInit {
125                workflow_type: W::definition().workflow_type,
126            });
127        }
128
129        let factory = Arc::new(move |input| {
130            let (payloads, payload_converter, base_ctx) = workflow_input_parts(input);
131            let context_data =
132                SerializationContextData::Workflow(WorkflowSerializationContext::new());
133            let ser_ctx = SerializationContext::new(&context_data, &payload_converter);
134            let input: <W::Run as WorkflowDefinition>::Input =
135                payload_converter.from_payloads(&ser_ctx, payloads)?;
136
137            let workflow = user_factory();
138            Ok(Box::new(GuestWorkflowInstance::<W>::new_with_workflow(
139                workflow,
140                base_ctx,
141                Some(input),
142            )) as Box<dyn WorkflowInstance>)
143        });
144
145        self.insert_workflow(W::definition(), factory)?;
146        Ok(self)
147    }
148
149    /// Check if any workflows are registered.
150    pub fn is_empty(&self) -> bool {
151        self.workflows.is_empty()
152    }
153
154    pub(crate) fn insert_workflow(
155        &mut self,
156        definition: WorkflowDefinitionDescriptor,
157        factory: WorkflowExecutionFactory,
158    ) -> Result<(), WorkflowRegistrationError> {
159        let workflow_type = definition.workflow_type.clone();
160        if self.workflows.contains_key(&workflow_type) {
161            return Err(WorkflowRegistrationError::DuplicateWorkflowType { workflow_type });
162        }
163        self.workflows.insert(
164            workflow_type,
165            RegisteredWorkflow {
166                definition,
167                factory,
168            },
169        );
170        Ok(())
171    }
172
173    pub(crate) fn get_workflow(&self, workflow_type: &str) -> Option<WorkflowExecutionFactory> {
174        self.workflows
175            .get(workflow_type)
176            .map(|wf| wf.factory.clone())
177    }
178
179    /// Returns an iterator over registered workflow definitions.
180    pub fn workflow_definitions(&self) -> impl Iterator<Item = &WorkflowDefinitionDescriptor> + '_ {
181        self.workflows.values().map(|wf| &wf.definition)
182    }
183}
184
185fn workflow_input_parts(
186    input: WorkflowExecutionInput,
187) -> (Vec<Payload>, PayloadConverter, BaseWorkflowContext) {
188    let WorkflowExecutionInput {
189        namespace,
190        task_queue,
191        run_id,
192        init_workflow_job,
193        data_converter,
194        host,
195        patch_activation_callback,
196        workflow_interceptor_constructors,
197    } = input;
198    let payloads = init_workflow_job.arguments.clone();
199    let payload_converter = data_converter.payload_converter().clone();
200    let init = WorkflowInit {
201        namespace,
202        task_queue,
203        run_id,
204        initialize_workflow: init_workflow_job,
205    };
206    let base_ctx = BaseWorkflowContext::from_raw(
207        init,
208        data_converter,
209        host,
210        patch_activation_callback,
211        workflow_interceptor_constructors,
212    );
213    (payloads, payload_converter, base_ctx)
214}
215
216impl Debug for WorkflowDefinitions {
217    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
218        f.debug_struct("WorkflowDefinitions")
219            .field("workflows", &self.workflows.keys().collect::<Vec<_>>())
220            .finish()
221    }
222}