Skip to main content

temporalio_common_wasm/
memo.rs

1use crate::{
2    data_converters::{
3        GenericPayloadConverter, PayloadConversionError, PayloadConverter, SerializationContext,
4        SerializationContextData, TemporalDeserializable, TemporalSerializable,
5    },
6    protos::temporal::api::common::v1::{Memo as ProtoMemo, Payload},
7};
8use std::{collections::BTreeMap, sync::Arc};
9
10/// A collection of memo payloads that can be deserialized into typed values.
11#[derive(Clone, Debug)]
12#[non_exhaustive]
13pub struct Memo {
14    raw: ProtoMemo,
15    ordered_keys: BTreeMap<String, ()>,
16    payload_converter: PayloadConverter,
17    context: SerializationContextData,
18}
19
20impl Memo {
21    /// Construct a memo with the payload converter and serialization context associated with its
22    /// source.
23    pub fn from_raw(
24        raw: Option<ProtoMemo>,
25        payload_converter: PayloadConverter,
26        context: SerializationContextData,
27    ) -> Self {
28        let raw = raw.unwrap_or_default();
29        let ordered_keys = raw.fields.keys().cloned().map(|key| (key, ())).collect();
30        Self {
31            raw,
32            ordered_keys,
33            payload_converter,
34            context,
35        }
36    }
37
38    /// Decode a memo value as `T`, returning `None` when the key is absent.
39    pub fn get<T: TemporalDeserializable + 'static>(
40        &self,
41        key: &str,
42    ) -> Result<Option<T>, PayloadConversionError> {
43        let Some(payload) = self.raw.fields.get(key) else {
44            return Ok(None);
45        };
46        self.payload_converter
47            .from_payload(
48                &SerializationContext::new(&self.context, &self.payload_converter),
49                payload.clone(),
50            )
51            .map(Some)
52    }
53
54    /// Returns whether the memo contains `key`.
55    pub fn contains_key(&self, key: &str) -> bool {
56        self.raw.fields.contains_key(key)
57    }
58
59    /// Returns the number of memo entries.
60    pub fn len(&self) -> usize {
61        self.raw.fields.len()
62    }
63
64    /// Returns whether the memo has no entries.
65    pub fn is_empty(&self) -> bool {
66        self.raw.fields.is_empty()
67    }
68
69    /// Iterates over memo keys in lexicographic order.
70    pub fn keys(&self) -> impl Iterator<Item = &str> {
71        self.ordered_keys.keys().map(String::as_str)
72    }
73
74    /// Returns the underlying payload without applying payload conversion.
75    pub fn raw_value(&self, key: &str) -> Option<&Payload> {
76        self.raw.fields.get(key)
77    }
78
79    /// Access the underlying memo protobuf.
80    pub fn raw(&self) -> &ProtoMemo {
81        &self.raw
82    }
83
84    /// Consume this wrapper and return the underlying memo protobuf.
85    pub fn into_raw(self) -> ProtoMemo {
86        self.raw
87    }
88}
89
90trait SerializableMemoValue: Send + Sync {
91    fn to_payload(
92        &self,
93        context: &SerializationContext<'_>,
94    ) -> Result<Payload, PayloadConversionError>;
95}
96
97impl<T> SerializableMemoValue for T
98where
99    T: TemporalSerializable + Send + Sync + 'static,
100{
101    fn to_payload(
102        &self,
103        context: &SerializationContext<'_>,
104    ) -> Result<Payload, PayloadConversionError> {
105        context.converter.to_payload(context, self)
106    }
107}
108
109/// A typed value used in a workflow memo update.
110#[derive(Clone, derive_more::Debug)]
111#[non_exhaustive]
112pub struct MemoValue {
113    #[debug(skip)]
114    value: Arc<dyn SerializableMemoValue>,
115}
116
117impl MemoValue {
118    /// Create a memo value that will be serialized with the workflow's data converter.
119    pub fn new<T: TemporalSerializable + Send + Sync + 'static>(value: T) -> Self {
120        Self {
121            value: Arc::new(value),
122        }
123    }
124}
125
126impl TemporalSerializable for MemoValue {
127    fn to_payload(
128        &self,
129        context: &SerializationContext<'_>,
130    ) -> Result<Payload, PayloadConversionError> {
131        self.value.to_payload(context)
132    }
133}
134
135/// A complete set of memo values for a new workflow execution.
136#[derive(Clone, Debug, Default)]
137#[non_exhaustive]
138pub struct MemoValues {
139    values: BTreeMap<String, MemoValue>,
140}
141
142impl MemoValues {
143    /// Create an empty set of memo values.
144    pub fn new() -> Self {
145        Self::default()
146    }
147
148    /// Add or replace a memo value.
149    pub fn insert<T>(&mut self, key: impl Into<String>, value: T) -> &mut Self
150    where
151        T: TemporalSerializable + Send + Sync + 'static,
152    {
153        self.values.insert(key.into(), MemoValue::new(value));
154        self
155    }
156
157    /// Returns the value for `key`, if present.
158    pub fn get(&self, key: &str) -> Option<&MemoValue> {
159        self.values.get(key)
160    }
161
162    /// Iterates over the memo entries in key order.
163    pub fn iter(&self) -> impl Iterator<Item = (&str, &MemoValue)> {
164        self.values.iter().map(|(key, value)| (key.as_str(), value))
165    }
166}
167
168#[cfg(test)]
169mod tests {
170    use super::*;
171    use crate::data_converters::WorkflowSerializationContext;
172    use std::collections::HashMap;
173
174    #[test]
175    fn memo_decodes_serialized_values() {
176        let payload_converter = PayloadConverter::default();
177        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
178        let context = SerializationContext::new(&context_data, &payload_converter);
179        let payload = payload_converter.to_payload(&context, &7_u32).unwrap();
180        let raw = ProtoMemo {
181            fields: HashMap::from([("count".to_owned(), payload.clone())]),
182        };
183        let memo = Memo::from_raw(
184            Some(raw.clone()),
185            payload_converter,
186            SerializationContextData::Workflow(WorkflowSerializationContext::new()),
187        );
188
189        assert_eq!(memo.get::<u32>("count").unwrap(), Some(7));
190        assert_eq!(memo.get::<u32>("missing").unwrap(), None);
191        assert_eq!(memo.raw_value("count"), Some(&payload));
192        assert_eq!(memo.into_raw(), raw);
193    }
194
195    #[test]
196    fn memo_reports_deserialization_errors() {
197        let payload_converter = PayloadConverter::default();
198        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
199        let context = SerializationContext::new(&context_data, &payload_converter);
200        let payload = payload_converter.to_payload(&context, &7_u32).unwrap();
201        let memo = Memo::from_raw(
202            Some(ProtoMemo {
203                fields: HashMap::from([("count".to_owned(), payload)]),
204            }),
205            payload_converter,
206            SerializationContextData::Workflow(WorkflowSerializationContext::new()),
207        );
208
209        assert!(memo.get::<String>("count").is_err());
210    }
211
212    #[test]
213    fn memo_keys_have_replay_stable_order() {
214        let memo = Memo::from_raw(
215            Some(ProtoMemo {
216                fields: HashMap::from([
217                    ("zebra".to_owned(), Payload::default()),
218                    ("alpha".to_owned(), Payload::default()),
219                    ("middle".to_owned(), Payload::default()),
220                ]),
221            }),
222            PayloadConverter::default(),
223            SerializationContextData::Workflow(WorkflowSerializationContext::new()),
224        );
225
226        assert_eq!(
227            memo.keys().collect::<Vec<_>>(),
228            vec!["alpha", "middle", "zebra"]
229        );
230    }
231
232    #[test]
233    fn memo_values_serialize_heterogeneous_values() {
234        let payload_converter = PayloadConverter::default();
235        let mut values = MemoValues::new();
236        values
237            .insert("count", 7_u32)
238            .insert("label", "hello".to_string());
239
240        let context_data = SerializationContextData::Workflow(WorkflowSerializationContext::new());
241        let context = SerializationContext::new(&context_data, &payload_converter);
242        let fields = values
243            .iter()
244            .map(|(key, value)| {
245                (
246                    key.to_owned(),
247                    payload_converter.to_payload(&context, value).unwrap(),
248                )
249            })
250            .collect();
251
252        let memo = Memo::from_raw(
253            Some(ProtoMemo { fields }),
254            payload_converter.clone(),
255            SerializationContextData::Workflow(WorkflowSerializationContext::new()),
256        );
257
258        assert_eq!(memo.get::<u32>("count").unwrap(), Some(7));
259        assert_eq!(
260            memo.get::<String>("label").unwrap(),
261            Some("hello".to_string())
262        );
263    }
264}