temporalio_common_wasm/
memo.rs1use 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#[derive(Clone, Debug)]
12#[non_exhaustive]
13pub struct Memo {
14 raw: ProtoMemo,
15 payload_converter: PayloadConverter,
16 context: SerializationContextData,
17}
18
19impl Memo {
20 #[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 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 pub fn contains_key(&self, key: &str) -> bool {
53 self.raw.fields.contains_key(key)
54 }
55
56 pub fn len(&self) -> usize {
58 self.raw.fields.len()
59 }
60
61 pub fn is_empty(&self) -> bool {
63 self.raw.fields.is_empty()
64 }
65
66 pub fn keys(&self) -> impl Iterator<Item = &str> {
68 self.raw.fields.keys().map(String::as_str)
69 }
70
71 pub fn raw_value(&self, key: &str) -> Option<&Payload> {
73 self.raw.fields.get(key)
74 }
75
76 pub fn raw(&self) -> &ProtoMemo {
78 &self.raw
79 }
80
81 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#[derive(Clone, derive_more::Debug)]
108#[non_exhaustive]
109pub struct MemoValue {
110 #[debug(skip)]
111 value: Arc<dyn SerializableMemoValue>,
112}
113
114impl MemoValue {
115 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#[derive(Clone, Debug, Default)]
134#[non_exhaustive]
135pub struct MemoValues {
136 values: BTreeMap<String, MemoValue>,
137}
138
139impl MemoValues {
140 pub fn new() -> Self {
142 Self::default()
143 }
144
145 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 pub fn get(&self, key: &str) -> Option<&MemoValue> {
156 self.values.get(key)
157 }
158
159 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}