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 __private::sdk::{GuestWorkflowInstance, WorkflowHost, WorkflowInit, WorkflowInstance},
16 BaseWorkflowContext, InternalPatchActivationCallback as PatchActivationCallback,
17 workflow_interceptors::WorkflowInterceptorConstructor,
18 workflows::{WorkflowDefinitionDescriptor, WorkflowImplementation},
19};
20
21pub(crate) struct WorkflowExecutionInput {
23 pub namespace: String,
24 pub task_queue: String,
25 pub run_id: String,
26 pub init_workflow_job: InitializeWorkflow,
27 pub data_converter: DataConverter,
28 pub host: Rc<dyn WorkflowHost>,
29 pub patch_activation_callback: Option<PatchActivationCallback>,
30 pub workflow_interceptor_constructors: Vec<WorkflowInterceptorConstructor>,
31}
32
33pub(crate) type WorkflowExecutionFactory = Arc<
35 dyn Fn(WorkflowExecutionInput) -> Result<Box<dyn WorkflowInstance>, anyhow::Error>
36 + Send
37 + Sync,
38>;
39
40#[derive(Clone)]
41struct RegisteredWorkflow {
42 definition: WorkflowDefinitionDescriptor,
43 factory: WorkflowExecutionFactory,
44}
45
46#[derive(Debug, thiserror::Error, Clone, PartialEq, Eq)]
48pub enum WorkflowRegistrationError {
49 #[error("Workflow type {workflow_type} is already registered")]
51 DuplicateWorkflowType {
52 workflow_type: String,
54 },
55
56 #[error(
58 "Workflow type {workflow_type} must not define an #[init] method when registered with a factory"
59 )]
60 FactoryRegistrationWithInit {
61 workflow_type: String,
63 },
64}
65
66#[derive(Default, Clone)]
68pub struct WorkflowDefinitions {
69 workflows: HashMap<String, RegisteredWorkflow>,
70}
71
72impl WorkflowDefinitions {
73 #[cfg(feature = "experimental")]
75 pub(crate) fn extend(&mut self, other: &Self) -> Result<(), WorkflowRegistrationError> {
76 for workflow in other.workflows.values() {
77 self.insert_workflow(workflow.definition.clone(), workflow.factory.clone())?;
78 }
79 Ok(())
80 }
81
82 pub fn new() -> Self {
84 Self::default()
85 }
86
87 pub fn register_workflow<W: WorkflowImplementation>(
91 &mut self,
92 ) -> Result<&mut Self, WorkflowRegistrationError>
93 where
94 <W::Run as WorkflowDefinition>::Input: Send,
95 {
96 let factory = Arc::new(move |input| {
97 let (payloads, payload_converter, base_ctx) = workflow_input_parts(input);
98 GuestWorkflowInstance::<W>::instantiate(payloads, payload_converter, base_ctx)
99 .context("Failed to instantiate native workflow")
100 });
101 self.insert_workflow(W::definition(), factory)?;
102 Ok(self)
103 }
104
105 pub fn register_workflow_run_with_factory<W, F>(
110 &mut self,
111 user_factory: F,
112 ) -> Result<&mut Self, WorkflowRegistrationError>
113 where
114 W: WorkflowImplementation,
115 <W::Run as WorkflowDefinition>::Input: Send,
116 F: Fn() -> W + Send + Sync + 'static,
117 {
118 if W::HAS_INIT {
119 return Err(WorkflowRegistrationError::FactoryRegistrationWithInit {
120 workflow_type: W::definition().workflow_type,
121 });
122 }
123
124 let factory = Arc::new(move |input| {
125 let (payloads, payload_converter, base_ctx) = workflow_input_parts(input);
126 let context_data =
127 SerializationContextData::Workflow(WorkflowSerializationContext::new());
128 let ser_ctx = SerializationContext::new(&context_data, &payload_converter);
129 let input: <W::Run as WorkflowDefinition>::Input =
130 payload_converter.from_payloads(&ser_ctx, payloads)?;
131
132 let workflow = user_factory();
133 Ok(Box::new(GuestWorkflowInstance::<W>::new_with_workflow(
134 workflow,
135 base_ctx,
136 Some(input),
137 )) as Box<dyn WorkflowInstance>)
138 });
139
140 self.insert_workflow(W::definition(), factory)?;
141 Ok(self)
142 }
143
144 pub fn is_empty(&self) -> bool {
146 self.workflows.is_empty()
147 }
148
149 pub(crate) fn insert_workflow(
150 &mut self,
151 definition: WorkflowDefinitionDescriptor,
152 factory: WorkflowExecutionFactory,
153 ) -> Result<(), WorkflowRegistrationError> {
154 let workflow_type = definition.workflow_type.clone();
155 if self.workflows.contains_key(&workflow_type) {
156 return Err(WorkflowRegistrationError::DuplicateWorkflowType { workflow_type });
157 }
158 self.workflows.insert(
159 workflow_type,
160 RegisteredWorkflow {
161 definition,
162 factory,
163 },
164 );
165 Ok(())
166 }
167
168 pub(crate) fn get_workflow(&self, workflow_type: &str) -> Option<WorkflowExecutionFactory> {
169 self.workflows
170 .get(workflow_type)
171 .map(|wf| wf.factory.clone())
172 }
173
174 pub fn workflow_definitions(&self) -> impl Iterator<Item = &WorkflowDefinitionDescriptor> + '_ {
176 self.workflows.values().map(|wf| &wf.definition)
177 }
178}
179
180fn workflow_input_parts(
181 input: WorkflowExecutionInput,
182) -> (Vec<Payload>, PayloadConverter, BaseWorkflowContext) {
183 let WorkflowExecutionInput {
184 namespace,
185 task_queue,
186 run_id,
187 init_workflow_job,
188 data_converter,
189 host,
190 patch_activation_callback,
191 workflow_interceptor_constructors,
192 } = input;
193 let payloads = init_workflow_job.arguments.clone();
194 let payload_converter = data_converter.payload_converter().clone();
195 let init = WorkflowInit {
196 namespace,
197 task_queue,
198 run_id,
199 initialize_workflow: init_workflow_job,
200 };
201 let base_ctx = BaseWorkflowContext::from_raw(
202 init,
203 data_converter,
204 host,
205 patch_activation_callback,
206 workflow_interceptor_constructors,
207 );
208 (payloads, payload_converter, base_ctx)
209}
210
211impl Debug for WorkflowDefinitions {
212 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
213 f.debug_struct("WorkflowDefinitions")
214 .field("workflows", &self.workflows.keys().collect::<Vec<_>>())
215 .finish()
216 }
217}