1use crate::{
8 data_converters::{PayloadCodec, PayloadConversionError, SerializationContextData},
9 protos::temporal::api::common::v1::{Payload, Payloads},
10};
11use futures::future::BoxFuture;
12
13pub struct PayloadField<'a> {
16 pub path: &'static str,
19 pub data: PayloadFieldData<'a>,
21}
22
23pub enum PayloadFieldData<'a> {
25 Single(&'a mut Payload),
27 Repeated(&'a mut Vec<Payload>),
29 Payloads(&'a mut Payloads),
31}
32
33pub trait AsyncPayloadVisitor {
35 fn visit<'a>(
37 &'a mut self,
38 field: PayloadField<'a>,
39 ) -> BoxFuture<'a, Result<(), PayloadConversionError>>;
40}
41
42pub trait PayloadVisitable: Send {
45 fn visit_payloads_mut<'a>(
49 &'a mut self,
50 visitor: &'a mut (dyn AsyncPayloadVisitor + Send),
51 ) -> BoxFuture<'a, Result<(), PayloadConversionError>>;
52}
53
54fn is_search_attributes_path(path: &str) -> bool {
57 path.contains("SearchAttributes.indexed_fields")
59}
60
61fn should_encode(path: &str) -> bool {
62 !is_search_attributes_path(path)
63}
64
65pub struct EncodeVisitor<'a> {
67 codec: &'a (dyn PayloadCodec + Send + Sync),
68 context: &'a SerializationContextData,
69}
70
71impl AsyncPayloadVisitor for EncodeVisitor<'_> {
72 fn visit<'a>(
73 &'a mut self,
74 field: PayloadField<'a>,
75 ) -> BoxFuture<'a, Result<(), PayloadConversionError>> {
76 Box::pin(async move {
77 if !should_encode(field.path) {
78 return Ok(());
79 }
80 match field.data {
81 PayloadFieldData::Single(payload) => {
82 let encoded = self
83 .codec
84 .encode(self.context, vec![std::mem::take(payload)])
85 .await?;
86 if let Some(p) = encoded.into_iter().next() {
87 *payload = p;
88 }
89 }
90 PayloadFieldData::Repeated(payloads) => {
91 *payloads = self
92 .codec
93 .encode(self.context, std::mem::take(payloads))
94 .await?;
95 }
96 PayloadFieldData::Payloads(payloads_msg) => {
97 payloads_msg.payloads = self
98 .codec
99 .encode(self.context, std::mem::take(&mut payloads_msg.payloads))
100 .await?;
101 }
102 }
103 Ok(())
104 })
105 }
106}
107
108pub struct DecodeVisitor<'a> {
110 codec: &'a (dyn PayloadCodec + Send + Sync),
111 context: &'a SerializationContextData,
112}
113
114impl AsyncPayloadVisitor for DecodeVisitor<'_> {
115 fn visit<'a>(
116 &'a mut self,
117 field: PayloadField<'a>,
118 ) -> BoxFuture<'a, Result<(), PayloadConversionError>> {
119 Box::pin(async move {
120 if !should_encode(field.path) {
121 return Ok(());
122 }
123 match field.data {
124 PayloadFieldData::Single(payload) => {
125 let decoded = self
126 .codec
127 .decode(self.context, vec![std::mem::take(payload)])
128 .await?;
129 if let Some(p) = decoded.into_iter().next() {
130 *payload = p;
131 }
132 }
133 PayloadFieldData::Repeated(payloads) => {
134 *payloads = self
135 .codec
136 .decode(self.context, std::mem::take(payloads))
137 .await?;
138 }
139 PayloadFieldData::Payloads(payloads_msg) => {
140 payloads_msg.payloads = self
141 .codec
142 .decode(self.context, std::mem::take(&mut payloads_msg.payloads))
143 .await?;
144 }
145 }
146 Ok(())
147 })
148 }
149}
150
151pub async fn encode_payloads<M: PayloadVisitable + Send>(
153 msg: &mut M,
154 codec: &(dyn PayloadCodec + Send + Sync),
155 context: &SerializationContextData,
156) -> Result<(), PayloadConversionError> {
157 let mut visitor = EncodeVisitor { codec, context };
158 msg.visit_payloads_mut(&mut visitor).await
159}
160
161pub async fn decode_payloads<M: PayloadVisitable + Send>(
163 msg: &mut M,
164 codec: &(dyn PayloadCodec + Send + Sync),
165 context: &SerializationContextData,
166) -> Result<(), PayloadConversionError> {
167 let mut visitor = DecodeVisitor { codec, context };
168 msg.visit_payloads_mut(&mut visitor).await
169}
170
171impl PayloadVisitable for Payload {
173 fn visit_payloads_mut<'a>(
174 &'a mut self,
175 visitor: &'a mut (dyn AsyncPayloadVisitor + Send),
176 ) -> BoxFuture<'a, Result<(), PayloadConversionError>> {
177 Box::pin(async move {
178 visitor
179 .visit(PayloadField {
180 path: "temporal.api.common.v1.Payload",
181 data: PayloadFieldData::Single(self),
182 })
183 .await
184 })
185 }
186}
187
188impl PayloadVisitable for Payloads {
190 fn visit_payloads_mut<'a>(
191 &'a mut self,
192 visitor: &'a mut (dyn AsyncPayloadVisitor + Send),
193 ) -> BoxFuture<'a, Result<(), PayloadConversionError>> {
194 Box::pin(async move {
195 visitor
196 .visit(PayloadField {
197 path: "temporal.api.common.v1.Payloads",
198 data: PayloadFieldData::Payloads(self),
199 })
200 .await
201 })
202 }
203}
204
205include!(concat!(env!("OUT_DIR"), "/payload_visitor_impl.rs"));
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211 use crate::protos::{
212 coresdk::{
213 activity_result::{
214 ActivityResolution, Success, activity_resolution::Status as ActivityStatus,
215 },
216 workflow_activation::{
217 InitializeWorkflow, ResolveActivity, WorkflowActivation, WorkflowActivationJob,
218 workflow_activation_job::Variant,
219 },
220 workflow_commands::{
221 ContinueAsNewWorkflowExecution, ScheduleActivity, StartChildWorkflowExecution,
222 UpsertWorkflowSearchAttributes, WorkflowCommand,
223 workflow_command::Variant as CmdVariant,
224 },
225 workflow_completion::{
226 WorkflowActivationCompletion, workflow_activation_completion::Status,
227 },
228 },
229 temporal::api::{
230 common::v1::{Memo, SearchAttributes},
231 failure::v1::failure::FailureInfo,
232 workflow::v1::WorkflowExecutionInfo,
233 workflowservice::v1::DescribeWorkflowExecutionResponse,
234 },
235 };
236 use futures::FutureExt;
237 use std::{
238 collections::HashMap,
239 sync::atomic::{AtomicUsize, Ordering},
240 };
241 use temporalio_common_wasm::{
242 data_converters::{DefaultFailureConverter, FailureConverter, PayloadConverter},
243 error::{ApplicationFailure, OutgoingError, OutgoingWorkflowError},
244 };
245
246 struct MarkingCodec;
247 impl PayloadCodec for MarkingCodec {
248 fn encode(
249 &self,
250 _: &SerializationContextData,
251 payloads: Vec<Payload>,
252 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
253 async move {
254 Ok(payloads
255 .into_iter()
256 .map(|mut p| {
257 p.metadata.insert("encoded".to_string(), b"true".to_vec());
258 p
259 })
260 .collect())
261 }
262 .boxed()
263 }
264
265 fn decode(
266 &self,
267 _: &SerializationContextData,
268 payloads: Vec<Payload>,
269 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
270 async move {
271 Ok(payloads
272 .into_iter()
273 .map(|mut p| {
274 p.metadata.insert("decoded".to_string(), b"true".to_vec());
275 p
276 })
277 .collect())
278 }
279 .boxed()
280 }
281 }
282
283 #[derive(Default)]
284 struct FailingCodec {
285 encode_calls: AtomicUsize,
286 decode_calls: AtomicUsize,
287 }
288
289 impl PayloadCodec for FailingCodec {
290 fn encode(
291 &self,
292 _: &SerializationContextData,
293 _: Vec<Payload>,
294 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
295 self.encode_calls.fetch_add(1, Ordering::SeqCst);
296 async move {
297 Err(PayloadConversionError::EncodingError(
298 "visitor encode failed".into(),
299 ))
300 }
301 .boxed()
302 }
303
304 fn decode(
305 &self,
306 _: &SerializationContextData,
307 _: Vec<Payload>,
308 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
309 self.decode_calls.fetch_add(1, Ordering::SeqCst);
310 async move {
311 Err(PayloadConversionError::EncodingError(
312 "visitor decode failed".into(),
313 ))
314 }
315 .boxed()
316 }
317 }
318
319 struct PathRecordingVisitor {
320 visited_paths: Vec<String>,
321 }
322 impl PathRecordingVisitor {
323 fn new() -> Self {
324 Self {
325 visited_paths: Vec::new(),
326 }
327 }
328
329 fn paths(&self) -> Vec<String> {
330 self.visited_paths.clone()
331 }
332 }
333
334 impl AsyncPayloadVisitor for PathRecordingVisitor {
335 fn visit<'a>(
336 &'a mut self,
337 field: PayloadField<'a>,
338 ) -> BoxFuture<'a, Result<(), PayloadConversionError>> {
339 let path = field.path.to_string();
340 self.visited_paths.push(path);
341 async move { Ok(()) }.boxed()
342 }
343 }
344
345 fn make_payload(data: &str) -> Payload {
346 Payload {
347 metadata: HashMap::new(),
348 data: data.as_bytes().to_vec(),
349 external_payloads: vec![],
350 }
351 }
352
353 fn is_encoded(p: &Payload) -> bool {
354 p.metadata.contains_key("encoded")
355 }
356
357 fn is_decoded(p: &Payload) -> bool {
358 p.metadata.contains_key("decoded")
359 }
360
361 #[tokio::test]
362 async fn test_direct_visitor_records_paths() {
363 let mut activation = WorkflowActivation {
364 run_id: "test-run".to_string(),
365 jobs: vec![WorkflowActivationJob {
366 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
367 workflow_type: "test-workflow".to_string(),
368 arguments: vec![make_payload("input1")],
369 headers: {
370 let mut h = HashMap::new();
371 h.insert("header-key".to_string(), make_payload("header-value"));
372 h
373 },
374 memo: Some(Memo {
375 fields: {
376 let mut m = HashMap::new();
377 m.insert("memo-key".to_string(), make_payload("memo-value"));
378 m
379 },
380 }),
381 ..Default::default()
382 })),
383 }],
384 ..Default::default()
385 };
386
387 let mut visitor = PathRecordingVisitor::new();
388 activation.visit_payloads_mut(&mut visitor).await.unwrap();
389
390 let paths = visitor.paths();
391 assert!(
392 paths
393 .iter()
394 .any(|p| p.contains("InitializeWorkflow.arguments")),
395 "should visit arguments, got: {:?}",
396 paths
397 );
398 assert!(
399 paths
400 .iter()
401 .any(|p| p.contains("InitializeWorkflow.headers")),
402 "should visit headers, got: {:?}",
403 paths
404 );
405 assert!(
406 paths.iter().any(|p| p.contains("Memo.fields")),
407 "should visit memo fields, got: {:?}",
408 paths
409 );
410 }
411
412 #[tokio::test]
413 async fn test_encode_workflow_activation_completion_with_schedule_activity() {
414 let mut completion = WorkflowActivationCompletion {
415 run_id: "test-run".to_string(),
416 status: Some(Status::Successful(
417 crate::protos::coresdk::workflow_completion::Success {
418 commands: vec![WorkflowCommand {
419 variant: Some(CmdVariant::ScheduleActivity(ScheduleActivity {
420 activity_id: "act-1".to_string(),
421 activity_type: "test-activity".to_string(),
422 arguments: vec![make_payload("arg1"), make_payload("arg2")],
423 headers: {
424 let mut h = HashMap::new();
425 h.insert("header-key".to_string(), make_payload("header-value"));
426 h
427 },
428 ..Default::default()
429 })),
430 user_metadata: None,
431 }],
432 ..Default::default()
433 },
434 )),
435 };
436
437 encode_payloads(
438 &mut completion,
439 &MarkingCodec,
440 &SerializationContextData::Workflow,
441 )
442 .await
443 .unwrap();
444
445 let status = completion.status.as_ref().unwrap();
446 let Status::Successful(success) = status else {
447 panic!("Expected successful status")
448 };
449 let cmd = &success.commands[0];
450 let CmdVariant::ScheduleActivity(schedule) = cmd.variant.as_ref().unwrap() else {
451 panic!("Expected ScheduleActivity")
452 };
453
454 assert!(is_encoded(&schedule.arguments[0]), "arg1 should be encoded");
455 assert!(is_encoded(&schedule.arguments[1]), "arg2 should be encoded");
456 assert!(
457 is_encoded(schedule.headers.get("header-key").unwrap()),
458 "header should be encoded"
459 );
460 }
461
462 #[tokio::test]
463 async fn test_decode_workflow_activation_with_initialize() {
464 let mut activation = WorkflowActivation {
465 run_id: "test-run".to_string(),
466 jobs: vec![WorkflowActivationJob {
467 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
468 workflow_type: "test-workflow".to_string(),
469 arguments: vec![make_payload("input1"), make_payload("input2")],
470 headers: {
471 let mut h = HashMap::new();
472 h.insert("header-key".to_string(), make_payload("header-value"));
473 h
474 },
475 ..Default::default()
476 })),
477 }],
478 ..Default::default()
479 };
480
481 decode_payloads(
482 &mut activation,
483 &MarkingCodec,
484 &SerializationContextData::Workflow,
485 )
486 .await
487 .unwrap();
488
489 let job = &activation.jobs[0];
490 let Variant::InitializeWorkflow(init) = job.variant.as_ref().unwrap() else {
491 panic!("Expected InitializeWorkflow")
492 };
493
494 assert!(is_decoded(&init.arguments[0]), "arg1 should be decoded");
495 assert!(is_decoded(&init.arguments[1]), "arg2 should be decoded");
496 assert!(
497 is_decoded(init.headers.get("header-key").unwrap()),
498 "header should be decoded"
499 );
500 }
501
502 #[tokio::test]
503 async fn test_decode_workflow_activation_with_resolve_activity() {
504 let mut activation = WorkflowActivation {
505 run_id: "test-run".to_string(),
506 jobs: vec![WorkflowActivationJob {
507 variant: Some(Variant::ResolveActivity(ResolveActivity {
508 seq: 1,
509 result: Some(ActivityResolution {
510 status: Some(ActivityStatus::Completed(Success {
511 result: Some(make_payload("activity-result")),
512 })),
513 }),
514 ..Default::default()
515 })),
516 }],
517 ..Default::default()
518 };
519
520 decode_payloads(
521 &mut activation,
522 &MarkingCodec,
523 &SerializationContextData::Workflow,
524 )
525 .await
526 .unwrap();
527
528 let job = &activation.jobs[0];
529 let Variant::ResolveActivity(resolve) = job.variant.as_ref().unwrap() else {
530 panic!("Expected ResolveActivity")
531 };
532 let ActivityStatus::Completed(success) =
533 resolve.result.as_ref().unwrap().status.as_ref().unwrap()
534 else {
535 panic!("Expected Completed status")
536 };
537
538 assert!(
539 is_decoded(success.result.as_ref().unwrap()),
540 "activity result should be decoded"
541 );
542 }
543
544 #[tokio::test]
545 async fn test_search_attributes_skipped_on_encode() {
546 let mut completion = WorkflowActivationCompletion {
548 run_id: "test-run".to_string(),
549 status: Some(Status::Successful(
550 crate::protos::coresdk::workflow_completion::Success {
551 commands: vec![
552 WorkflowCommand {
554 variant: Some(CmdVariant::UpsertWorkflowSearchAttributes(
555 UpsertWorkflowSearchAttributes {
556 search_attributes: Some(SearchAttributes {
557 indexed_fields: {
558 let mut sa = HashMap::new();
559 sa.insert(
560 "CustomField".to_string(),
561 make_payload("search-value"),
562 );
563 sa
564 },
565 }),
566 },
567 )),
568 user_metadata: None,
569 },
570 WorkflowCommand {
572 variant: Some(CmdVariant::ContinueAsNewWorkflowExecution(
573 ContinueAsNewWorkflowExecution {
574 arguments: vec![make_payload("continue-arg")],
575 search_attributes: Some(SearchAttributes {
576 indexed_fields: {
577 let mut sa = HashMap::new();
578 sa.insert(
579 "CustomField".to_string(),
580 make_payload("continue-search-value"),
581 );
582 sa
583 },
584 }),
585 ..Default::default()
586 },
587 )),
588 user_metadata: None,
589 },
590 WorkflowCommand {
592 variant: Some(CmdVariant::StartChildWorkflowExecution(
593 StartChildWorkflowExecution {
594 seq: 1,
595 workflow_type: "child-workflow".to_string(),
596 input: vec![make_payload("child-arg")],
597 search_attributes: Some(SearchAttributes {
598 indexed_fields: {
599 let mut sa = HashMap::new();
600 sa.insert(
601 "CustomField".to_string(),
602 make_payload("child-search-value"),
603 );
604 sa
605 },
606 }),
607 ..Default::default()
608 },
609 )),
610 user_metadata: None,
611 },
612 ],
613 ..Default::default()
614 },
615 )),
616 };
617
618 encode_payloads(
619 &mut completion,
620 &MarkingCodec,
621 &SerializationContextData::Workflow,
622 )
623 .await
624 .unwrap();
625
626 let status = completion.status.as_ref().unwrap();
627 let Status::Successful(success) = status else {
628 panic!("Expected successful status")
629 };
630
631 let CmdVariant::UpsertWorkflowSearchAttributes(upsert) =
633 success.commands[0].variant.as_ref().unwrap()
634 else {
635 panic!("Expected UpsertWorkflowSearchAttributes")
636 };
637 let sa = upsert.search_attributes.as_ref().unwrap();
638 assert!(
639 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
640 "search attributes should NOT be encoded"
641 );
642
643 let CmdVariant::ContinueAsNewWorkflowExecution(continue_as_new) =
645 success.commands[1].variant.as_ref().unwrap()
646 else {
647 panic!("Expected ContinueAsNewWorkflowExecution")
648 };
649 assert!(
650 is_encoded(&continue_as_new.arguments[0]),
651 "arguments should be encoded"
652 );
653 let sa = continue_as_new.search_attributes.as_ref().unwrap();
654 assert!(
655 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
656 "search attributes should NOT be encoded"
657 );
658
659 let CmdVariant::StartChildWorkflowExecution(start_child) =
661 success.commands[2].variant.as_ref().unwrap()
662 else {
663 panic!("Expected StartChildWorkflowExecution")
664 };
665 assert!(is_encoded(&start_child.input[0]), "input should be encoded");
666 let sa = start_child.search_attributes.as_ref().unwrap();
667 assert!(
668 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
669 "search attributes should NOT be encoded"
670 );
671 }
672
673 #[tokio::test]
674 async fn test_search_attributes_skipped_on_decode() {
675 let mut response = DescribeWorkflowExecutionResponse {
676 workflow_execution_info: Some(WorkflowExecutionInfo {
677 memo: Some(Memo {
678 fields: {
679 let mut memo = HashMap::new();
680 memo.insert("tracked".to_string(), make_payload("memo-value"));
681 memo
682 },
683 }),
684 search_attributes: Some(SearchAttributes {
685 indexed_fields: {
686 let mut sa = HashMap::new();
687 sa.insert("CustomField".to_string(), make_payload("search-value"));
688 sa
689 },
690 }),
691 ..Default::default()
692 }),
693 ..Default::default()
694 };
695
696 decode_payloads(
697 &mut response,
698 &MarkingCodec,
699 &SerializationContextData::Workflow,
700 )
701 .await
702 .unwrap();
703
704 let info = response.workflow_execution_info.as_ref().unwrap();
705 assert!(
706 is_decoded(info.memo.as_ref().unwrap().fields.get("tracked").unwrap()),
707 "memo should be decoded"
708 );
709 assert!(
710 !is_decoded(
711 info.search_attributes
712 .as_ref()
713 .unwrap()
714 .indexed_fields
715 .get("CustomField")
716 .unwrap()
717 ),
718 "search attributes should NOT be decoded"
719 );
720 }
721
722 #[tokio::test]
723 async fn test_encode_single_payload() {
724 let mut payload = make_payload("test-data");
725
726 encode_payloads(
727 &mut payload,
728 &MarkingCodec,
729 &SerializationContextData::Workflow,
730 )
731 .await
732 .unwrap();
733
734 assert!(is_encoded(&payload), "single payload should be encoded");
735 }
736
737 #[tokio::test]
738 async fn test_decode_single_payload() {
739 let mut payload = make_payload("test-data");
740
741 decode_payloads(
742 &mut payload,
743 &MarkingCodec,
744 &SerializationContextData::Workflow,
745 )
746 .await
747 .unwrap();
748
749 assert!(is_decoded(&payload), "single payload should be decoded");
750 }
751
752 #[tokio::test]
753 async fn test_encode_payloads_message() {
754 let mut payloads = Payloads {
755 payloads: vec![make_payload("p1"), make_payload("p2"), make_payload("p3")],
756 };
757
758 encode_payloads(
759 &mut payloads,
760 &MarkingCodec,
761 &SerializationContextData::Workflow,
762 )
763 .await
764 .unwrap();
765
766 for (i, p) in payloads.payloads.iter().enumerate() {
767 assert!(is_encoded(p), "payload {} should be encoded", i);
768 }
769 }
770
771 #[tokio::test]
772 async fn test_codec_errors_stop_subsequent_codec_invocations() {
773 let codec = FailingCodec::default();
774 let mut activation = WorkflowActivation {
775 jobs: vec![WorkflowActivationJob {
776 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
777 arguments: vec![make_payload("input")],
778 headers: HashMap::from([("header".to_string(), make_payload("value"))]),
779 memo: Some(Memo {
780 fields: HashMap::from([("memo".to_string(), make_payload("memo-value"))]),
781 }),
782 ..Default::default()
783 })),
784 }],
785 ..Default::default()
786 };
787
788 let err = decode_payloads(&mut activation, &codec, &SerializationContextData::Workflow)
789 .await
790 .unwrap_err();
791
792 assert_eq!(err.to_string(), "Encoding error: visitor decode failed");
793 assert_eq!(codec.decode_calls.load(Ordering::SeqCst), 1);
794 }
795
796 #[tokio::test]
797 async fn test_encode_failure_encodes_application_failure_details() {
798 let mut failure = DefaultFailureConverter.to_failure(
799 OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new(
800 ApplicationFailure::builder(anyhow::anyhow!("app boom"))
801 .details(crate::data_converters::RawValue::new(vec![make_payload(
802 "detail",
803 )]))
804 .build(),
805 ))),
806 &PayloadConverter::default(),
807 &SerializationContextData::Workflow,
808 );
809
810 encode_payloads(
811 &mut failure,
812 &MarkingCodec,
813 &SerializationContextData::Workflow,
814 )
815 .await
816 .unwrap();
817
818 let Some(FailureInfo::ApplicationFailureInfo(info)) = failure.failure_info else {
819 panic!("expected application failure info")
820 };
821 assert!(is_encoded(&info.details.unwrap().payloads[0]));
822 }
823}