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