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::{
212 data_converters::WorkflowSerializationContext,
213 protos::{
214 coresdk::{
215 activity_result::{
216 ActivityResolution, Success, activity_resolution::Status as ActivityStatus,
217 },
218 workflow_activation::{
219 InitializeWorkflow, ResolveActivity, WorkflowActivation, WorkflowActivationJob,
220 workflow_activation_job::Variant,
221 },
222 workflow_commands::{
223 ContinueAsNewWorkflowExecution, ScheduleActivity, StartChildWorkflowExecution,
224 UpsertWorkflowSearchAttributes, WorkflowCommand,
225 workflow_command::Variant as CmdVariant,
226 },
227 workflow_completion::{
228 WorkflowActivationCompletion, workflow_activation_completion::Status,
229 },
230 },
231 temporal::api::{
232 common::v1::{Memo, SearchAttributes},
233 failure::v1::failure::FailureInfo,
234 workflow::v1::WorkflowExecutionInfo,
235 workflowservice::v1::DescribeWorkflowExecutionResponse,
236 },
237 },
238 };
239 use futures::FutureExt;
240 use std::{
241 collections::HashMap,
242 sync::atomic::{AtomicUsize, Ordering},
243 };
244 use temporalio_common_wasm::{
245 data_converters::{DefaultFailureConverter, FailureConverter, PayloadConverter},
246 error::{ApplicationFailure, OutgoingError, OutgoingWorkflowError},
247 };
248
249 struct MarkingCodec;
250 impl PayloadCodec for MarkingCodec {
251 fn encode(
252 &self,
253 _: &SerializationContextData,
254 payloads: Vec<Payload>,
255 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
256 async move {
257 Ok(payloads
258 .into_iter()
259 .map(|mut p| {
260 p.metadata.insert("encoded".to_string(), b"true".to_vec());
261 p
262 })
263 .collect())
264 }
265 .boxed()
266 }
267
268 fn decode(
269 &self,
270 _: &SerializationContextData,
271 payloads: Vec<Payload>,
272 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
273 async move {
274 Ok(payloads
275 .into_iter()
276 .map(|mut p| {
277 p.metadata.insert("decoded".to_string(), b"true".to_vec());
278 p
279 })
280 .collect())
281 }
282 .boxed()
283 }
284 }
285
286 #[derive(Default)]
287 struct FailingCodec {
288 encode_calls: AtomicUsize,
289 decode_calls: AtomicUsize,
290 }
291
292 impl PayloadCodec for FailingCodec {
293 fn encode(
294 &self,
295 _: &SerializationContextData,
296 _: Vec<Payload>,
297 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
298 self.encode_calls.fetch_add(1, Ordering::SeqCst);
299 async move {
300 Err(PayloadConversionError::EncodingError(
301 "visitor encode failed".into(),
302 ))
303 }
304 .boxed()
305 }
306
307 fn decode(
308 &self,
309 _: &SerializationContextData,
310 _: Vec<Payload>,
311 ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
312 self.decode_calls.fetch_add(1, Ordering::SeqCst);
313 async move {
314 Err(PayloadConversionError::EncodingError(
315 "visitor decode failed".into(),
316 ))
317 }
318 .boxed()
319 }
320 }
321
322 struct PathRecordingVisitor {
323 visited_paths: Vec<String>,
324 }
325 impl PathRecordingVisitor {
326 fn new() -> Self {
327 Self {
328 visited_paths: Vec::new(),
329 }
330 }
331
332 fn paths(&self) -> Vec<String> {
333 self.visited_paths.clone()
334 }
335 }
336
337 impl AsyncPayloadVisitor for PathRecordingVisitor {
338 fn visit<'a>(
339 &'a mut self,
340 field: PayloadField<'a>,
341 ) -> BoxFuture<'a, Result<(), PayloadConversionError>> {
342 let path = field.path.to_string();
343 self.visited_paths.push(path);
344 async move { Ok(()) }.boxed()
345 }
346 }
347
348 fn make_payload(data: &str) -> Payload {
349 Payload {
350 metadata: HashMap::new(),
351 data: data.as_bytes().to_vec(),
352 external_payloads: vec![],
353 }
354 }
355
356 fn is_encoded(p: &Payload) -> bool {
357 p.metadata.contains_key("encoded")
358 }
359
360 fn is_decoded(p: &Payload) -> bool {
361 p.metadata.contains_key("decoded")
362 }
363
364 #[tokio::test]
365 async fn test_direct_visitor_records_paths() {
366 let mut activation = WorkflowActivation {
367 run_id: "test-run".to_string(),
368 jobs: vec![WorkflowActivationJob {
369 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
370 workflow_type: "test-workflow".to_string(),
371 arguments: vec![make_payload("input1")],
372 headers: {
373 let mut h = HashMap::new();
374 h.insert("header-key".to_string(), make_payload("header-value"));
375 h
376 },
377 memo: Some(Memo {
378 fields: {
379 let mut m = HashMap::new();
380 m.insert("memo-key".to_string(), make_payload("memo-value"));
381 m
382 },
383 }),
384 ..Default::default()
385 })),
386 }],
387 ..Default::default()
388 };
389
390 let mut visitor = PathRecordingVisitor::new();
391 activation.visit_payloads_mut(&mut visitor).await.unwrap();
392
393 let paths = visitor.paths();
394 assert!(
395 paths
396 .iter()
397 .any(|p| p.contains("InitializeWorkflow.arguments")),
398 "should visit arguments, got: {:?}",
399 paths
400 );
401 assert!(
402 paths
403 .iter()
404 .any(|p| p.contains("InitializeWorkflow.headers")),
405 "should visit headers, got: {:?}",
406 paths
407 );
408 assert!(
409 paths.iter().any(|p| p.contains("Memo.fields")),
410 "should visit memo fields, got: {:?}",
411 paths
412 );
413 }
414
415 #[tokio::test]
416 async fn test_encode_workflow_activation_completion_with_schedule_activity() {
417 let mut completion = WorkflowActivationCompletion {
418 run_id: "test-run".to_string(),
419 status: Some(Status::Successful(
420 crate::protos::coresdk::workflow_completion::Success {
421 commands: vec![WorkflowCommand {
422 variant: Some(CmdVariant::ScheduleActivity(ScheduleActivity {
423 activity_id: "act-1".to_string(),
424 activity_type: "test-activity".to_string(),
425 arguments: vec![make_payload("arg1"), make_payload("arg2")],
426 headers: {
427 let mut h = HashMap::new();
428 h.insert("header-key".to_string(), make_payload("header-value"));
429 h
430 },
431 ..Default::default()
432 })),
433 ..Default::default()
434 }],
435 ..Default::default()
436 },
437 )),
438 ..Default::default()
439 };
440
441 encode_payloads(
442 &mut completion,
443 &MarkingCodec,
444 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
445 )
446 .await
447 .unwrap();
448
449 let status = completion.status.as_ref().unwrap();
450 let Status::Successful(success) = status else {
451 panic!("Expected successful status")
452 };
453 let cmd = &success.commands[0];
454 let CmdVariant::ScheduleActivity(schedule) = cmd.variant.as_ref().unwrap() else {
455 panic!("Expected ScheduleActivity")
456 };
457
458 assert!(is_encoded(&schedule.arguments[0]), "arg1 should be encoded");
459 assert!(is_encoded(&schedule.arguments[1]), "arg2 should be encoded");
460 assert!(
461 is_encoded(schedule.headers.get("header-key").unwrap()),
462 "header should be encoded"
463 );
464 }
465
466 #[tokio::test]
467 async fn test_decode_workflow_activation_with_initialize() {
468 let mut activation = WorkflowActivation {
469 run_id: "test-run".to_string(),
470 jobs: vec![WorkflowActivationJob {
471 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
472 workflow_type: "test-workflow".to_string(),
473 arguments: vec![make_payload("input1"), make_payload("input2")],
474 headers: {
475 let mut h = HashMap::new();
476 h.insert("header-key".to_string(), make_payload("header-value"));
477 h
478 },
479 ..Default::default()
480 })),
481 }],
482 ..Default::default()
483 };
484
485 decode_payloads(
486 &mut activation,
487 &MarkingCodec,
488 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
489 )
490 .await
491 .unwrap();
492
493 let job = &activation.jobs[0];
494 let Variant::InitializeWorkflow(init) = job.variant.as_ref().unwrap() else {
495 panic!("Expected InitializeWorkflow")
496 };
497
498 assert!(is_decoded(&init.arguments[0]), "arg1 should be decoded");
499 assert!(is_decoded(&init.arguments[1]), "arg2 should be decoded");
500 assert!(
501 is_decoded(init.headers.get("header-key").unwrap()),
502 "header should be decoded"
503 );
504 }
505
506 #[tokio::test]
507 async fn test_decode_workflow_activation_with_resolve_activity() {
508 let mut activation = WorkflowActivation {
509 run_id: "test-run".to_string(),
510 jobs: vec![WorkflowActivationJob {
511 variant: Some(Variant::ResolveActivity(ResolveActivity {
512 seq: 1,
513 result: Some(ActivityResolution {
514 status: Some(ActivityStatus::Completed(Success {
515 result: Some(make_payload("activity-result")),
516 })),
517 }),
518 ..Default::default()
519 })),
520 }],
521 ..Default::default()
522 };
523
524 decode_payloads(
525 &mut activation,
526 &MarkingCodec,
527 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
528 )
529 .await
530 .unwrap();
531
532 let job = &activation.jobs[0];
533 let Variant::ResolveActivity(resolve) = job.variant.as_ref().unwrap() else {
534 panic!("Expected ResolveActivity")
535 };
536 let ActivityStatus::Completed(success) =
537 resolve.result.as_ref().unwrap().status.as_ref().unwrap()
538 else {
539 panic!("Expected Completed status")
540 };
541
542 assert!(
543 is_decoded(success.result.as_ref().unwrap()),
544 "activity result should be decoded"
545 );
546 }
547
548 #[tokio::test]
549 async fn test_search_attributes_skipped_on_encode() {
550 let mut completion = WorkflowActivationCompletion {
552 run_id: "test-run".to_string(),
553 status: Some(Status::Successful(
554 crate::protos::coresdk::workflow_completion::Success {
555 commands: vec![
556 WorkflowCommand {
558 variant: Some(CmdVariant::UpsertWorkflowSearchAttributes(
559 UpsertWorkflowSearchAttributes {
560 search_attributes: Some(SearchAttributes {
561 indexed_fields: {
562 let mut sa = HashMap::new();
563 sa.insert(
564 "CustomField".to_string(),
565 make_payload("search-value"),
566 );
567 sa
568 },
569 }),
570 },
571 )),
572 ..Default::default()
573 },
574 WorkflowCommand {
576 variant: Some(CmdVariant::ContinueAsNewWorkflowExecution(
577 ContinueAsNewWorkflowExecution {
578 arguments: vec![make_payload("continue-arg")],
579 search_attributes: Some(SearchAttributes {
580 indexed_fields: {
581 let mut sa = HashMap::new();
582 sa.insert(
583 "CustomField".to_string(),
584 make_payload("continue-search-value"),
585 );
586 sa
587 },
588 }),
589 ..Default::default()
590 },
591 )),
592 ..Default::default()
593 },
594 WorkflowCommand {
596 variant: Some(CmdVariant::StartChildWorkflowExecution(
597 StartChildWorkflowExecution {
598 seq: 1,
599 workflow_type: "child-workflow".to_string(),
600 input: vec![make_payload("child-arg")],
601 search_attributes: Some(SearchAttributes {
602 indexed_fields: {
603 let mut sa = HashMap::new();
604 sa.insert(
605 "CustomField".to_string(),
606 make_payload("child-search-value"),
607 );
608 sa
609 },
610 }),
611 ..Default::default()
612 },
613 )),
614 ..Default::default()
615 },
616 ],
617 ..Default::default()
618 },
619 )),
620 ..Default::default()
621 };
622
623 encode_payloads(
624 &mut completion,
625 &MarkingCodec,
626 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
627 )
628 .await
629 .unwrap();
630
631 let status = completion.status.as_ref().unwrap();
632 let Status::Successful(success) = status else {
633 panic!("Expected successful status")
634 };
635
636 let CmdVariant::UpsertWorkflowSearchAttributes(upsert) =
638 success.commands[0].variant.as_ref().unwrap()
639 else {
640 panic!("Expected UpsertWorkflowSearchAttributes")
641 };
642 let sa = upsert.search_attributes.as_ref().unwrap();
643 assert!(
644 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
645 "search attributes should NOT be encoded"
646 );
647
648 let CmdVariant::ContinueAsNewWorkflowExecution(continue_as_new) =
650 success.commands[1].variant.as_ref().unwrap()
651 else {
652 panic!("Expected ContinueAsNewWorkflowExecution")
653 };
654 assert!(
655 is_encoded(&continue_as_new.arguments[0]),
656 "arguments should be encoded"
657 );
658 let sa = continue_as_new.search_attributes.as_ref().unwrap();
659 assert!(
660 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
661 "search attributes should NOT be encoded"
662 );
663
664 let CmdVariant::StartChildWorkflowExecution(start_child) =
666 success.commands[2].variant.as_ref().unwrap()
667 else {
668 panic!("Expected StartChildWorkflowExecution")
669 };
670 assert!(is_encoded(&start_child.input[0]), "input should be encoded");
671 let sa = start_child.search_attributes.as_ref().unwrap();
672 assert!(
673 !is_encoded(sa.indexed_fields.get("CustomField").unwrap()),
674 "search attributes should NOT be encoded"
675 );
676 }
677
678 #[tokio::test]
679 async fn test_search_attributes_skipped_on_decode() {
680 let mut response = DescribeWorkflowExecutionResponse {
681 workflow_execution_info: Some(WorkflowExecutionInfo {
682 memo: Some(Memo {
683 fields: {
684 let mut memo = HashMap::new();
685 memo.insert("tracked".to_string(), make_payload("memo-value"));
686 memo
687 },
688 }),
689 search_attributes: Some(SearchAttributes {
690 indexed_fields: {
691 let mut sa = HashMap::new();
692 sa.insert("CustomField".to_string(), make_payload("search-value"));
693 sa
694 },
695 }),
696 ..Default::default()
697 }),
698 ..Default::default()
699 };
700
701 decode_payloads(
702 &mut response,
703 &MarkingCodec,
704 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
705 )
706 .await
707 .unwrap();
708
709 let info = response.workflow_execution_info.as_ref().unwrap();
710 assert!(
711 is_decoded(info.memo.as_ref().unwrap().fields.get("tracked").unwrap()),
712 "memo should be decoded"
713 );
714 assert!(
715 !is_decoded(
716 info.search_attributes
717 .as_ref()
718 .unwrap()
719 .indexed_fields
720 .get("CustomField")
721 .unwrap()
722 ),
723 "search attributes should NOT be decoded"
724 );
725 }
726
727 #[tokio::test]
728 async fn test_encode_single_payload() {
729 let mut payload = make_payload("test-data");
730
731 encode_payloads(
732 &mut payload,
733 &MarkingCodec,
734 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
735 )
736 .await
737 .unwrap();
738
739 assert!(is_encoded(&payload), "single payload should be encoded");
740 }
741
742 #[tokio::test]
743 async fn test_decode_single_payload() {
744 let mut payload = make_payload("test-data");
745
746 decode_payloads(
747 &mut payload,
748 &MarkingCodec,
749 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
750 )
751 .await
752 .unwrap();
753
754 assert!(is_decoded(&payload), "single payload should be decoded");
755 }
756
757 #[tokio::test]
758 async fn test_encode_payloads_message() {
759 let mut payloads = Payloads {
760 payloads: vec![make_payload("p1"), make_payload("p2"), make_payload("p3")],
761 };
762
763 encode_payloads(
764 &mut payloads,
765 &MarkingCodec,
766 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
767 )
768 .await
769 .unwrap();
770
771 for (i, p) in payloads.payloads.iter().enumerate() {
772 assert!(is_encoded(p), "payload {} should be encoded", i);
773 }
774 }
775
776 #[tokio::test]
777 async fn test_codec_errors_stop_subsequent_codec_invocations() {
778 let codec = FailingCodec::default();
779 let mut activation = WorkflowActivation {
780 jobs: vec![WorkflowActivationJob {
781 variant: Some(Variant::InitializeWorkflow(InitializeWorkflow {
782 arguments: vec![make_payload("input")],
783 headers: HashMap::from([("header".to_string(), make_payload("value"))]),
784 memo: Some(Memo {
785 fields: HashMap::from([("memo".to_string(), make_payload("memo-value"))]),
786 }),
787 ..Default::default()
788 })),
789 }],
790 ..Default::default()
791 };
792
793 let err = decode_payloads(
794 &mut activation,
795 &codec,
796 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
797 )
798 .await
799 .unwrap_err();
800
801 assert_eq!(err.to_string(), "Encoding error: visitor decode failed");
802 assert_eq!(codec.decode_calls.load(Ordering::SeqCst), 1);
803 }
804
805 #[tokio::test]
806 async fn test_encode_failure_encodes_application_failure_details() {
807 let mut failure = DefaultFailureConverter::default().to_failure(
808 OutgoingError::Workflow(OutgoingWorkflowError::Application(Box::new(
809 ApplicationFailure::builder(anyhow::anyhow!("app boom"))
810 .details(crate::data_converters::RawValue::new(vec![make_payload(
811 "detail",
812 )]))
813 .build(),
814 ))),
815 &PayloadConverter::default(),
816 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
817 );
818
819 encode_payloads(
820 &mut failure,
821 &MarkingCodec,
822 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
823 )
824 .await
825 .unwrap();
826
827 let Some(FailureInfo::ApplicationFailureInfo(info)) = failure.failure_info else {
828 panic!("expected application failure info")
829 };
830 assert!(is_encoded(&info.details.unwrap().payloads[0]));
831 }
832}