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