Skip to main content

temporalio_common_wasm/
memo.rs

1use crate::{
2    data_converters::{
3        GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext,
4        SerializationContextData, TemporalDeserializable,
5    },
6    protos::temporal::api::common::v1::{Memo as ProtoMemo, Payload},
7};
8
9/// A collection of memo payloads that can be deserialized into typed values.
10#[derive(Clone, Debug)]
11#[non_exhaustive]
12pub struct Memo {
13    raw: ProtoMemo,
14    payload_converter: PayloadConverter,
15    context: SerializationContextData,
16}
17
18impl Memo {
19    /// Construct a memo with the payload converter and serialization context associated with its
20    /// source.
21    #[doc(hidden)]
22    pub fn from_raw(
23        raw: Option<ProtoMemo>,
24        payload_converter: PayloadConverter,
25        context: SerializationContextData,
26    ) -> Self {
27        Self {
28            raw: raw.unwrap_or_default(),
29            payload_converter,
30            context,
31        }
32    }
33
34    /// Decode a memo value as `T`, returning `None` when the key is absent.
35    pub fn get<T: TemporalDeserializable + 'static>(
36        &self,
37        key: &str,
38    ) -> Result<Option<T>, PayloadConversionError> {
39        let Some(payload) = self.raw.fields.get(key) else {
40            return Ok(None);
41        };
42        self.payload_converter
43            .from_payload(
44                &SerializationContext {
45                    data: &self.context,
46                    converter: &self.payload_converter,
47                },
48                payload.clone(),
49            )
50            .map(Some)
51    }
52
53    /// Returns whether the memo contains `key`.
54    pub fn contains_key(&self, key: &str) -> bool {
55        self.raw.fields.contains_key(key)
56    }
57
58    /// Returns the number of memo entries.
59    pub fn len(&self) -> usize {
60        self.raw.fields.len()
61    }
62
63    /// Returns whether the memo has no entries.
64    pub fn is_empty(&self) -> bool {
65        self.raw.fields.is_empty()
66    }
67
68    /// Iterates over memo keys.
69    pub fn keys(&self) -> impl Iterator<Item = &str> {
70        self.raw.fields.keys().map(String::as_str)
71    }
72
73    /// Returns the underlying payload without applying payload conversion.
74    pub fn raw_value(&self, key: &str) -> Option<&Payload> {
75        self.raw.fields.get(key)
76    }
77
78    /// Access the underlying memo protobuf.
79    pub fn raw(&self) -> &ProtoMemo {
80        &self.raw
81    }
82
83    /// Consume this wrapper and return the underlying memo protobuf.
84    pub fn into_raw(self) -> ProtoMemo {
85        self.raw
86    }
87}
88
89#[cfg(test)]
90mod tests {
91    use super::*;
92    use std::collections::HashMap;
93
94    #[test]
95    fn memo_decodes_serialized_values() {
96        let payload_converter = PayloadConverter::default();
97        let context = SerializationContext {
98            data: &SerializationContextData::Workflow,
99            converter: &payload_converter,
100        };
101        let payload = payload_converter.to_payload(&context, &7_u32).unwrap();
102        let raw = ProtoMemo {
103            fields: HashMap::from([("count".to_owned(), payload.clone())]),
104        };
105        let memo = Memo::from_raw(
106            Some(raw.clone()),
107            payload_converter,
108            SerializationContextData::Workflow,
109        );
110
111        assert_eq!(memo.get::<u32>("count").unwrap(), Some(7));
112        assert_eq!(memo.get::<u32>("missing").unwrap(), None);
113        assert_eq!(memo.raw_value("count"), Some(&payload));
114        assert_eq!(memo.into_raw(), raw);
115    }
116
117    #[test]
118    fn memo_reports_deserialization_errors() {
119        let payload_converter = PayloadConverter::default();
120        let context = SerializationContext {
121            data: &SerializationContextData::Workflow,
122            converter: &payload_converter,
123        };
124        let payload = payload_converter.to_payload(&context, &7_u32).unwrap();
125        let memo = Memo::from_raw(
126            Some(ProtoMemo {
127                fields: HashMap::from([("count".to_owned(), payload)]),
128            }),
129            payload_converter,
130            SerializationContextData::Workflow,
131        );
132
133        assert!(memo.get::<String>("count").is_err());
134    }
135}