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
26pub(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
38pub(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#[derive(Debug, thiserror::Error, Clone, PartialEq, Eq)]
53pub enum WorkflowRegistrationError {
54 #[error("Workflow type {workflow_type} is already registered")]
56 DuplicateWorkflowType {
57 workflow_type: String,
59 },
60
61 #[error(
63 "Workflow type {workflow_type} must not define an #[init] method when registered with a factory"
64 )]
65 FactoryRegistrationWithInit {
66 workflow_type: String,
68 },
69}
70
71#[derive(Default, Clone)]
73pub struct WorkflowDefinitions {
74 workflows: HashMap<String, RegisteredWorkflow>,
75}
76
77impl WorkflowDefinitions {
78 #[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 pub fn new() -> Self {
89 Self::default()
90 }
91
92 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 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 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 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}