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 ..Default::default()
431 }],
432 ..Default::default()
433 },
434 )),
435 ..Default::default()
436 };
437
438 encode_payloads(
439 &mut completion,
440 &MarkingCodec,
441 &SerializationContextData::Workflow,
442 )
443 .await
444 .unwrap();
445
446 let status = completion.status.as_ref().unwrap();
447 let Status::Successful(success) = status else {
448 panic!("Expected successful status")
449 };
450 let cmd = &success.commands[0];
451 let CmdVariant::ScheduleActivity(schedule) = cmd.variant.as_ref().unwrap() else {
452 panic!("Expected ScheduleActivity")
453 };
454
455 assert!(is_encoded(&schedule.arguments[0]), "arg1 should be encoded");
456 assert!(is_encoded(&schedule.arguments[1]), "arg2 should be encoded");
457 assert!(
458 is_encoded(schedule.headers.get("header-key").unwrap()),
459 "header should be encoded"
460 );
461 }
462
463 #[tokio::test]
464 async fn test_decode_workflow_activation_with_initialize() {
465 let mut activation = WorkflowActivation {
466 run_id: "test-run".to_string(),
467 jobs: vec![WorkflowActivationJob {
468 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
469 workflow_type: "test-workflow".to_string(),
470 arguments: vec![make_payload("input1"), make_payload("input2")],
471 headers: {
472 let mut h = HashMap::new();
473 h.insert("header-key".to_string(), make_payload("header-value"));
474 h
475 },
476 ..Default::default()
477 })),
478 }],
479 ..Default::default()
480 };
481
482 decode_payloads(
483 &mut activation,
484 &MarkingCodec,
485 &SerializationContextData::Workflow,
486 )
487 .await
488 .unwrap();
489
490 let job = &activation.jobs[0];
491 let Variant::InitializeWorkflow(init) = job.variant.as_ref().unwrap() else {
492 panic!("Expected InitializeWorkflow")
493 };
494
495 assert!(is_decoded(&init.arguments[0]), "arg1 should be decoded");
496 assert!(is_decoded(&init.arguments[1]), "arg2 should be decoded");
497 assert!(
498 is_decoded(init.headers.get("header-key").unwrap()),
499 "header should be decoded"
500 );
501 }
502
503 #[tokio::test]
504 async fn test_decode_workflow_activation_with_resolve_activity() {
505 let mut activation = WorkflowActivation {
506 run_id: "test-run".to_string(),
507 jobs: vec![WorkflowActivationJob {
508 variant: Some(Variant::ResolveActivity(ResolveActivity {
509 seq: 1,
510 result: Some(ActivityResolution {
511 status: Some(ActivityStatus::Completed(Success {
512 result: Some(make_payload("activity-result")),
513 })),
514 }),
515 ..Default::default()
516 })),
517 }],
518 ..Default::default()
519 };
520
521 decode_payloads(
522 &mut activation,
523 &MarkingCodec,
524 &SerializationContextData::Workflow,
525 )
526 .await
527 .unwrap();
528
529 let job = &activation.jobs[0];
530 let Variant::ResolveActivity(resolve) = job.variant.as_ref().unwrap() else {
531 panic!("Expected ResolveActivity")
532 };
533 let ActivityStatus::Completed(success) =
534 resolve.result.as_ref().unwrap().status.as_ref().unwrap()
535 else {
536 panic!("Expected Completed status")
537 };
538
539 assert!(
540 is_decoded(success.result.as_ref().unwrap()),
541 "activity result should be decoded"
542 );
543 }
544
545 #[tokio::test]
546 async fn test_search_attributes_skipped_on_encode() {
547 let mut completion = WorkflowActivationCompletion {
549 run_id: "test-run".to_string(),
550 status: Some(Status::Successful(
551 crate::protos::coresdk::workflow_completion::Success {
552 commands: vec![
553 WorkflowCommand {
555 variant: Some(CmdVariant::UpsertWorkflowSearchAttributes(
556 UpsertWorkflowSearchAttributes {
557 search_attributes: Some(SearchAttributes {
558 indexed_fields: {
559 let mut sa = HashMap::new();
560 sa.insert(
561 "CustomField".to_string(),
562 make_payload("search-value"),
563 );
564 sa
565 },
566 }),
567 },
568 )),
569 ..Default::default()
570 },
571 WorkflowCommand {
573 variant: Some(CmdVariant::ContinueAsNewWorkflowExecution(
574 ContinueAsNewWorkflowExecution {
575 arguments: vec![make_payload("continue-arg")],
576 search_attributes: Some(SearchAttributes {
577 indexed_fields: {
578 let mut sa = HashMap::new();
579 sa.insert(
580 "CustomField".to_string(),
581 make_payload("continue-search-value"),
582 );
583 sa
584 },
585 }),
586 ..Default::default()
587 },
588 )),
589 ..Default::default()
590 },
591 WorkflowCommand {
593 variant: Some(CmdVariant::StartChildWorkflowExecution(
594 StartChildWorkflowExecution {
595 seq: 1,
596 workflow_type: "child-workflow".to_string(),
597 input: vec![make_payload("child-arg")],
598 search_attributes: Some(SearchAttributes {
599 indexed_fields: {
600 let mut sa = HashMap::new();
601 sa.insert(
602 "CustomField".to_string(),
603 make_payload("child-search-value"),
604 );
605 sa
606 },
607 }),
608 ..Default::default()
609 },
610 )),
611 ..Default::default()
612 },
613 ],
614 ..Default::default()
615 },
616 )),
617 ..Default::default()
618 };
619
620 encode_payloads(
621 &mut completion,
622 &MarkingCodec,
623 &SerializationContextData::Workflow,
624 )
625 .await
626 .unwrap();
627
628 let status = completion.status.as_ref().unwrap();
629 let Status::Successful(success) = status else {
630 panic!("Expected successful status")
631 };
632
633 let CmdVariant::UpsertWorkflowSearchAttributes(upsert) =
635 success.commands[0].variant.as_ref().unwrap()
636 else {
637 panic!("Expected UpsertWorkflowSearchAttributes")
638 };
639 let sa = upsert.search_attributes.as_ref().unwrap();
640 assert!(
641 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
642 "search attributes should NOT be encoded"
643 );
644
645 let CmdVariant::ContinueAsNewWorkflowExecution(continue_as_new) =
647 success.commands[1].variant.as_ref().unwrap()
648 else {
649 panic!("Expected ContinueAsNewWorkflowExecution")
650 };
651 assert!(
652 is_encoded(&continue_as_new.arguments[0]),
653 "arguments should be encoded"
654 );
655 let sa = continue_as_new.search_attributes.as_ref().unwrap();
656 assert!(
657 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
658 "search attributes should NOT be encoded"
659 );
660
661 let CmdVariant::StartChildWorkflowExecution(start_child) =
663 success.commands[2].variant.as_ref().unwrap()
664 else {
665 panic!("Expected StartChildWorkflowExecution")
666 };
667 assert!(is_encoded(&start_child.input[0]), "input should be encoded");
668 let sa = start_child.search_attributes.as_ref().unwrap();
669 assert!(
670 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
671 "search attributes should NOT be encoded"
672 );
673 }
674
675 #[tokio::test]
676 async fn test_search_attributes_skipped_on_decode() {
677 let mut response = DescribeWorkflowExecutionResponse {
678 workflow_execution_info: Some(WorkflowExecutionInfo {
679 memo: Some(Memo {
680 fields: {
681 let mut memo = HashMap::new();
682 memo.insert("tracked".to_string(), make_payload("memo-value"));
683 memo
684 },
685 }),
686 search_attributes: Some(SearchAttributes {
687 indexed_fields: {
688 let mut sa = HashMap::new();
689 sa.insert("CustomField".to_string(), make_payload("search-value"));
690 sa
691 },
692 }),
693 ..Default::default()
694 }),
695 ..Default::default()
696 };
697
698 decode_payloads(
699 &mut response,
700 &MarkingCodec,
701 &SerializationContextData::Workflow,
702 )
703 .await
704 .unwrap();
705
706 let info = response.workflow_execution_info.as_ref().unwrap();
707 assert!(
708 is_decoded(info.memo.as_ref().unwrap().fields.get("tracked").unwrap()),
709 "memo should be decoded"
710 );
711 assert!(
712 !is_decoded(
713 info.search_attributes
714 .as_ref()
715 .unwrap()
716 .indexed_fields
717 .get("CustomField")
718 .unwrap()
719 ),
720 "search attributes should NOT be decoded"
721 );
722 }
723
724 #[tokio::test]
725 async fn test_encode_single_payload() {
726 let mut payload = make_payload("test-data");
727
728 encode_payloads(
729 &mut payload,
730 &MarkingCodec,
731 &SerializationContextData::Workflow,
732 )
733 .await
734 .unwrap();
735
736 assert!(is_encoded(&payload), "single payload should be encoded");
737 }
738
739 #[tokio::test]
740 async fn test_decode_single_payload() {
741 let mut payload = make_payload("test-data");
742
743 decode_payloads(
744 &mut payload,
745 &MarkingCodec,
746 &SerializationContextData::Workflow,
747 )
748 .await
749 .unwrap();
750
751 assert!(is_decoded(&payload), "single payload should be decoded");
752 }
753
754 #[tokio::test]
755 async fn test_encode_payloads_message() {
756 let mut payloads = Payloads {
757 payloads: vec![make_payload("p1"), make_payload("p2"), make_payload("p3")],
758 };
759
760 encode_payloads(
761 &mut payloads,
762 &MarkingCodec,
763 &SerializationContextData::Workflow,
764 )
765 .await
766 .unwrap();
767
768 for (i, p) in payloads.payloads.iter().enumerate() {
769 assert!(is_encoded(p), "payload {} should be encoded", i);
770 }
771 }
772
773 #[tokio::test]
774 async fn test_codec_errors_stop_subsequent_codec_invocations() {
775 let codec = FailingCodec::default();
776 let mut activation = WorkflowActivation {
777 jobs: vec![WorkflowActivationJob {
778 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
779 arguments: vec![make_payload("input")],
780 headers: HashMap::from([("header".to_string(), make_payload("value"))]),
781 memo: Some(Memo {
782 fields: HashMap::from([("memo".to_string(), make_payload("memo-value"))]),
783 }),
784 ..Default::default()
785 })),
786 }],
787 ..Default::default()
788 };
789
790 let err = decode_payloads(&mut activation, &codec, &SerializationContextData::Workflow)
791 .await
792 .unwrap_err();
793
794 assert_eq!(err.to_string(), "Encoding error: visitor decode failed");
795 assert_eq!(codec.decode_calls.load(Ordering::SeqCst), 1);
796 }
797
798 #[tokio::test]
799 async fn test_encode_failure_encodes_application_failure_details() {
800 let mut failure = DefaultFailureConverter.to_failure(
801 OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new(
802 ApplicationFailure::builder(anyhow::anyhow!("app boom"))
803 .details(crate::data_converters::RawValue::new(vec![make_payload(
804 "detail",
805 )]))
806 .build(),
807 ))),
808 &PayloadConverter::default(),
809 &SerializationContextData::Workflow,
810 );
811
812 encode_payloads(
813 &mut failure,
814 &MarkingCodec,
815 &SerializationContextData::Workflow,
816 )
817 .await
818 .unwrap();
819
820 let Some(FailureInfo::ApplicationFailureInfo(info)) = failure.failure_info else {
821 panic!("expected application failure info")
822 };
823 assert!(is_encoded(&info.details.unwrap().payloads[0]));
824 }
825}