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