1use crate::error::DecodeError;
11use crate::limits::MAX_FRAME_BYTES;
12use serde::Serialize;
13use serde::de::DeserializeOwned;
14
15pub fn encode_named<T: Serialize + ?Sized>(value: &T) -> Result<Vec<u8>, DecodeError> {
17 let mut buffer = Vec::new();
18 encode_named_into(value, &mut buffer)?;
19 Ok(buffer)
20}
21
22pub fn encode_named_into<T: Serialize + ?Sized>(
27 value: &T,
28 buffer: &mut Vec<u8>,
29) -> Result<(), DecodeError> {
30 buffer.clear();
31 ciborium::into_writer(value, &mut *buffer)
32 .map_err(|error| DecodeError::Encode(error.to_string()))
33}
34
35pub fn decode_named<T: DeserializeOwned>(payload: &[u8]) -> Result<T, DecodeError> {
40 let mut reader = payload;
41 let value = ciborium::from_reader(&mut reader)
42 .map_err(|error| DecodeError::Decode(error.to_string()))?;
43 if !reader.is_empty() {
44 return Err(DecodeError::Decode(format!(
45 "{} trailing byte(s) after the CBOR value",
46 reader.len()
47 )));
48 }
49 Ok(value)
50}
51
52pub fn frame_encode(payload: &[u8]) -> Result<Vec<u8>, DecodeError> {
55 if payload.len() > MAX_FRAME_BYTES {
56 return Err(DecodeError::Frame(format!(
57 "payload is {}B, exceeds frame cap {MAX_FRAME_BYTES}B",
58 payload.len()
59 )));
60 }
61 let mut framed = Vec::with_capacity(4 + payload.len());
62 framed.extend_from_slice(&(payload.len() as u32).to_le_bytes());
63 framed.extend_from_slice(payload);
64 Ok(framed)
65}
66
67pub fn frame_decode(buf: &[u8]) -> Result<Option<(&[u8], usize)>, DecodeError> {
72 let Some(prefix) = buf.get(..4) else {
73 return Ok(None);
74 };
75 let len = u32::from_le_bytes(prefix.try_into().expect("4-byte slice")) as usize;
76 if len > MAX_FRAME_BYTES {
77 return Err(DecodeError::Frame(format!(
78 "frame length {len}B exceeds cap {MAX_FRAME_BYTES}B"
79 )));
80 }
81 match buf.get(4..4 + len) {
82 Some(payload) => Ok(Some((payload, 4 + len))),
83 None => Ok(None),
84 }
85}
86
87#[cfg(test)]
88mod tests {
89 use super::*;
90
91 #[test]
92 fn given_a_payload_when_framed_then_should_round_trip() {
93 let framed = frame_encode(b"hello").expect("frames");
94 assert_eq!(&framed[..4], &5u32.to_le_bytes());
95 let (payload, consumed) = frame_decode(&framed)
96 .expect("decodes")
97 .expect("whole frame present");
98 assert_eq!(payload, b"hello");
99 assert_eq!(consumed, framed.len());
100 }
101
102 #[test]
103 fn given_a_partial_frame_when_decoded_then_should_ask_for_more() {
104 let framed = frame_encode(b"hello").expect("frames");
105 assert!(frame_decode(&framed[..3]).expect("no error").is_none());
106 assert!(frame_decode(&framed[..6]).expect("no error").is_none());
107 }
108
109 #[test]
110 fn given_an_oversized_length_prefix_when_decoded_then_should_error() {
111 let mut bad = Vec::new();
112 bad.extend_from_slice(&((MAX_FRAME_BYTES as u32) + 1).to_le_bytes());
113 assert!(matches!(frame_decode(&bad), Err(DecodeError::Frame(_))));
114 }
115
116 #[test]
117 fn given_named_encoding_when_round_tripped_then_should_preserve_fields() {
118 #[derive(serde::Serialize, serde::Deserialize, PartialEq, Debug)]
119 struct Body {
120 id: u32,
121 name: String,
122 }
123 let body = Body {
124 id: 7,
125 name: "alice".to_owned(),
126 };
127 let bytes = encode_named(&body).expect("encodes");
128 let back: Body = decode_named(&bytes).expect("decodes");
129 assert_eq!(back, body);
130 }
131
132 #[test]
133 fn given_trailing_bytes_when_decoded_then_should_error() {
134 let mut bytes = encode_named(&7u32).expect("encodes");
135 bytes.push(0x00);
136 assert!(matches!(
137 decode_named::<u32>(&bytes),
138 Err(DecodeError::Decode(_))
139 ));
140 }
141 #[test]
142 fn given_the_in_place_encoder_when_reused_then_should_match_encode_named_bytes() {
143 let value = serde_json::json!({"a": 1, "b": "two"});
144 let mut buffer = vec![0xff; 64]; super::encode_named_into(&value, &mut buffer).expect("encodes");
146 assert_eq!(buffer, super::encode_named(&value).expect("encodes"));
147 }
148}