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,
9 },
10 protos::{
11 coresdk::workflow_activation::InitializeWorkflow, temporal::api::common::v1::Payload,
12 },
13};
14use temporalio_workflow::{
15 BaseWorkflowContext, 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 pub(crate) fn extend(&mut self, other: &Self) -> Result<(), WorkflowRegistrationError> {
79 for workflow in other.workflows.values() {
80 self.insert_workflow(workflow.definition.clone(), workflow.factory.clone())?;
81 }
82 Ok(())
83 }
84
85 pub fn new() -> Self {
87 Self::default()
88 }
89
90 pub fn register_workflow<W: WorkflowImplementation>(
94 &mut self,
95 ) -> Result<&mut Self, WorkflowRegistrationError>
96 where
97 <W::Run as WorkflowDefinition>::Input: Send,
98 {
99 let factory = Arc::new(move |input| {
100 let (payloads, payload_converter, base_ctx) = workflow_input_parts(input);
101 instantiate_workflow::<W>(payloads, payload_converter, base_ctx)
102 .context("Failed to instantiate native workflow")
103 });
104 self.insert_workflow(W::definition(), factory)?;
105 Ok(self)
106 }
107
108 pub fn register_workflow_run_with_factory<W, F>(
113 &mut self,
114 user_factory: F,
115 ) -> Result<&mut Self, WorkflowRegistrationError>
116 where
117 W: WorkflowImplementation,
118 <W::Run as WorkflowDefinition>::Input: Send,
119 F: Fn() -> W + Send + Sync + 'static,
120 {
121 if W::HAS_INIT {
122 return Err(WorkflowRegistrationError::FactoryRegistrationWithInit {
123 workflow_type: W::definition().workflow_type,
124 });
125 }
126
127 let factory = Arc::new(move |input| {
128 let (payloads, payload_converter, base_ctx) = workflow_input_parts(input);
129 let ser_ctx = SerializationContext {
130 data: &SerializationContextData::Workflow,
131 converter: &payload_converter,
132 };
133 let input: <W::Run as WorkflowDefinition>::Input =
134 payload_converter.from_payloads(&ser_ctx, payloads)?;
135
136 let workflow = user_factory();
137 Ok(Box::new(GuestWorkflowInstance::<W>::new_with_workflow(
138 workflow,
139 base_ctx,
140 Some(input),
141 )) as Box<dyn WorkflowInstance>)
142 });
143
144 self.insert_workflow(W::definition(), factory)?;
145 Ok(self)
146 }
147
148 pub fn is_empty(&self) -> bool {
150 self.workflows.is_empty()
151 }
152
153 pub(crate) fn insert_workflow(
154 &mut self,
155 definition: WorkflowDefinitionDescriptor,
156 factory: WorkflowExecutionFactory,
157 ) -> Result<(), WorkflowRegistrationError> {
158 let workflow_type = definition.workflow_type.clone();
159 if self.workflows.contains_key(&workflow_type) {
160 return Err(WorkflowRegistrationError::DuplicateWorkflowType { workflow_type });
161 }
162 self.workflows.insert(
163 workflow_type,
164 RegisteredWorkflow {
165 definition,
166 factory,
167 },
168 );
169 Ok(())
170 }
171
172 pub(crate) fn get_workflow(&self, workflow_type: &str) -> Option<WorkflowExecutionFactory> {
173 self.workflows
174 .get(workflow_type)
175 .map(|wf| wf.factory.clone())
176 }
177
178 pub fn workflow_definitions(&self) -> impl Iterator<Item = &WorkflowDefinitionDescriptor> + '_ {
180 self.workflows.values().map(|wf| &wf.definition)
181 }
182}
183
184fn workflow_input_parts(
185 input: WorkflowExecutionInput,
186) -> (Vec<Payload>, PayloadConverter, BaseWorkflowContext) {
187 let WorkflowExecutionInput {
188 namespace,
189 task_queue,
190 run_id,
191 init_workflow_job,
192 data_converter,
193 host,
194 patch_activation_callback,
195 workflow_interceptor_constructors,
196 } = input;
197 let payloads = init_workflow_job.arguments.clone();
198 let payload_converter = data_converter.payload_converter().clone();
199 let init = WorkflowInit {
200 namespace,
201 task_queue,
202 run_id,
203 initialize_workflow: init_workflow_job,
204 };
205 let base_ctx = BaseWorkflowContext::from_raw(
206 init,
207 data_converter,
208 host,
209 patch_activation_callback,
210 workflow_interceptor_constructors,
211 );
212 (payloads, payload_converter, base_ctx)
213}
214
215impl Debug for WorkflowDefinitions {
216 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
217 f.debug_struct("WorkflowDefinitions")
218 .field("workflows", &self.workflows.keys().collect::<Vec<_>>())
219 .finish()
220 }
221}