Skip to main content

temporalio_common_wasm/
data_converters.rs

1//! Contains traits for and default implementations of data converters, codecs, and other
2//! serialization related functionality.
3
4mod failure_converter;
5mod well_known;
6
7pub use failure_converter::{
8    ActivityExecutionDecodeHint, CancelExternalWorkflowDecodeHint,
9    ChildWorkflowExecutionDecodeHint, ChildWorkflowStartDecodeHint, CommonAttributes,
10    DefaultFailureConverter, FailureConverter, FailureDecodeHint, NoopDecodeHint,
11    WorkflowSignalDecodeHint,
12};
13use well_known::{BINARY_NULL_ENCODING_VAL, WellKnownType, binary_null_payload};
14
15use crate::protos::{ENCODING_PAYLOAD_KEY, JSON_ENCODING_VAL, temporal::api::common::v1::Payload};
16use futures::{FutureExt, future::BoxFuture};
17use std::{collections::HashMap, sync::Arc};
18
19const PROTOBUF_ENCODING_VAL: &str = "binary/protobuf";
20
21/// Combines a [`PayloadConverter`], [`FailureConverter`], and [`PayloadCodec`] to handle all
22/// serialization needs for communicating with the Temporal server.
23#[derive(Clone)]
24pub struct DataConverter {
25    payload_converter: PayloadConverter,
26    #[allow(dead_code)] // Will be used for failure conversion
27    failure_converter: Arc<dyn FailureConverter + Send + Sync>,
28    codec: Arc<dyn PayloadCodec + Send + Sync>,
29}
30
31impl std::fmt::Debug for DataConverter {
32    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
33        f.debug_struct("DataConverter")
34            .field("payload_converter", &self.payload_converter)
35            .finish_non_exhaustive()
36    }
37}
38
39impl DataConverter {
40    /// Create a new DataConverter with the given payload converter, failure converter, and codec.
41    pub fn new(
42        payload_converter: PayloadConverter,
43        failure_converter: impl FailureConverter + Send + Sync + 'static,
44        codec: impl PayloadCodec + Send + Sync + 'static,
45    ) -> Self {
46        Self {
47            payload_converter,
48            failure_converter: Arc::new(failure_converter),
49            codec: Arc::new(codec),
50        }
51    }
52
53    /// Serialize a value into a single payload, applying the codec.
54    pub async fn to_payload<T: TemporalSerializable + 'static>(
55        &self,
56        data: &SerializationContextData,
57        val: &T,
58    ) -> Result<Payload, PayloadConversionError> {
59        let context = SerializationContext::new(data, &self.payload_converter);
60        let payload = self.payload_converter.to_payload(&context, val)?;
61        let encoded = self.codec.encode(data, vec![payload]).await?;
62        encoded
63            .into_iter()
64            .next()
65            .ok_or(PayloadConversionError::WrongEncoding)
66    }
67
68    /// Deserialize a value from a single payload, applying the codec.
69    pub async fn from_payload<T: TemporalDeserializable + 'static>(
70        &self,
71        data: &SerializationContextData,
72        payload: Payload,
73    ) -> Result<T, PayloadConversionError> {
74        let context = SerializationContext::new(data, &self.payload_converter);
75        let decoded = self.codec.decode(data, vec![payload]).await?;
76        let payload = decoded
77            .into_iter()
78            .next()
79            .ok_or(PayloadConversionError::WrongEncoding)?;
80        self.payload_converter.from_payload(&context, payload)
81    }
82
83    /// Serialize a value into multiple payloads (e.g. for multi-arg support), applying the codec.
84    pub async fn to_payloads<T: TemporalSerializable + 'static>(
85        &self,
86        data: &SerializationContextData,
87        val: &T,
88    ) -> Result<Vec<Payload>, PayloadConversionError> {
89        let context = SerializationContext::new(data, &self.payload_converter);
90        let payloads = self.payload_converter.to_payloads(&context, val)?;
91        self.codec.encode(data, payloads).await
92    }
93
94    /// Deserialize a value from multiple payloads (e.g. for multi-arg support), applying the codec.
95    pub async fn from_payloads<T: TemporalDeserializable + 'static>(
96        &self,
97        data: &SerializationContextData,
98        payloads: Vec<Payload>,
99    ) -> Result<T, PayloadConversionError> {
100        let context = SerializationContext::new(data, &self.payload_converter);
101        let decoded = self.codec.decode(data, payloads).await?;
102        self.payload_converter.from_payloads(&context, decoded)
103    }
104
105    /// Returns the payload converter component of this data converter.
106    pub fn payload_converter(&self) -> &PayloadConverter {
107        &self.payload_converter
108    }
109
110    /// Returns the failure converter component of this data converter.
111    pub fn failure_converter(&self) -> &(dyn FailureConverter + Send + Sync) {
112        self.failure_converter.as_ref()
113    }
114
115    /// Decode a Temporal failure into a caller-facing Rust error surface.
116    pub fn to_error<H: FailureDecodeHint>(
117        &self,
118        context: &SerializationContextData,
119        failure: crate::protos::temporal::api::failure::v1::Failure,
120        hint: H,
121    ) -> Result<H::Output, PayloadConversionError> {
122        let normalized =
123            self.failure_converter
124                .to_error(failure, &self.payload_converter, context)?;
125        Ok(hint.adapt(normalized))
126    }
127
128    /// Encode a typed Rust error surface into a Temporal failure.
129    pub fn to_failure(
130        &self,
131        context: &SerializationContextData,
132        error: crate::error::OutgoingError,
133    ) -> crate::protos::temporal::api::failure::v1::Failure {
134        self.failure_converter
135            .to_failure(error, &self.payload_converter, context)
136    }
137
138    /// Returns the codec component of this data converter.
139    pub fn codec(&self) -> &(dyn PayloadCodec + Send + Sync) {
140        self.codec.as_ref()
141    }
142}
143
144/// Data available when serializing in a workflow context.
145#[derive(Clone, Debug, PartialEq, Eq)]
146#[non_exhaustive]
147pub struct WorkflowSerializationContext {}
148
149#[allow(clippy::new_without_default)]
150impl WorkflowSerializationContext {
151    /// Creates an empty workflow serialization context.
152    ///
153    /// **Experimental:** This constructor may change when workflow context data is added.
154    pub fn new() -> Self {
155        Self {}
156    }
157}
158
159/// Data available when serializing in an activity context.
160#[derive(Clone, Debug, PartialEq, Eq)]
161#[non_exhaustive]
162pub struct ActivitySerializationContext {}
163
164#[allow(clippy::new_without_default)]
165impl ActivitySerializationContext {
166    /// Creates an empty activity serialization context.
167    ///
168    /// **Experimental:** This constructor may change when activity context data is added.
169    pub fn new() -> Self {
170        Self {}
171    }
172}
173
174/// Data available when serializing in a Nexus context.
175#[derive(Clone, Debug, PartialEq, Eq)]
176#[non_exhaustive]
177pub struct NexusSerializationContext {}
178
179#[allow(clippy::new_without_default)]
180impl NexusSerializationContext {
181    /// Creates an empty Nexus serialization context.
182    ///
183    /// **Experimental:** This constructor may change when Nexus context data is added.
184    pub fn new() -> Self {
185        Self {}
186    }
187}
188
189/// Data about the serialization context, indicating where the serialization is occurring.
190#[derive(Clone, Debug, PartialEq, Eq)]
191#[non_exhaustive]
192pub enum SerializationContextData {
193    /// Serialization is occurring in a workflow context.
194    Workflow(WorkflowSerializationContext),
195    /// Serialization is occurring in an activity context.
196    Activity(ActivitySerializationContext),
197    /// Serialization is occurring in a nexus context.
198    Nexus(NexusSerializationContext),
199    /// No specific serialization context.
200    None,
201}
202
203/// Context for serialization operations, including the kind of context and the
204/// payload converter for nested serialization.
205#[derive(Clone, Copy)]
206#[non_exhaustive]
207pub struct SerializationContext<'a> {
208    /// The kind of serialization context (workflow, activity, etc.).
209    pub data: &'a SerializationContextData,
210    /// Allows nested types to serialize their contents using the same converter.
211    pub converter: &'a PayloadConverter,
212}
213
214impl<'a> SerializationContext<'a> {
215    /// Creates a serialization context for the given execution context and payload converter.
216    pub fn new(data: &'a SerializationContextData, converter: &'a PayloadConverter) -> Self {
217        Self { data, converter }
218    }
219}
220
221/// Converts values to and from [`Payload`]s using different encoding strategies.
222#[derive(Clone)]
223#[non_exhaustive]
224pub enum PayloadConverter {
225    /// Uses a serde-based converter for encoding/decoding.
226    Serde(Arc<dyn ErasedSerdePayloadConverter>),
227    /// This variant signals the user wants to delegate to wrapper types
228    UseWrappers,
229    /// Tries multiple converters in order until one succeeds.
230    Composite(Arc<CompositePayloadConverter>),
231}
232
233impl std::fmt::Debug for PayloadConverter {
234    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
235        match self {
236            PayloadConverter::Serde(_) => write!(f, "PayloadConverter::Serde(...)"),
237            PayloadConverter::UseWrappers => write!(f, "PayloadConverter::UseWrappers"),
238            PayloadConverter::Composite(_) => write!(f, "PayloadConverter::Composite(...)"),
239        }
240    }
241}
242impl PayloadConverter {
243    /// Create a payload converter that uses JSON serialization via serde.
244    pub fn serde_json() -> Self {
245        Self::Serde(Arc::new(SerdeJsonPayloadConverter))
246    }
247    // TODO [rust-sdk-branch]: Proto binary, other standard built-ins
248}
249
250impl Default for PayloadConverter {
251    fn default() -> Self {
252        Self::Composite(Arc::new(CompositePayloadConverter {
253            converters: vec![Self::UseWrappers, Self::serde_json()],
254        }))
255    }
256}
257
258/// Errors that can occur during payload conversion.
259#[derive(Debug)]
260#[non_exhaustive]
261pub enum PayloadConversionError {
262    /// The payload's encoding does not match what the converter expects.
263    WrongEncoding,
264    /// An error occurred during encoding or decoding.
265    EncodingError(Box<dyn std::error::Error + Send + Sync>),
266}
267
268impl std::fmt::Display for PayloadConversionError {
269    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
270        match self {
271            PayloadConversionError::WrongEncoding => write!(f, "Wrong encoding"),
272            PayloadConversionError::EncodingError(err) => write!(f, "Encoding error: {}", err),
273        }
274    }
275}
276
277impl std::error::Error for PayloadConversionError {
278    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
279        match self {
280            PayloadConversionError::WrongEncoding => None,
281            PayloadConversionError::EncodingError(err) => Some(err.as_ref()),
282        }
283    }
284}
285
286/// Encodes and decodes payloads, enabling encryption or compression.
287///
288/// Operational codec failures should be returned as
289/// [`PayloadConversionError::EncodingError`].
290pub trait PayloadCodec {
291    /// Encode payloads before they are sent to the server.
292    fn encode(
293        &self,
294        context: &SerializationContextData,
295        payloads: Vec<Payload>,
296    ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>>;
297    /// Decode payloads after they are received from the server.
298    fn decode(
299        &self,
300        context: &SerializationContextData,
301        payloads: Vec<Payload>,
302    ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>>;
303}
304
305impl<T: PayloadCodec> PayloadCodec for Arc<T> {
306    fn encode(
307        &self,
308        context: &SerializationContextData,
309        payloads: Vec<Payload>,
310    ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
311        (**self).encode(context, payloads)
312    }
313    fn decode(
314        &self,
315        context: &SerializationContextData,
316        payloads: Vec<Payload>,
317    ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
318        (**self).decode(context, payloads)
319    }
320}
321
322/// A no-op codec that passes payloads through unchanged.
323pub struct DefaultPayloadCodec;
324
325/// Indicates some type can be serialized for use with Temporal.
326///
327/// You don't need to implement this unless you are using a non-serde-compatible custom converter,
328/// in which case you should implement the to/from_payload functions on some wrapper type.
329pub trait TemporalSerializable {
330    /// Return a reference to this value as a serde-serializable trait object.
331    fn as_serde(&self) -> Result<&dyn erased_serde::Serialize, PayloadConversionError> {
332        Err(PayloadConversionError::WrongEncoding)
333    }
334    /// Convert this value into a single [`Payload`].
335    fn to_payload(&self, _: &SerializationContext<'_>) -> Result<Payload, PayloadConversionError> {
336        Err(PayloadConversionError::WrongEncoding)
337    }
338    /// Convert to multiple payloads. Override this for types representing multiple arguments.
339    fn to_payloads(
340        &self,
341        ctx: &SerializationContext<'_>,
342    ) -> Result<Vec<Payload>, PayloadConversionError> {
343        Ok(vec![self.to_payload(ctx)?])
344    }
345}
346
347/// Indicates some type can be deserialized for use with Temporal.
348///
349/// You don't need to implement this unless you are using a non-serde-compatible custom converter,
350/// in which case you should implement the to/from_payload functions on some wrapper type.
351pub trait TemporalDeserializable: Sized {
352    /// Deserialize from a serde-based payload converter.
353    fn from_serde(
354        _: &dyn ErasedSerdePayloadConverter,
355        _ctx: &SerializationContext<'_>,
356        _: Payload,
357    ) -> Result<Self, PayloadConversionError> {
358        Err(PayloadConversionError::WrongEncoding)
359    }
360    /// Deserialize from a single [`Payload`].
361    fn from_payload(
362        ctx: &SerializationContext<'_>,
363        payload: Payload,
364    ) -> Result<Self, PayloadConversionError> {
365        let _ = (ctx, payload);
366        Err(PayloadConversionError::WrongEncoding)
367    }
368    /// Convert from multiple payloads. Override this for types representing multiple arguments.
369    fn from_payloads(
370        ctx: &SerializationContext<'_>,
371        payloads: Vec<Payload>,
372    ) -> Result<Self, PayloadConversionError> {
373        if payloads.len() != 1 {
374            return Err(PayloadConversionError::WrongEncoding);
375        }
376        Self::from_payload(ctx, payloads.into_iter().next().unwrap())
377    }
378}
379
380/// A codec-decoded set of payloads that can be deserialized later with to a user provided type.
381#[derive(Clone, Debug)]
382pub struct DecodablePayloads {
383    payloads: Vec<Payload>,
384    payload_converter: PayloadConverter,
385    context: SerializationContextData,
386}
387
388impl DecodablePayloads {
389    /// Create a new decodable payload set from raw payloads and the converter context needed to
390    /// deserialize them later.
391    pub fn new(
392        payloads: Vec<Payload>,
393        payload_converter: PayloadConverter,
394        context: SerializationContextData,
395    ) -> Self {
396        Self {
397            payloads,
398            payload_converter,
399            context,
400        }
401    }
402
403    /// Deserialize these payloads into a typed value using the stored payload converter.
404    pub fn deserialize<T: TemporalDeserializable + 'static>(
405        &self,
406    ) -> Result<T, PayloadConversionError> {
407        self.payload_converter.from_payloads(
408            &SerializationContext::new(&self.context, &self.payload_converter),
409            self.payloads.clone(),
410        )
411    }
412
413    /// Returns the underlying payloads.
414    pub fn raw(&self) -> &[Payload] {
415        &self.payloads
416    }
417
418    /// Consume this value and return the underlying payloads as a [`RawValue`].
419    pub fn into_raw(self) -> RawValue {
420        RawValue::new(self.payloads)
421    }
422}
423
424/// An unconverted set of payloads, used when the caller wants to defer deserialization.
425#[derive(Clone, Debug, Default)]
426pub struct RawValue {
427    /// The underlying payloads.
428    pub payloads: Vec<Payload>,
429}
430impl RawValue {
431    /// A RawValue representing no meaningful data, containing a single default payload.
432    /// This ensures the value can still be serialized as a single payload.
433    pub fn empty() -> Self {
434        Self {
435            payloads: vec![Payload::default()],
436        }
437    }
438
439    /// Create a new RawValue from a vector of payloads.
440    pub fn new(payloads: Vec<Payload>) -> Self {
441        Self { payloads }
442    }
443
444    /// Create a [`RawValue`] by serializing a value with the given converter.
445    pub fn from_value<T: TemporalSerializable + 'static>(
446        value: &T,
447        converter: &PayloadConverter,
448    ) -> RawValue {
449        RawValue::new(vec![
450            converter
451                .to_payload(
452                    &SerializationContext::new(&SerializationContextData::None, converter),
453                    value,
454                )
455                .unwrap(),
456        ])
457    }
458
459    /// Deserialize this [`RawValue`] into a typed value using the given converter.
460    pub fn to_value<T: TemporalDeserializable + 'static>(self, converter: &PayloadConverter) -> T {
461        converter
462            .from_payload(
463                &SerializationContext::new(&SerializationContextData::None, converter),
464                self.payloads.into_iter().next().unwrap(),
465            )
466            .unwrap()
467    }
468}
469
470impl TemporalSerializable for RawValue {
471    fn to_payload(&self, _: &SerializationContext<'_>) -> Result<Payload, PayloadConversionError> {
472        Ok(self.payloads.first().cloned().unwrap_or_default())
473    }
474    fn to_payloads(
475        &self,
476        _: &SerializationContext<'_>,
477    ) -> Result<Vec<Payload>, PayloadConversionError> {
478        Ok(self.payloads.clone())
479    }
480}
481
482impl TemporalDeserializable for RawValue {
483    fn from_payload(
484        _: &SerializationContext<'_>,
485        p: Payload,
486    ) -> Result<Self, PayloadConversionError> {
487        Ok(RawValue { payloads: vec![p] })
488    }
489    fn from_payloads(
490        _: &SerializationContext<'_>,
491        payloads: Vec<Payload>,
492    ) -> Result<Self, PayloadConversionError> {
493        Ok(RawValue { payloads })
494    }
495}
496
497/// Generic interface for converting between typed values and [`Payload`]s.
498pub trait GenericPayloadConverter {
499    /// Serialize a value into a single [`Payload`].
500    fn to_payload<T: TemporalSerializable + 'static>(
501        &self,
502        context: &SerializationContext<'_>,
503        val: &T,
504    ) -> Result<Payload, PayloadConversionError>;
505    /// Deserialize a value from a single [`Payload`].
506    #[allow(clippy::wrong_self_convention)]
507    fn from_payload<T: TemporalDeserializable + 'static>(
508        &self,
509        context: &SerializationContext<'_>,
510        payload: Payload,
511    ) -> Result<T, PayloadConversionError>;
512    /// Serialize a value into multiple [`Payload`]s.
513    fn to_payloads<T: TemporalSerializable + 'static>(
514        &self,
515        context: &SerializationContext<'_>,
516        val: &T,
517    ) -> Result<Vec<Payload>, PayloadConversionError> {
518        Ok(vec![self.to_payload(context, val)?])
519    }
520    /// Deserialize a value from multiple [`Payload`]s.
521    #[allow(clippy::wrong_self_convention)]
522    fn from_payloads<T: TemporalDeserializable + 'static>(
523        &self,
524        context: &SerializationContext<'_>,
525        payloads: Vec<Payload>,
526    ) -> Result<T, PayloadConversionError> {
527        if payloads.len() != 1 {
528            return Err(PayloadConversionError::WrongEncoding);
529        }
530        self.from_payload(context, payloads.into_iter().next().unwrap())
531    }
532}
533
534impl GenericPayloadConverter for PayloadConverter {
535    fn to_payload<T: TemporalSerializable + 'static>(
536        &self,
537        context: &SerializationContext<'_>,
538        val: &T,
539    ) -> Result<Payload, PayloadConversionError> {
540        match self {
541            PayloadConverter::Serde(pc) => {
542                if let Some(well_known_type) = WellKnownType::of::<T>() {
543                    Ok(well_known_type.to_payload(val))
544                } else {
545                    pc.to_payload(context.data, val.as_serde()?)
546                }
547            }
548            PayloadConverter::UseWrappers => T::to_payload(val, context),
549            PayloadConverter::Composite(composite) => {
550                for converter in &composite.converters {
551                    match converter.to_payload(context, val) {
552                        Ok(payload) => return Ok(payload),
553                        Err(PayloadConversionError::WrongEncoding) => continue,
554                        Err(e) => return Err(e),
555                    }
556                }
557                Err(PayloadConversionError::WrongEncoding)
558            }
559        }
560    }
561
562    fn from_payload<T: TemporalDeserializable + 'static>(
563        &self,
564        context: &SerializationContext<'_>,
565        mut payload: Payload,
566    ) -> Result<T, PayloadConversionError> {
567        match self {
568            PayloadConverter::Serde(pc) => {
569                if let Some(well_known_type) = WellKnownType::of::<T>() {
570                    payload = match well_known_type.try_from_payload(payload) {
571                        Ok(value) => return Ok(value),
572                        Err(payload) => payload,
573                    };
574                }
575                T::from_serde(pc.as_ref(), context, payload)
576            }
577            PayloadConverter::UseWrappers => T::from_payload(context, payload),
578            PayloadConverter::Composite(composite) => {
579                for converter in &composite.converters {
580                    match converter.from_payload(context, payload.clone()) {
581                        Ok(value) => return Ok(value),
582                        Err(PayloadConversionError::WrongEncoding) => continue,
583                        Err(e) => return Err(e),
584                    }
585                }
586                Err(PayloadConversionError::WrongEncoding)
587            }
588        }
589    }
590
591    fn to_payloads<T: TemporalSerializable + 'static>(
592        &self,
593        context: &SerializationContext<'_>,
594        val: &T,
595    ) -> Result<Vec<Payload>, PayloadConversionError> {
596        match self {
597            PayloadConverter::Serde(pc) => {
598                if let Some(well_known_type) = WellKnownType::of::<T>() {
599                    Ok(well_known_type.to_payloads(val))
600                } else {
601                    Ok(vec![pc.to_payload(context.data, val.as_serde()?)?])
602                }
603            }
604            PayloadConverter::UseWrappers => T::to_payloads(val, context),
605            PayloadConverter::Composite(composite) => {
606                for converter in &composite.converters {
607                    match converter.to_payloads(context, val) {
608                        Ok(payloads) => return Ok(payloads),
609                        Err(PayloadConversionError::WrongEncoding) => continue,
610                        Err(e) => return Err(e),
611                    }
612                }
613                Err(PayloadConversionError::WrongEncoding)
614            }
615        }
616    }
617
618    fn from_payloads<T: TemporalDeserializable + 'static>(
619        &self,
620        context: &SerializationContext<'_>,
621        mut payloads: Vec<Payload>,
622    ) -> Result<T, PayloadConversionError> {
623        match self {
624            PayloadConverter::Serde(pc) => {
625                if let Some(well_known_type) = WellKnownType::of::<T>() {
626                    payloads = match well_known_type.try_from_payloads(payloads) {
627                        Ok(value) => return Ok(value),
628                        Err(payloads) => payloads,
629                    };
630                }
631                if payloads.len() != 1 {
632                    return Err(PayloadConversionError::WrongEncoding);
633                }
634                let payload = payloads.into_iter().next().unwrap();
635                T::from_serde(pc.as_ref(), context, payload)
636            }
637            PayloadConverter::UseWrappers => T::from_payloads(context, payloads),
638            PayloadConverter::Composite(composite) => {
639                for converter in &composite.converters {
640                    match converter.from_payloads(context, payloads.clone()) {
641                        Ok(val) => return Ok(val),
642                        Err(PayloadConversionError::WrongEncoding) => continue,
643                        Err(e) => return Err(e),
644                    }
645                }
646                Err(PayloadConversionError::WrongEncoding)
647            }
648        }
649    }
650}
651
652// TODO [rust-sdk-branch]: Potentially allow opt-out / no-serde compile flags
653impl<T> TemporalSerializable for T
654where
655    T: serde::Serialize,
656{
657    fn as_serde(&self) -> Result<&dyn erased_serde::Serialize, PayloadConversionError> {
658        Ok(self)
659    }
660}
661impl<T> TemporalDeserializable for T
662where
663    T: serde::de::DeserializeOwned,
664{
665    fn from_serde(
666        pc: &dyn ErasedSerdePayloadConverter,
667        context: &SerializationContext<'_>,
668        payload: Payload,
669    ) -> Result<Self, PayloadConversionError>
670    where
671        Self: Sized,
672    {
673        let mut de = pc.from_payload(context.data, payload)?;
674        erased_serde::deserialize(&mut de)
675            .map_err(|e| PayloadConversionError::EncodingError(Box::new(e)))
676    }
677}
678
679struct SerdeJsonPayloadConverter;
680impl ErasedSerdePayloadConverter for SerdeJsonPayloadConverter {
681    fn to_payload(
682        &self,
683        _: &SerializationContextData,
684        value: &dyn erased_serde::Serialize,
685    ) -> Result<Payload, PayloadConversionError> {
686        let as_json = serde_json::to_vec(value)
687            .map_err(|e| PayloadConversionError::EncodingError(e.into()))?;
688        if as_json.as_slice() == b"null" {
689            return Ok(binary_null_payload());
690        }
691        Ok(Payload {
692            metadata: {
693                let mut hm = HashMap::new();
694                hm.insert(
695                    ENCODING_PAYLOAD_KEY.to_string(),
696                    JSON_ENCODING_VAL.as_bytes().to_vec(),
697                );
698                hm
699            },
700            data: as_json,
701            external_payloads: vec![],
702        })
703    }
704
705    fn from_payload(
706        &self,
707        _: &SerializationContextData,
708        payload: Payload,
709    ) -> Result<Box<dyn erased_serde::Deserializer<'static>>, PayloadConversionError> {
710        let encoding = payload
711            .metadata
712            .get(ENCODING_PAYLOAD_KEY)
713            .map(|v| v.as_slice());
714        let json_v = if encoding == Some(JSON_ENCODING_VAL.as_bytes()) {
715            serde_json::from_slice(&payload.data)
716                .map_err(|e| PayloadConversionError::EncodingError(Box::new(e)))?
717        } else if encoding == Some(BINARY_NULL_ENCODING_VAL.as_bytes()) {
718            serde_json::Value::Null
719        } else {
720            return Err(PayloadConversionError::WrongEncoding);
721        };
722        Ok(Box::new(<dyn erased_serde::Deserializer>::erase(json_v)))
723    }
724}
725/// Type-erased serde-based payload converter for use behind `dyn` trait objects.
726pub trait ErasedSerdePayloadConverter: Send + Sync {
727    /// Serialize a type-erased serde value into a [`Payload`].
728    fn to_payload(
729        &self,
730        context: &SerializationContextData,
731        value: &dyn erased_serde::Serialize,
732    ) -> Result<Payload, PayloadConversionError>;
733    /// Deserialize a [`Payload`] into a type-erased serde deserializer.
734    #[allow(clippy::wrong_self_convention)]
735    fn from_payload(
736        &self,
737        context: &SerializationContextData,
738        payload: Payload,
739    ) -> Result<Box<dyn erased_serde::Deserializer<'static>>, PayloadConversionError>;
740}
741
742// TODO [rust-sdk-branch]: All prost things should be behind a compile flag
743
744/// Wrapper for protobuf messages that implements [`TemporalSerializable`]/[`TemporalDeserializable`]
745/// using `binary/protobuf` encoding.
746pub struct ProstSerializable<T: prost::Message>(pub T);
747impl<T> TemporalSerializable for ProstSerializable<T>
748where
749    T: prost::Message + Default + 'static,
750{
751    fn to_payload(&self, _: &SerializationContext<'_>) -> Result<Payload, PayloadConversionError> {
752        let as_proto = prost::Message::encode_to_vec(&self.0);
753        Ok(Payload {
754            metadata: {
755                let mut hm = HashMap::new();
756                hm.insert(
757                    ENCODING_PAYLOAD_KEY.to_string(),
758                    PROTOBUF_ENCODING_VAL.as_bytes().to_vec(),
759                );
760                hm
761            },
762            data: as_proto,
763            external_payloads: vec![],
764        })
765    }
766}
767impl<T> TemporalDeserializable for ProstSerializable<T>
768where
769    T: prost::Message + Default + 'static,
770{
771    fn from_payload(
772        _: &SerializationContext<'_>,
773        p: Payload,
774    ) -> Result<Self, PayloadConversionError>
775    where
776        Self: Sized,
777    {
778        let encoding = p.metadata.get(ENCODING_PAYLOAD_KEY).map(|v| v.as_slice());
779        if encoding != Some(PROTOBUF_ENCODING_VAL.as_bytes()) {
780            return Err(PayloadConversionError::WrongEncoding);
781        }
782        T::decode(p.data.as_slice())
783            .map(ProstSerializable)
784            .map_err(|e| PayloadConversionError::EncodingError(Box::new(e)))
785    }
786}
787
788/// A payload converter that delegates to an ordered list of inner converters.
789#[derive(Clone)]
790pub struct CompositePayloadConverter {
791    converters: Vec<PayloadConverter>,
792}
793
794impl Default for DataConverter {
795    fn default() -> Self {
796        Self::new(
797            PayloadConverter::default(),
798            DefaultFailureConverter::default(),
799            DefaultPayloadCodec,
800        )
801    }
802}
803impl PayloadCodec for DefaultPayloadCodec {
804    fn encode(
805        &self,
806        _: &SerializationContextData,
807        payloads: Vec<Payload>,
808    ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
809        async move { Ok(payloads) }.boxed()
810    }
811    fn decode(
812        &self,
813        _: &SerializationContextData,
814        payloads: Vec<Payload>,
815    ) -> BoxFuture<'static, Result<Vec<Payload>, PayloadConversionError>> {
816        async move { Ok(payloads) }.boxed()
817    }
818}
819
820/// Represents multiple arguments for workflows/activities that accept more than one argument.
821/// Use this when interoperating with other language SDKs that allow multiple arguments.
822macro_rules! impl_multi_args {
823    ($name:ident; $count:expr; $($idx:tt: $ty:ident),+) => {
824        #[doc = concat!("Wrapper for ", stringify!($count), " typed arguments, enabling multi-arg serialization.")]
825        #[derive(Clone, Debug, PartialEq, Eq)]
826        pub struct $name<$($ty),+>($(pub $ty),+);
827
828        impl<$($ty),+> TemporalSerializable for $name<$($ty),+>
829        where
830            $($ty: TemporalSerializable + 'static),+
831        {
832            fn to_payload(&self, _: &SerializationContext<'_>) -> Result<Payload, PayloadConversionError> {
833                Err(PayloadConversionError::WrongEncoding)
834            }
835            fn to_payloads(
836                &self,
837                ctx: &SerializationContext<'_>,
838            ) -> Result<Vec<Payload>, PayloadConversionError> {
839                Ok(vec![$(ctx.converter.to_payload(ctx, &self.$idx)?),+])
840            }
841        }
842
843        #[allow(non_snake_case)]
844        impl<$($ty),+> From<($($ty),+,)> for $name<$($ty),+> {
845            fn from(t: ($($ty),+,)) -> Self {
846                $name($(t.$idx),+)
847            }
848        }
849
850        impl<$($ty),+> TemporalDeserializable for $name<$($ty),+>
851        where
852            $($ty: TemporalDeserializable + 'static),+
853        {
854            fn from_payload(_: &SerializationContext<'_>, _: Payload) -> Result<Self, PayloadConversionError> {
855                Err(PayloadConversionError::WrongEncoding)
856            }
857            fn from_payloads(
858                ctx: &SerializationContext<'_>,
859                payloads: Vec<Payload>,
860            ) -> Result<Self, PayloadConversionError> {
861                if payloads.len() != $count {
862                    return Err(PayloadConversionError::WrongEncoding);
863                }
864                let mut iter = payloads.into_iter();
865                Ok($name(
866                    $(ctx.converter.from_payload::<$ty>(ctx, iter.next().unwrap())?),+
867                ))
868            }
869        }
870    };
871}
872
873impl_multi_args!(MultiArgs2; 2; 0: A, 1: B);
874impl_multi_args!(MultiArgs3; 3; 0: A, 1: B, 2: C);
875impl_multi_args!(MultiArgs4; 4; 0: A, 1: B, 2: C, 3: D);
876impl_multi_args!(MultiArgs5; 5; 0: A, 1: B, 2: C, 3: D, 4: E);
877impl_multi_args!(MultiArgs6; 6; 0: A, 1: B, 2: C, 3: D, 4: E, 5: F);
878
879#[cfg(test)]
880mod tests {
881    use super::*;
882    use crate::data_converters::well_known::BINARY_PLAIN_ENCODING_VAL;
883    use rstest::rstest;
884
885    #[test]
886    fn unit_payloads_roundtrip() {
887        let converter = PayloadConverter::serde_json();
888        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
889        let ctx = SerializationContext::new(&context_data, &converter);
890
891        let payloads = converter.to_payloads(&ctx, &()).unwrap();
892        assert!(payloads.is_empty());
893
894        let result: () = converter.from_payloads(&ctx, payloads).unwrap();
895        assert_eq!(result, ());
896    }
897
898    #[rstest]
899    #[case::unit((), BINARY_NULL_ENCODING_VAL, b"")]
900    #[case::none_string(Option::<String>::None, BINARY_NULL_ENCODING_VAL, b"")]
901    #[case::some_string(
902        Some("value".to_string()),
903        JSON_ENCODING_VAL,
904        br#""value""#
905    )]
906    #[case::bytes(vec![0_u8, 1, 2, 255], BINARY_PLAIN_ENCODING_VAL, &[0, 1, 2, 255])]
907    #[case::some_bytes(
908        Some(vec![1_u8, 2, 3]),
909        BINARY_PLAIN_ENCODING_VAL,
910        &[1, 2, 3]
911    )]
912    #[case::none_bytes(Option::<Vec<u8>>::None, BINARY_NULL_ENCODING_VAL, b"")]
913    fn value_encodes_as<T>(
914        #[case] value: T,
915        #[case] expected_encoding: &str,
916        #[case] expected_data: &[u8],
917    ) where
918        T: TemporalSerializable + std::fmt::Debug + 'static,
919    {
920        let converter = PayloadConverter::default();
921        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
922        let ctx = SerializationContext::new(&context_data, &converter);
923
924        let payload = converter.to_payload(&ctx, &value).unwrap();
925
926        assert_eq!(
927            payload.metadata.get(ENCODING_PAYLOAD_KEY).unwrap(),
928            expected_encoding.as_bytes()
929        );
930        assert_eq!(payload.data, expected_data);
931    }
932
933    #[rstest]
934    #[case::unit(BINARY_NULL_ENCODING_VAL, b"", ())]
935    #[case::none_string(BINARY_NULL_ENCODING_VAL, b"", Option::<String>::None)]
936    #[case::legacy_none_string(JSON_ENCODING_VAL, b"null", Option::<String>::None)]
937    #[case::bytes(BINARY_PLAIN_ENCODING_VAL, &[0, 1, 2, 255], vec![0_u8, 1, 2, 255])]
938    #[case::legacy_bytes(JSON_ENCODING_VAL, b"[3,2,1]", vec![3_u8, 2, 1])]
939    #[case::some_bytes(
940        BINARY_PLAIN_ENCODING_VAL,
941        &[1, 2, 3],
942        Some(vec![1_u8, 2, 3])
943    )]
944    #[case::none_bytes(BINARY_NULL_ENCODING_VAL, b"", Option::<Vec<u8>>::None)]
945    #[case::legacy_some_bytes(
946        JSON_ENCODING_VAL,
947        b"[3,2,1]",
948        Some(vec![3_u8, 2, 1])
949    )]
950    #[case::legacy_none_bytes(JSON_ENCODING_VAL, b"null", Option::<Vec<u8>>::None)]
951    fn payload_decodes_as<T>(#[case] encoding: &str, #[case] data: &[u8], #[case] expected: T)
952    where
953        T: TemporalDeserializable + std::fmt::Debug + PartialEq + 'static,
954    {
955        let converter = PayloadConverter::default();
956        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
957        let ctx = SerializationContext::new(&context_data, &converter);
958
959        let actual: T = converter
960            .from_payload(
961                &ctx,
962                Payload {
963                    metadata: HashMap::from([(
964                        ENCODING_PAYLOAD_KEY.to_string(),
965                        encoding.as_bytes().to_vec(),
966                    )]),
967                    data: data.to_vec(),
968                    external_payloads: vec![],
969                },
970            )
971            .unwrap();
972        assert_eq!(actual, expected);
973    }
974
975    #[test]
976    fn use_wrappers_returns_wrong_encoding_for_standard_types() {
977        let converter = PayloadConverter::UseWrappers;
978        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
979        let ctx = SerializationContext::new(&context_data, &converter);
980
981        let result = converter.to_payload(&ctx, &());
982        assert!(
983            matches!(result, Err(PayloadConversionError::WrongEncoding)),
984            "{result:?}"
985        );
986
987        let result = converter.to_payloads(&ctx, &());
988        assert!(
989            matches!(result, Err(PayloadConversionError::WrongEncoding)),
990            "{result:?}"
991        );
992
993        let result = converter.to_payloads(&ctx, &vec![1_u8, 2, 3]);
994        assert!(
995            matches!(result, Err(PayloadConversionError::WrongEncoding)),
996            "{result:?}"
997        );
998
999        let result: Result<(), _> = converter.from_payload(&ctx, binary_null_payload());
1000        assert!(
1001            matches!(result, Err(PayloadConversionError::WrongEncoding)),
1002            "{result:?}"
1003        );
1004    }
1005
1006    #[test]
1007    fn multi_args_round_trip() {
1008        let converter = PayloadConverter::default();
1009        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
1010        let ctx = SerializationContext::new(&context_data, &converter);
1011
1012        let args = MultiArgs2("hello".to_string(), 42i32);
1013        let payloads = converter.to_payloads(&ctx, &args).unwrap();
1014        assert_eq!(payloads.len(), 2);
1015
1016        let result: MultiArgs2<String, i32> = converter.from_payloads(&ctx, payloads).unwrap();
1017        assert_eq!(result, args);
1018    }
1019
1020    #[test]
1021    fn empty_payloads_do_not_decode_as_option() {
1022        let converter = PayloadConverter::default();
1023        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
1024        let ctx = SerializationContext::new(&context_data, &converter);
1025
1026        let result: Result<Option<String>, _> = converter.from_payloads(&ctx, vec![]);
1027        assert!(matches!(result, Err(PayloadConversionError::WrongEncoding)));
1028    }
1029
1030    #[test]
1031    fn multi_args_from_tuple() {
1032        let args: MultiArgs2<String, i32> = ("hello".to_string(), 42i32).into();
1033        assert_eq!(args, MultiArgs2("hello".to_string(), 42));
1034    }
1035
1036    #[rstest]
1037    #[case::string("hello".to_string())]
1038    #[case::some_string(Some("hello".to_string()))]
1039    #[case::none_string(Option::<String>::None)]
1040    #[case::unit(())]
1041    #[case::strings(vec!["hello".to_string(), "world".to_string()])]
1042    #[case::bytes(vec![1_u8, 2, 3])]
1043    #[case::some_bytes(Some(vec![1_u8, 2, 3]))]
1044    #[case::none_bytes(Option::<Vec<u8>>::None)]
1045    fn decodable_payloads_roundtrip<T>(#[case] value: T)
1046    where
1047        T: TemporalSerializable + TemporalDeserializable + std::fmt::Debug + PartialEq + 'static,
1048    {
1049        let converter = PayloadConverter::default();
1050        let payloads = converter
1051            .to_payloads(
1052                &SerializationContext::new(
1053                    &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
1054                    &converter,
1055                ),
1056                &value,
1057            )
1058            .unwrap();
1059        let payloads = DecodablePayloads::new(
1060            payloads,
1061            converter,
1062            SerializationContextData::Workflow(WorkflowSerializationContext::new()),
1063        );
1064
1065        let result: T = payloads.deserialize().unwrap();
1066        assert_eq!(result, value);
1067    }
1068}