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#[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 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 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 pub fn contains_key(&self, key: &str) -> bool {
56 self.raw.fields.contains_key(key)
57 }
58
59 pub fn len(&self) -> usize {
61 self.raw.fields.len()
62 }
63
64 pub fn is_empty(&self) -> bool {
66 self.raw.fields.is_empty()
67 }
68
69 pub fn keys(&self) -> impl Iterator<Item = &str> {
71 self.ordered_keys.keys().map(String::as_str)
72 }
73
74 pub fn raw_value(&self, key: &str) -> Option<&Payload> {
76 self.raw.fields.get(key)
77 }
78
79 pub fn raw(&self) -> &ProtoMemo {
81 &self.raw
82 }
83
84 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#[derive(Clone, derive_more::Debug)]
111#[non_exhaustive]
112pub struct MemoValue {
113 #[debug(skip)]
114 value: Arc<dyn SerializableMemoValue>,
115}
116
117impl MemoValue {
118 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#[derive(Clone, Debug, Default)]
137#[non_exhaustive]
138pub struct MemoValues {
139 values: BTreeMap<String, MemoValue>,
140}
141
142impl MemoValues {
143 pub fn new() -> Self {
145 Self::default()
146 }
147
148 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 pub fn get(&self, key: &str) -> Option<&MemoValue> {
159 self.values.get(key)
160 }
161
162 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}