1use std::fmt;
8
9use crate::varint::{self, VarintError, decode_varint};
10
11pub const MAGIC: u8 = 0x57;
13
14pub const MAX_PREAMBLE_LEN: usize = 2 + varint::MAX_ENCODED_LEN;
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
19pub enum FrameKind {
20 Hello,
22 Data,
24 Error,
26 Subscribe,
28 Unsubscribe,
30 Credit,
38 Cursor,
45 Flow,
54}
55
56impl FrameKind {
57 pub const fn to_u8(self) -> u8 {
59 match self {
60 FrameKind::Hello => 0,
61 FrameKind::Data => 1,
62 FrameKind::Error => 2,
63 FrameKind::Subscribe => 3,
64 FrameKind::Unsubscribe => 4,
65 FrameKind::Credit => 5,
66 FrameKind::Cursor => 6,
67 FrameKind::Flow => 7,
68 }
69 }
70
71 pub const fn from_u8(code: u8) -> Option<FrameKind> {
75 match code {
76 0 => Some(FrameKind::Hello),
77 1 => Some(FrameKind::Data),
78 2 => Some(FrameKind::Error),
79 3 => Some(FrameKind::Subscribe),
80 4 => Some(FrameKind::Unsubscribe),
81 5 => Some(FrameKind::Credit),
82 6 => Some(FrameKind::Cursor),
83 7 => Some(FrameKind::Flow),
84 _ => None,
85 }
86 }
87
88 pub const fn has_payload(self) -> bool {
90 matches!(self, FrameKind::Data)
91 }
92}
93
94impl fmt::Display for FrameKind {
95 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
96 let s = match self {
97 FrameKind::Hello => "HELLO",
98 FrameKind::Data => "DATA",
99 FrameKind::Error => "ERROR",
100 FrameKind::Subscribe => "SUBSCRIBE",
101 FrameKind::Unsubscribe => "UNSUBSCRIBE",
102 FrameKind::Credit => "CREDIT",
103 FrameKind::Cursor => "CURSOR",
104 FrameKind::Flow => "FLOW",
105 };
106 f.write_str(s)
107 }
108}
109
110#[derive(Clone, Copy, Debug, PartialEq, Eq)]
112pub struct Preamble {
113 pub kind: FrameKind,
115 pub header_len: u64,
117}
118
119#[derive(Clone, Copy, Debug, PartialEq, Eq)]
121pub enum PreambleError {
122 Incomplete,
124 BadMagic(u8),
126 UnknownKind(u8),
128 HeaderTooLarge {
131 len: u64,
133 max: u64,
135 },
136}
137
138impl PreambleError {
139 pub const fn is_violation(self) -> bool {
142 !matches!(self, PreambleError::Incomplete)
143 }
144}
145
146impl fmt::Display for PreambleError {
147 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
148 match self {
149 PreambleError::Incomplete => f.write_str("incomplete preamble"),
150 PreambleError::BadMagic(b) => write!(f, "bad magic byte {b:#04x}, expected 0x57"),
151 PreambleError::UnknownKind(k) => write!(f, "unknown frame kind {k}"),
152 PreambleError::HeaderTooLarge { len, max } => {
153 write!(f, "header length {len} exceeds the limit of {max} bytes")
154 }
155 }
156 }
157}
158
159impl std::error::Error for PreambleError {}
160
161pub fn parse_preamble(
167 input: &[u8],
168 max_header_bytes: u64,
169) -> Result<(Preamble, usize), PreambleError> {
170 if input.len() < 2 {
171 return Err(PreambleError::Incomplete);
172 }
173 if input[0] != MAGIC {
174 return Err(PreambleError::BadMagic(input[0]));
175 }
176 let kind = FrameKind::from_u8(input[1]).ok_or(PreambleError::UnknownKind(input[1]))?;
177 let (header_len, used) = match decode_varint(&input[2..]) {
178 Ok(v) => v,
179 Err(VarintError::Truncated) => return Err(PreambleError::Incomplete),
180 Err(VarintError::OutOfRange) => unreachable!("decoding cannot overflow"),
183 };
184 if header_len > max_header_bytes {
185 return Err(PreambleError::HeaderTooLarge {
186 len: header_len,
187 max: max_header_bytes,
188 });
189 }
190 Ok((Preamble { kind, header_len }, 2 + used))
191}
192
193pub fn preamble_bytes(kind: FrameKind, header_len: u64) -> ([u8; MAX_PREAMBLE_LEN], usize) {
203 let mut bytes = [0u8; MAX_PREAMBLE_LEN];
204 bytes[0] = MAGIC;
205 bytes[1] = kind.to_u8();
206 let len = crate::varint::write_varint(header_len, &mut bytes[2..])
207 .expect("header lengths are bounded far below 2^62");
208 (bytes, 2 + len)
209}
210
211pub fn encode_preamble(kind: FrameKind, header_len: u64, out: &mut Vec<u8>) {
213 let (bytes, len) = preamble_bytes(kind, header_len);
214 out.extend_from_slice(&bytes[..len]);
215}
216
217pub fn encode_frame(kind: FrameKind, header: &[u8]) -> Vec<u8> {
219 let mut out = Vec::with_capacity(MAX_PREAMBLE_LEN + header.len());
220 encode_preamble(kind, header.len() as u64, &mut out);
221 out.extend_from_slice(header);
222 out
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228
229 const CAP: u64 = 16 * 1024;
230
231 #[test]
232 fn kind_codes_match_the_protocol_document() {
233 let all = [
234 (FrameKind::Hello, 0u8),
235 (FrameKind::Data, 1),
236 (FrameKind::Error, 2),
237 (FrameKind::Subscribe, 3),
238 (FrameKind::Unsubscribe, 4),
239 (FrameKind::Credit, 5),
240 (FrameKind::Cursor, 6),
241 (FrameKind::Flow, 7),
242 ];
243 for (kind, code) in all {
244 assert_eq!(kind.to_u8(), code);
245 assert_eq!(FrameKind::from_u8(code), Some(kind));
246 }
247 for code in 8u8..=255 {
248 assert_eq!(FrameKind::from_u8(code), None, "kind {code}");
249 }
250 }
251
252 #[test]
253 fn only_data_carries_payload() {
254 assert!(FrameKind::Data.has_payload());
255 for k in [
256 FrameKind::Hello,
257 FrameKind::Error,
258 FrameKind::Subscribe,
259 FrameKind::Unsubscribe,
260 FrameKind::Credit,
261 FrameKind::Cursor,
262 FrameKind::Flow,
263 ] {
264 assert!(!k.has_payload(), "{k}");
265 }
266 }
267
268 #[test]
269 fn roundtrip_through_encode_and_parse() {
270 for len in [0u64, 1, 63, 64, 16_383, 16_384] {
271 let mut buf = Vec::new();
272 encode_preamble(FrameKind::Data, len, &mut buf);
273 let (p, used) = parse_preamble(&buf, CAP).unwrap();
274 assert_eq!(p.kind, FrameKind::Data);
275 assert_eq!(p.header_len, len);
276 assert_eq!(used, buf.len());
277 }
278 }
279
280 #[test]
281 fn every_kind_encodes_its_documented_kind_byte() {
282 for (kind, code) in [
287 (FrameKind::Hello, 0x00u8),
288 (FrameKind::Data, 0x01),
289 (FrameKind::Error, 0x02),
290 (FrameKind::Subscribe, 0x03),
291 (FrameKind::Unsubscribe, 0x04),
292 (FrameKind::Credit, 0x05),
293 (FrameKind::Cursor, 0x06),
294 (FrameKind::Flow, 0x07),
295 ] {
296 for header_len in [0usize, 3, 5, 9, 11, 16, 18] {
297 let frame = encode_frame(kind, &vec![0; header_len]);
298 assert_eq!(frame[0], MAGIC, "{kind}: magic");
299 assert_eq!(frame[1], code, "{kind}: kind byte");
300 let (preamble, used) = parse_preamble(&frame, CAP).unwrap();
301 assert_eq!(preamble.kind, kind);
302 assert_eq!(preamble.header_len as usize, header_len);
303 assert_eq!(frame.len() - used, header_len, "{kind}: header follows");
304 }
305 }
306 }
307
308 #[test]
309 fn incomplete_input_is_not_a_violation() {
310 for prefix in [
311 &[][..],
312 &[MAGIC][..],
313 &[MAGIC, 1, 0x80][..],
314 &[MAGIC, 1, 0xc0, 0, 0][..],
315 ] {
316 let err = parse_preamble(prefix, CAP).unwrap_err();
317 assert_eq!(err, PreambleError::Incomplete, "{prefix:?}");
318 assert!(!err.is_violation());
319 }
320 }
321
322 #[test]
323 fn bad_magic_is_a_violation() {
324 let err = parse_preamble(&[0x58, 0x01, 0x00], CAP).unwrap_err();
325 assert_eq!(err, PreambleError::BadMagic(0x58));
326 assert!(err.is_violation());
327 }
328
329 #[test]
330 fn unknown_kind_is_a_violation() {
331 let err = parse_preamble(&[MAGIC, 0x08, 0x00], CAP).unwrap_err();
337 assert_eq!(err, PreambleError::UnknownKind(8));
338 assert!(err.is_violation());
339 }
340
341 #[test]
342 fn oversized_header_is_rejected_before_allocation() {
343 let mut buf = vec![MAGIC, FrameKind::Data.to_u8()];
345 crate::varint::encode_varint(1024 * 1024, &mut buf).unwrap();
346 let err = parse_preamble(&buf, CAP).unwrap_err();
347 assert_eq!(
348 err,
349 PreambleError::HeaderTooLarge {
350 len: 1024 * 1024,
351 max: CAP
352 }
353 );
354 assert!(err.is_violation());
355 }
356
357 #[test]
358 fn a_header_exactly_at_the_cap_is_accepted() {
359 let mut buf = vec![MAGIC, FrameKind::Hello.to_u8()];
360 crate::varint::encode_varint(CAP, &mut buf).unwrap();
361 assert_eq!(parse_preamble(&buf, CAP).unwrap().0.header_len, CAP);
362 }
363
364 #[test]
365 fn non_minimal_length_encodings_are_accepted() {
366 let buf = [MAGIC, FrameKind::Error.to_u8(), 0xc0, 0, 0, 0, 0, 0, 0, 5];
368 let (p, used) = parse_preamble(&buf, CAP).unwrap();
369 assert_eq!(p.header_len, 5);
370 assert_eq!(used, 10);
371 }
372
373 #[test]
374 fn payload_bytes_after_the_preamble_are_untouched() {
375 let frame = encode_frame(FrameKind::Data, &[0xaa, 0xbb]);
376 let (p, used) = parse_preamble(&frame, CAP).unwrap();
377 assert_eq!(p.header_len, 2);
378 assert_eq!(&frame[used..], &[0xaa, 0xbb]);
379 }
380}