1#![deny(missing_docs)]
2
3use 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#[derive(Debug)]
16pub enum Error {
17 Serialize(postcard::Error),
19 Deserialize(postcard::Error),
21 TrailingBytes {
23 decoded: usize,
25 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
65pub type Result<T> = std::result::Result<T, Error>;
67
68pub fn serialize<T>(value: &T) -> Result<Vec<u8>>
70where T: Serialize {
71 postcard::to_allocvec(value).map_err(Error::from)
72}
73
74pub 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
87pub 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
121pub 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
132pub 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}