Skip to main content

pure_stage/
serde.rs

1#![expect(clippy::borrowed_box)]
2// Copyright 2025 PRAGMA
3//
4// Licensed under the Apache License, Version 2.0 (the "License");
5// you may not use this file except in compliance with the License.
6// You may obtain a copy of the License at
7//
8//     http://www.apache.org/licenses/LICENSE-2.0
9//
10// Unless required by applicable law or agreed to in writing, software
11// distributed under the License is distributed on an "AS IS" BASIS,
12// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13// See the License for the specific language governing permissions and
14// limitations under the License.
15
16//! This module contains some serialization and deserialization code for the Pure Stage library.
17
18use std::{cell::RefCell, fmt};
19
20use cbor4ii::{
21    core::{Value, utils::BufWriter},
22    serde::{from_slice, to_writer},
23};
24use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Error};
25
26use crate::SendData;
27
28/// Helper type to wrap futures/functions/etc. and thus avoid having to handroll
29/// a `Debug` implementation for a type containing the wrapped value.
30#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
31pub struct NoDebug<T>(pub T);
32impl<T> NoDebug<T> {
33    pub fn new(t: T) -> Self {
34        Self(t)
35    }
36
37    pub fn into_inner(self) -> T {
38        self.0
39    }
40}
41impl<T> std::ops::Deref for NoDebug<T> {
42    type Target = T;
43    fn deref(&self) -> &Self::Target {
44        &self.0
45    }
46}
47impl<T> std::ops::DerefMut for NoDebug<T> {
48    fn deref_mut(&mut self) -> &mut Self::Target {
49        &mut self.0
50    }
51}
52impl<T> std::fmt::Debug for NoDebug<T> {
53    fn fmt(&self, _f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54        Ok(())
55    }
56}
57
58/// Diverging helper used in type contexts that expect `!`.
59/// Panics intentionally; for non-resolving futures, use `std::future::pending()`.
60#[cold]
61#[inline(never)]
62#[track_caller]
63#[expect(clippy::panic)]
64pub fn never() -> ! {
65    panic!("unreachable: never")
66}
67
68/// `#[serde(with = "pure_stage::serde::serialize_error")]` for serializing [`anyhow::Error`].
69pub mod serialize_error {
70    use super::*;
71
72    pub fn serialize<S: Serializer>(error: &anyhow::Error, serializer: S) -> Result<S::Ok, S::Error> {
73        serializer.serialize_str(error.to_string().as_str())
74    }
75
76    pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<anyhow::Error, D::Error> {
77        let s = String::deserialize(deserializer)?;
78        Ok(anyhow::Error::msg(s))
79    }
80}
81
82/// A trait to allow keeping collections of deserializer guards.
83pub trait DeserializerGuard {}
84pub type DeserializerGuards = Vec<Box<dyn DeserializerGuard>>;
85
86enum Field {
87    Typetag,
88    Value,
89    Ignored,
90}
91impl<'de> serde::de::Visitor<'de> for Field {
92    type Value = Self;
93
94    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
95        write!(formatter, "field identifier for SendDataValue")
96    }
97    fn visit_u64<E>(self, v: u64) -> Result<Self, E>
98    where
99        E: Error,
100    {
101        match v {
102            0 => Ok(Field::Typetag),
103            1 => Ok(Field::Value),
104            _ => Ok(Field::Ignored),
105        }
106    }
107    fn visit_str<E>(self, v: &str) -> Result<Self, E>
108    where
109        E: Error,
110    {
111        match v {
112            "typetag" => Ok(Field::Typetag),
113            "value" => Ok(Field::Value),
114            _ => Ok(Field::Ignored),
115        }
116    }
117    fn visit_bytes<E>(self, v: &[u8]) -> Result<Self, E>
118    where
119        E: Error,
120    {
121        match v {
122            b"typetag" => Ok(Field::Typetag),
123            b"value" => Ok(Field::Value),
124            _ => Ok(Field::Ignored),
125        }
126    }
127}
128impl<'de> serde::de::Deserialize<'de> for Field {
129    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
130    where
131        D: Deserializer<'de>,
132    {
133        deserializer.deserialize_any(Field::Ignored)
134    }
135}
136
137/// `#[serde(with = "pure_stage::serde::serialize_send_data")]` for serializing [`Box<dyn SendData>`](crate::SendData).
138#[allow(clippy::disallowed_types)]
139pub mod serialize_send_data {
140    use std::{any::type_name, collections::HashMap, fmt, sync::Arc};
141
142    use serde::de::Error;
143
144    use super::*;
145    use crate::SendData;
146
147    pub fn serialize<S: Serializer>(data: &Box<dyn SendData>, serializer: S) -> Result<S::Ok, S::Error> {
148        data.serialize(serializer)
149    }
150
151    pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Box<dyn SendData>, D::Error> {
152        deserializer.deserialize_struct("SendData", &["typetag", "value"], Visitor)
153    }
154
155    /// TODO(network): needs proper docs
156    pub fn register_data_deserializer<T: SendData + serde::de::DeserializeOwned>() -> DropGuard {
157        let name = type_name::<T>();
158        TYPES.with_borrow_mut(|types| {
159            types.insert(
160                name.to_string(),
161                Arc::new(|deserializer| {
162                    let value = T::deserialize(deserializer)?;
163                    Ok(Box::new(value))
164                }),
165            );
166        });
167        DropGuard(name)
168    }
169
170    pub struct DropGuard(&'static str);
171    impl DeserializerGuard for DropGuard {}
172    impl DropGuard {
173        pub fn boxed(self) -> Box<dyn DeserializerGuard> {
174            Box::new(self)
175        }
176    }
177
178    impl Drop for DropGuard {
179        fn drop(&mut self) {
180            TYPES.with_borrow_mut(|types| {
181                types.remove(self.0);
182            });
183        }
184    }
185
186    type Deser = Arc<
187        dyn for<'de> Fn(&mut dyn erased_serde::Deserializer<'de>) -> Result<Box<dyn SendData>, erased_serde::Error>,
188    >;
189    thread_local! {
190        static TYPES: RefCell<HashMap<String, Deser>> = RefCell::new(HashMap::new());
191        static DESER: RefCell<Option<Deser>> = const { RefCell::new(None) };
192    }
193
194    struct Dessert(Box<dyn SendData>);
195    impl<'de> serde::de::Deserialize<'de> for Dessert {
196        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
197        where
198            D: Deserializer<'de>,
199        {
200            #[allow(clippy::expect_used)]
201            let deser = DESER.with(|deser| deser.borrow_mut().take()).expect("deser is set");
202            let mut deserializer = <dyn erased_serde::Deserializer<'de>>::erase(deserializer);
203            let value = deser(&mut deserializer).map_err(D::Error::custom)?;
204            Ok(Dessert(value))
205        }
206    }
207
208    struct Visitor;
209    impl<'de> serde::de::Visitor<'de> for Visitor {
210        type Value = Box<dyn SendData>;
211
212        fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
213            write!(formatter, "a SendData value")
214        }
215
216        fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
217        where
218            A: serde::de::SeqAccess<'de>,
219        {
220            let typetag =
221                seq.next_element::<String>()?.ok_or_else(|| serde::de::Error::invalid_length(0, &"typetag & value"))?;
222            if let Some(value) = TYPES.with(|types| {
223                if let Some(factory) = types.borrow().get(&typetag) {
224                    DESER.with_borrow_mut(|deser| deser.replace(factory.clone()));
225                    let value =
226                        seq.next_element::<Dessert>()?.ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?;
227                    Ok(Some(value.0))
228                } else {
229                    Ok(None)
230                }
231            })? {
232                return Ok(value);
233            }
234            let value = seq
235                .next_element::<cbor4ii::core::Value>()?
236                .ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?;
237            Ok(Box::new(SendDataValue { typetag, value }))
238        }
239
240        fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
241        where
242            A: serde::de::MapAccess<'de>,
243        {
244            let (Field::Typetag, typetag) =
245                map.next_entry::<Field, String>()?.ok_or_else(|| serde::de::Error::missing_field("typetag"))?
246            else {
247                return Err(serde::de::Error::custom("typetag must be encoded first"));
248            };
249            if let Some(value) = TYPES.with(|types| {
250                if let Some(factory) = types.borrow().get(&typetag) {
251                    DESER.with_borrow_mut(|deser| deser.replace(factory.clone()));
252                    let (Field::Value, value) = map
253                        .next_entry::<Field, Dessert>()?
254                        .ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?
255                    else {
256                        return Err(serde::de::Error::custom("value must be encoded after typetag"));
257                    };
258                    Ok(Some(value.0))
259                } else {
260                    Ok(None)
261                }
262            })? {
263                return Ok(value);
264            }
265            let (Field::Value, value) = map
266                .next_entry::<Field, cbor4ii::core::Value>()?
267                .ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?
268            else {
269                return Err(serde::de::Error::custom("value must be encoded after typetag"));
270            };
271            Ok(Box::new(SendDataValue { typetag, value }))
272        }
273    }
274}
275
276/// This is the wrapper representation of a [`SendData`] value after being deserialized.
277#[derive(Debug, PartialEq, serde::Serialize, serde::Deserialize)]
278pub struct SendDataValue {
279    pub typetag: String,
280    pub value: Value,
281}
282
283impl fmt::Display for SendDataValue {
284    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
285        SendDataValue::format_cbor_value(&self.value, f)
286    }
287}
288
289impl SendDataValue {
290    /// Try to format the SendDataValue CBOR value as a human-readable string.
291    fn format_cbor_value(value: &Value, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
292        match value {
293            Value::Text(s) => write!(f, "{}", s),
294            Value::Null => write!(f, "null"),
295            Value::Bool(v) => write!(f, "{v}"),
296            Value::Integer(v) => write!(f, "{v}"),
297            Value::Float(v) => write!(f, "{v}"),
298            Value::Bytes(_) => write!(f, "<bytes>"),
299            Value::Array(vs) => {
300                match SendDataValue::array_as_bytes(vs) {
301                    Some(bytes) => {
302                        write!(f, "{hash}", hash = hex::encode(bytes.as_slice()))?;
303                    }
304                    None => {
305                        write!(f, "[")?;
306                        let mut first = true;
307                        for v in vs {
308                            if first {
309                                first = false;
310                            } else {
311                                write!(f, ", ")?;
312                            }
313                            SendDataValue::format_cbor_value(v, f)?;
314                        }
315                        write!(f, "]")?;
316                    }
317                }
318                Ok(())
319            }
320            Value::Map(vs) => {
321                write!(f, "{{")?;
322                let mut first = true;
323                for (k, v) in vs {
324                    if first {
325                        first = false;
326                    } else {
327                        write!(f, ", ")?;
328                    }
329                    SendDataValue::format_cbor_value(k, f)?;
330                    write!(f, ": ")?;
331                    SendDataValue::format_cbor_value(v, f)?;
332                }
333                write!(f, "}}")?;
334                Ok(())
335            }
336            Value::Tag(_, v) => SendDataValue::format_cbor_value(v.as_ref(), f),
337            _ => Ok(()),
338        }
339    }
340
341    fn array_as_bytes(vs: &Vec<Value>) -> Option<Vec<u8>> {
342        let mut out = Vec::with_capacity(vs.len());
343        for v in vs {
344            if let Value::Integer(n) = v {
345                if *n < 0.into() || *n > (u8::MAX as u64).into() {
346                    return None;
347                } else {
348                    out.push(*n as u8);
349                }
350            } else {
351                return None;
352            }
353        }
354        Some(out)
355    }
356}
357
358impl SendDataValue {
359    pub fn new<T: SendData>(value: &T) -> Self {
360        Self::from(value as &dyn SendData)
361    }
362
363    /// Construct a boxed [`SendData`] value from a concrete type.
364    ///
365    /// This is a convenience function that serializes the value to a vector of bytes and then
366    /// deserializes it back into a boxed [`SendData`] value. It is mostly
367    /// useful in tests.
368    pub fn boxed<T: SendData>(value: &T) -> Box<dyn SendData> {
369        Box::new(Self::new(value))
370    }
371
372    pub fn from_json(tag: impl AsRef<str>, value: impl Serialize) -> Box<dyn SendData> {
373        #[expect(clippy::expect_used)]
374        Box::new(Self {
375            typetag: tag.as_ref().to_string(),
376            value: from_slice(&to_cbor(&value)).expect("round-trip serialization should not fail"),
377        })
378    }
379}
380
381impl From<&dyn SendData> for SendDataValue {
382    #[expect(clippy::expect_used)]
383    fn from(value: &dyn SendData) -> Self {
384        let mut buf = cbor4ii::serde::Serializer::new(BufWriter::new(Vec::new()));
385        value.serialize(&mut buf).expect("serialization should not fail");
386        let bytes = buf.into_inner().into_inner();
387        cbor4ii::serde::from_slice::<SendDataValue>(&bytes)
388            .expect("deserialization of serialized SendDataValue should not fail")
389    }
390}
391
392#[cfg(test)]
393mod test_send_data {
394    use cbor4ii::serde::from_slice;
395
396    use super::*;
397    use crate::SendData;
398
399    #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
400    struct TestTuple(String, u32);
401    #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
402    struct TestStruct {
403        a: String,
404        b: u32,
405    }
406    #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
407    enum TestEnum {
408        A(String),
409        B(u32),
410    }
411
412    #[derive(Debug, serde::Serialize, serde::Deserialize)]
413    struct Container(#[serde(with = "serialize_send_data")] Box<dyn SendData>);
414
415    #[test]
416    fn test_tuple() {
417        let value = TestTuple("hello".to_string(), 42);
418        let c1 = Container(Box::new(value.clone()));
419        let bytes = to_cbor(&c1);
420
421        let c2: Container = from_slice(&bytes).unwrap();
422        let cv = c2.0.cast_deserialize::<TestTuple>().unwrap();
423        assert_eq!(cv, value);
424
425        let c3: Container = from_slice(&bytes).unwrap();
426        let cv = cv.deserialize_value(&*c3.0).unwrap();
427        assert_eq!(&*cv, &*c1.0);
428    }
429
430    #[test]
431    fn test_struct() {
432        let value = TestStruct { a: "hello".to_string(), b: 42 };
433        let c1 = Container(Box::new(value.clone()));
434        let bytes = to_cbor(&c1);
435
436        let c2: Container = from_slice(&bytes).unwrap();
437        let cv = c2.0.cast_deserialize::<TestStruct>().unwrap();
438        assert_eq!(cv, value);
439
440        let c3: Container = from_slice(&bytes).unwrap();
441        let cv = cv.deserialize_value(&*c3.0).unwrap();
442        assert_eq!(&*cv, &*c1.0);
443    }
444
445    #[test]
446    fn test_enum_a() {
447        let value = TestEnum::A("hello".to_string());
448        let c1 = Container(Box::new(value.clone()));
449        let bytes = to_cbor(&c1);
450
451        let c2: Container = from_slice(&bytes).unwrap();
452        let cv = c2.0.cast_deserialize::<TestEnum>().unwrap();
453        assert_eq!(cv, value);
454
455        let c3: Container = from_slice(&bytes).unwrap();
456        let cv = cv.deserialize_value(&*c3.0).unwrap();
457        assert_eq!(&*cv, &*c1.0);
458    }
459
460    #[test]
461    fn test_enum_b() {
462        let value = TestEnum::B(42);
463        let c1 = Container(Box::new(value.clone()));
464        let bytes = to_cbor(&c1);
465
466        let c2: Container = from_slice(&bytes).unwrap();
467        let cv = c2.0.cast_deserialize::<TestEnum>().unwrap();
468        assert_eq!(cv, value);
469
470        let c3: Container = from_slice(&bytes).unwrap();
471        let cv = cv.deserialize_value(&*c3.0).unwrap();
472        assert_eq!(&*cv, &*c1.0);
473    }
474}
475
476/// `#[serde(with = "pure_stage::serde::serialize_external_effect")]` for serializing [`Box<dyn ExternalEffect>`](crate::ExternalEffect).
477#[allow(clippy::disallowed_types)]
478pub mod serialize_external_effect {
479    use std::{collections::HashMap, sync::Arc};
480
481    use super::*;
482    use crate::{ExternalEffect, effect::UnknownExternalEffect};
483
484    pub fn serialize<S: Serializer>(data: &Box<dyn ExternalEffect>, serializer: S) -> Result<S::Ok, S::Error> {
485        data.serialize(serializer)
486    }
487
488    pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Box<dyn ExternalEffect>, D::Error> {
489        deserializer.deserialize_struct("SendData", &["typetag", "value"], Visitor)
490    }
491
492    pub fn register_effect_deserializer<T: ExternalEffect + serde::de::DeserializeOwned>() -> DropGuard {
493        let name = std::any::type_name::<T>();
494        TYPES.with_borrow_mut(|types| {
495            types.insert(
496                name.to_string(),
497                Arc::new(|deserializer| {
498                    let value = T::deserialize(deserializer)?;
499                    Ok(Box::new(value))
500                }),
501            );
502        });
503        DropGuard(name)
504    }
505
506    pub struct DropGuard(&'static str);
507    impl DeserializerGuard for DropGuard {}
508    impl DropGuard {
509        pub fn boxed(self) -> Box<dyn DeserializerGuard> {
510            Box::new(self)
511        }
512    }
513
514    impl Drop for DropGuard {
515        fn drop(&mut self) {
516            TYPES.with_borrow_mut(|types| {
517                types.remove(self.0);
518            });
519        }
520    }
521
522    type Deser = Arc<
523        dyn for<'de> Fn(
524            &mut dyn erased_serde::Deserializer<'de>,
525        ) -> Result<Box<dyn ExternalEffect>, erased_serde::Error>,
526    >;
527    thread_local! {
528        static TYPES: RefCell<HashMap<String, Deser>> = RefCell::new(HashMap::new());
529        static DESER: RefCell<Option<Deser>> = const { RefCell::new(None) };
530    }
531
532    struct Dessert(Box<dyn ExternalEffect>);
533    impl<'de> serde::de::Deserialize<'de> for Dessert {
534        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
535        where
536            D: Deserializer<'de>,
537        {
538            #[expect(clippy::expect_used)]
539            let deser = DESER.with(|deser| deser.borrow_mut().take()).expect("deser is set");
540            let mut deserializer = <dyn erased_serde::Deserializer<'de>>::erase(deserializer);
541            let value = deser(&mut deserializer).map_err(D::Error::custom)?;
542            Ok(Dessert(value))
543        }
544    }
545
546    struct Visitor;
547    impl<'de> serde::de::Visitor<'de> for Visitor {
548        type Value = Box<dyn ExternalEffect>;
549
550        fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
551            write!(formatter, "a ExternalEffect value")
552        }
553
554        fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
555        where
556            A: serde::de::SeqAccess<'de>,
557        {
558            let typetag =
559                seq.next_element::<String>()?.ok_or_else(|| serde::de::Error::invalid_length(0, &"typetag & value"))?;
560            if let Some(value) = TYPES.with(|types| {
561                if let Some(factory) = types.borrow().get(&typetag) {
562                    DESER.with_borrow_mut(|deser| deser.replace(factory.clone()));
563                    let value =
564                        seq.next_element::<Dessert>()?.ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?;
565                    Ok(Some(value.0))
566                } else {
567                    Ok(None)
568                }
569            })? {
570                return Ok(value);
571            }
572            let value = seq
573                .next_element::<cbor4ii::core::Value>()?
574                .ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?;
575            Ok(Box::new(UnknownExternalEffect::new(SendDataValue { typetag, value })))
576        }
577
578        fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
579        where
580            A: serde::de::MapAccess<'de>,
581        {
582            let (Field::Typetag, typetag) =
583                map.next_entry::<Field, String>()?.ok_or_else(|| serde::de::Error::missing_field("typetag"))?
584            else {
585                return Err(serde::de::Error::custom("typetag must be encoded first"));
586            };
587            if let Some(value) = TYPES.with(|types| {
588                if let Some(factory) = types.borrow().get(&typetag) {
589                    DESER.with_borrow_mut(|deser| deser.replace(factory.clone()));
590                    let (Field::Value, value) = map
591                        .next_entry::<Field, Dessert>()?
592                        .ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?
593                    else {
594                        return Err(serde::de::Error::custom("value must be encoded after typetag"));
595                    };
596                    Ok(Some(value.0))
597                } else {
598                    Ok(None)
599                }
600            })? {
601                return Ok(value);
602            }
603            let (Field::Value, value) = map
604                .next_entry::<Field, cbor4ii::core::Value>()?
605                .ok_or_else(|| serde::de::Error::invalid_length(1, &"value"))?
606            else {
607                return Err(serde::de::Error::custom("value must be encoded after typetag"));
608            };
609            Ok(Box::new(UnknownExternalEffect::new(SendDataValue { typetag, value })))
610        }
611    }
612}
613
614/// Serialize a value to a vector of bytes, using a thread-local buffer to
615/// optimize allocations.
616pub fn to_cbor<T: serde::Serialize>(value: &T) -> Vec<u8> {
617    thread_local! {
618        static BUFFER: RefCell<Vec<u8>> = const { RefCell::new(Vec::new()) };
619    }
620    BUFFER.with_borrow_mut(|buffer| {
621        #[expect(clippy::expect_used)]
622        to_writer(&mut *buffer, value).expect("serialization should not fail");
623        let ret = Vec::from(buffer.as_slice());
624        buffer.clear();
625        ret
626    })
627}
628
629/// Deserialize a value from a vector of bytes
630pub fn from_cbor<'a, T: serde::Deserialize<'a>>(value: &'a Vec<u8>) -> anyhow::Result<T> {
631    Ok(cbor4ii::serde::from_slice::<T>(value.as_slice())?)
632}