Skip to main content

rings_codec/
lib.rs

1#![deny(missing_docs)]
2
3//! Serde wire helpers backed by the Rings postcard codec.
4
5use std::error::Error as StdError;
6use std::fmt;
7
8use serde::de::DeserializeOwned;
9use serde::de::EnumAccess;
10use serde::de::Visitor;
11use serde::Deserialize;
12use serde::Serialize;
13
14/// Codec error returned by Rings wire helper functions.
15#[derive(Debug)]
16pub enum Error {
17    /// Serialization failed.
18    Serialize(postcard::Error),
19    /// Deserialization failed.
20    Deserialize(postcard::Error),
21    /// Deserialization succeeded before all bytes were consumed.
22    TrailingBytes {
23        /// Number of bytes consumed by the decoder.
24        decoded: usize,
25        /// Total number of bytes provided to the decoder.
26        total: usize,
27    },
28}
29
30impl fmt::Display for Error {
31    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
32        match self {
33            Self::Serialize(error) => write!(f, "{error}"),
34            Self::Deserialize(error) => write!(f, "{error}"),
35            Self::TrailingBytes { decoded, total } => write!(
36                f,
37                "deserializer consumed {decoded} of {total} bytes; trailing bytes remain"
38            ),
39        }
40    }
41}
42
43impl StdError for Error {
44    fn source(&self) -> Option<&(dyn StdError + 'static)> {
45        match self {
46            Self::Serialize(error) => Some(error),
47            Self::Deserialize(error) => Some(error),
48            Self::TrailingBytes { .. } => None,
49        }
50    }
51}
52
53impl From<postcard::Error> for Error {
54    fn from(error: postcard::Error) -> Self {
55        Self::Serialize(error)
56    }
57}
58
59impl Error {
60    fn deserialize(error: postcard::Error) -> Self {
61        Self::Deserialize(error)
62    }
63}
64
65/// Codec result type.
66pub type Result<T> = std::result::Result<T, Error>;
67
68/// Serialize a serde value using the Rings wire encoding.
69pub fn serialize<T>(value: &T) -> Result<Vec<u8>>
70where T: Serialize {
71    postcard::to_allocvec(value).map_err(Error::from)
72}
73
74/// Deserialize a serde value using the Rings wire encoding.
75pub fn deserialize<T>(bytes: &[u8]) -> Result<T>
76where T: DeserializeOwned {
77    let (value, remaining) = postcard::take_from_bytes(bytes).map_err(Error::deserialize)?;
78    match remaining.is_empty() {
79        true => Ok(value),
80        false => Err(Error::TrailingBytes {
81            decoded: bytes.len() - remaining.len(),
82            total: bytes.len(),
83        }),
84    }
85}
86
87/// Deserialize one value from the start of `bytes` and return the unconsumed suffix.
88///
89/// Borrowing outputs such as `&[u8]` remain views into the input. This is used
90/// for allocation-free inspection of bounded-admission envelope prefixes.
91pub fn deserialize_prefix<'de, T>(bytes: &'de [u8]) -> Result<(T, &'de [u8])>
92where T: Deserialize<'de> {
93    postcard::take_from_bytes(bytes).map_err(Error::deserialize)
94}
95
96struct EnumVariant(u32);
97
98impl<'de> Deserialize<'de> for EnumVariant {
99    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
100    where D: serde::Deserializer<'de> {
101        deserializer.deserialize_enum("RingsEnum", &[], EnumVariantVisitor)
102    }
103}
104
105struct EnumVariantVisitor;
106
107impl<'de> Visitor<'de> for EnumVariantVisitor {
108    type Value = EnumVariant;
109
110    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
111        formatter.write_str("a Rings wire enum discriminant")
112    }
113
114    fn visit_enum<A>(self, data: A) -> std::result::Result<Self::Value, A::Error>
115    where A: EnumAccess<'de> {
116        let (variant, _) = data.variant::<u32>()?;
117        Ok(EnumVariant(variant))
118    }
119}
120
121/// Read an enum variant index without deserializing its body.
122///
123/// This supports bounded admission decisions that must inspect a Rings postcard
124/// envelope before allocating its potentially large variant payload.
125pub fn deserialize_enum_variant(bytes: &[u8]) -> Result<u32> {
126    let mut deserializer = postcard::Deserializer::from_bytes(bytes);
127    EnumVariant::deserialize(&mut deserializer)
128        .map(|variant| variant.0)
129        .map_err(Error::deserialize)
130}
131
132/// Return the serialized size of a serde value using the default wire encoding.
133pub fn serialized_size<T>(value: &T) -> Result<u64>
134where T: Serialize {
135    postcard::experimental::serialized_size(value)
136        .map(|bytes| bytes as u64)
137        .map_err(Error::from)
138}
139
140#[cfg(test)]
141mod tests {
142    use super::*;
143
144    #[derive(Debug, Deserialize, Eq, PartialEq, Serialize)]
145    struct Example {
146        id: u64,
147        label: String,
148        bytes: Vec<u8>,
149    }
150
151    #[derive(Debug, Deserialize, Eq, PartialEq, Serialize)]
152    struct WideIntegers {
153        signed: i128,
154        unsigned: u128,
155    }
156
157    #[derive(Debug, Deserialize, Eq, PartialEq, Serialize)]
158    enum TaggedBody {
159        Empty,
160        Bytes(Vec<u8>),
161    }
162
163    #[test]
164    fn test_roundtrip_and_size_match() -> std::result::Result<(), Box<dyn StdError>> {
165        let value = Example {
166            id: 42,
167            label: "rings".to_string(),
168            bytes: vec![1, 2, 3, 4],
169        };
170
171        let encoded = serialize(&value)?;
172        assert_eq!(
173            encoded,
174            vec![42, 5, b'r', b'i', b'n', b'g', b's', 4, 1, 2, 3, 4],
175            "wire encoding must stay stable for the Rings codec"
176        );
177        assert_eq!(u64::try_from(encoded.len())?, serialized_size(&value)?);
178        assert_eq!(deserialize::<Example>(&encoded)?, value);
179        Ok(())
180    }
181
182    #[test]
183    fn test_roundtrip_128_bit_integers() -> std::result::Result<(), Box<dyn StdError>> {
184        let value = WideIntegers {
185            signed: -123_456_789_012_345_678_901_234_567_890i128,
186            unsigned: 0x0102_0304_0506_0708_090a_0b0c_0d0e_0f10u128,
187        };
188
189        let encoded = serialize(&value)?;
190        assert_eq!(u64::try_from(encoded.len())?, serialized_size(&value)?);
191        assert_eq!(deserialize::<WideIntegers>(&encoded)?, value);
192        Ok(())
193    }
194
195    #[test]
196    fn test_enum_variant_decode_does_not_require_the_variant_body() -> Result<()> {
197        assert_eq!(
198            deserialize_enum_variant(&serialize(&TaggedBody::Empty)?)?,
199            0
200        );
201        let encoded = serialize(&TaggedBody::Bytes(vec![7; 1024]))?;
202
203        assert_eq!(deserialize_enum_variant(&encoded)?, 1);
204        assert!(encoded.starts_with(&[1]));
205        let variant_only = [1];
206        assert_eq!(deserialize_enum_variant(&variant_only)?, 1);
207        assert!(deserialize::<TaggedBody>(&variant_only).is_err());
208        Ok(())
209    }
210
211    #[test]
212    fn test_prefix_decode_borrows_byte_slices_without_consuming_suffix() -> Result<()> {
213        let mut encoded = serialize(&vec![1_u8, 2, 3])?;
214        encoded.push(99);
215
216        let (borrowed, remaining) = deserialize_prefix::<&[u8]>(&encoded)?;
217
218        assert_eq!(borrowed, &[1, 2, 3]);
219        assert_eq!(remaining, &[99]);
220        Ok(())
221    }
222
223    #[test]
224    fn test_deserialize_rejects_trailing_bytes() -> std::result::Result<(), Box<dyn StdError>> {
225        let mut encoded = serialize(&42u64)?;
226        encoded.extend_from_slice(&[1, 2, 3]);
227
228        match deserialize::<u64>(&encoded) {
229            Err(error) => {
230                assert!(matches!(error, Error::TrailingBytes {
231                    decoded: 1,
232                    total: 4
233                }));
234                Ok(())
235            }
236            Ok(_) => Err("trailing bytes must fail".into()),
237        }
238    }
239}