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 Report,
60}
61
62impl FrameKind {
63 pub const fn to_u8(self) -> u8 {
65 match self {
66 FrameKind::Hello => 0,
67 FrameKind::Data => 1,
68 FrameKind::Error => 2,
69 FrameKind::Subscribe => 3,
70 FrameKind::Unsubscribe => 4,
71 FrameKind::Credit => 5,
72 FrameKind::Cursor => 6,
73 FrameKind::Flow => 7,
74 FrameKind::Report => 8,
75 }
76 }
77
78 pub const fn from_u8(code: u8) -> Option<FrameKind> {
82 match code {
83 0 => Some(FrameKind::Hello),
84 1 => Some(FrameKind::Data),
85 2 => Some(FrameKind::Error),
86 3 => Some(FrameKind::Subscribe),
87 4 => Some(FrameKind::Unsubscribe),
88 5 => Some(FrameKind::Credit),
89 6 => Some(FrameKind::Cursor),
90 7 => Some(FrameKind::Flow),
91 8 => Some(FrameKind::Report),
92 _ => None,
93 }
94 }
95
96 pub const fn has_payload(self) -> bool {
98 matches!(self, FrameKind::Data)
99 }
100}
101
102impl fmt::Display for FrameKind {
103 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
104 let s = match self {
105 FrameKind::Hello => "HELLO",
106 FrameKind::Data => "DATA",
107 FrameKind::Error => "ERROR",
108 FrameKind::Subscribe => "SUBSCRIBE",
109 FrameKind::Unsubscribe => "UNSUBSCRIBE",
110 FrameKind::Credit => "CREDIT",
111 FrameKind::Cursor => "CURSOR",
112 FrameKind::Flow => "FLOW",
113 FrameKind::Report => "REPORT",
114 };
115 f.write_str(s)
116 }
117}
118
119#[derive(Clone, Copy, Debug, PartialEq, Eq)]
121pub struct Preamble {
122 pub kind: FrameKind,
124 pub header_len: u64,
126}
127
128#[derive(Clone, Copy, Debug, PartialEq, Eq)]
130pub enum PreambleError {
131 Incomplete,
133 BadMagic(u8),
135 UnknownKind(u8),
137 HeaderTooLarge {
140 len: u64,
142 max: u64,
144 },
145}
146
147impl PreambleError {
148 pub const fn is_violation(self) -> bool {
151 !matches!(self, PreambleError::Incomplete)
152 }
153}
154
155impl fmt::Display for PreambleError {
156 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
157 match self {
158 PreambleError::Incomplete => f.write_str("incomplete preamble"),
159 PreambleError::BadMagic(b) => write!(f, "bad magic byte {b:#04x}, expected 0x57"),
160 PreambleError::UnknownKind(k) => write!(f, "unknown frame kind {k}"),
161 PreambleError::HeaderTooLarge { len, max } => {
162 write!(f, "header length {len} exceeds the limit of {max} bytes")
163 }
164 }
165 }
166}
167
168impl std::error::Error for PreambleError {}
169
170pub fn parse_preamble(
176 input: &[u8],
177 max_header_bytes: u64,
178) -> Result<(Preamble, usize), PreambleError> {
179 if input.len() < 2 {
180 return Err(PreambleError::Incomplete);
181 }
182 if input[0] != MAGIC {
183 return Err(PreambleError::BadMagic(input[0]));
184 }
185 let kind = FrameKind::from_u8(input[1]).ok_or(PreambleError::UnknownKind(input[1]))?;
186 let (header_len, used) = match decode_varint(&input[2..]) {
187 Ok(v) => v,
188 Err(VarintError::Truncated) => return Err(PreambleError::Incomplete),
189 Err(VarintError::OutOfRange) => unreachable!("decoding cannot overflow"),
192 };
193 if header_len > max_header_bytes {
194 return Err(PreambleError::HeaderTooLarge {
195 len: header_len,
196 max: max_header_bytes,
197 });
198 }
199 Ok((Preamble { kind, header_len }, 2 + used))
200}
201
202pub fn preamble_bytes(kind: FrameKind, header_len: u64) -> ([u8; MAX_PREAMBLE_LEN], usize) {
212 let mut bytes = [0u8; MAX_PREAMBLE_LEN];
213 bytes[0] = MAGIC;
214 bytes[1] = kind.to_u8();
215 let len = crate::varint::write_varint(header_len, &mut bytes[2..])
216 .expect("header lengths are bounded far below 2^62");
217 (bytes, 2 + len)
218}
219
220pub fn encode_preamble(kind: FrameKind, header_len: u64, out: &mut Vec<u8>) {
222 let (bytes, len) = preamble_bytes(kind, header_len);
223 out.extend_from_slice(&bytes[..len]);
224}
225
226pub fn encode_frame(kind: FrameKind, header: &[u8]) -> Vec<u8> {
228 let mut out = Vec::with_capacity(MAX_PREAMBLE_LEN + header.len());
229 encode_preamble(kind, header.len() as u64, &mut out);
230 out.extend_from_slice(header);
231 out
232}
233
234#[cfg(test)]
235mod tests {
236 use super::*;
237
238 const CAP: u64 = 16 * 1024;
239
240 #[test]
241 fn kind_codes_match_the_protocol_document() {
242 let all = [
243 (FrameKind::Hello, 0u8),
244 (FrameKind::Data, 1),
245 (FrameKind::Error, 2),
246 (FrameKind::Subscribe, 3),
247 (FrameKind::Unsubscribe, 4),
248 (FrameKind::Credit, 5),
249 (FrameKind::Cursor, 6),
250 (FrameKind::Flow, 7),
251 (FrameKind::Report, 8),
252 ];
253 for (kind, code) in all {
254 assert_eq!(kind.to_u8(), code);
255 assert_eq!(FrameKind::from_u8(code), Some(kind));
256 }
257 for code in 9u8..=255 {
258 assert_eq!(FrameKind::from_u8(code), None, "kind {code}");
259 }
260 }
261
262 #[test]
263 fn only_data_carries_payload() {
264 assert!(FrameKind::Data.has_payload());
265 for k in [
266 FrameKind::Hello,
267 FrameKind::Error,
268 FrameKind::Subscribe,
269 FrameKind::Unsubscribe,
270 FrameKind::Credit,
271 FrameKind::Cursor,
272 FrameKind::Flow,
273 FrameKind::Report,
274 ] {
275 assert!(!k.has_payload(), "{k}");
276 }
277 }
278
279 #[test]
280 fn roundtrip_through_encode_and_parse() {
281 for len in [0u64, 1, 63, 64, 16_383, 16_384] {
282 let mut buf = Vec::new();
283 encode_preamble(FrameKind::Data, len, &mut buf);
284 let (p, used) = parse_preamble(&buf, CAP).unwrap();
285 assert_eq!(p.kind, FrameKind::Data);
286 assert_eq!(p.header_len, len);
287 assert_eq!(used, buf.len());
288 }
289 }
290
291 #[test]
292 fn every_kind_encodes_its_documented_kind_byte() {
293 for (kind, code) in [
298 (FrameKind::Hello, 0x00u8),
299 (FrameKind::Data, 0x01),
300 (FrameKind::Error, 0x02),
301 (FrameKind::Subscribe, 0x03),
302 (FrameKind::Unsubscribe, 0x04),
303 (FrameKind::Credit, 0x05),
304 (FrameKind::Cursor, 0x06),
305 (FrameKind::Flow, 0x07),
306 (FrameKind::Report, 0x08),
307 ] {
308 for header_len in [0usize, 3, 5, 9, 11, 16, 18] {
309 let frame = encode_frame(kind, &vec![0; header_len]);
310 assert_eq!(frame[0], MAGIC, "{kind}: magic");
311 assert_eq!(frame[1], code, "{kind}: kind byte");
312 let (preamble, used) = parse_preamble(&frame, CAP).unwrap();
313 assert_eq!(preamble.kind, kind);
314 assert_eq!(preamble.header_len as usize, header_len);
315 assert_eq!(frame.len() - used, header_len, "{kind}: header follows");
316 }
317 }
318 }
319
320 #[test]
321 fn incomplete_input_is_not_a_violation() {
322 for prefix in [
323 &[][..],
324 &[MAGIC][..],
325 &[MAGIC, 1, 0x80][..],
326 &[MAGIC, 1, 0xc0, 0, 0][..],
327 ] {
328 let err = parse_preamble(prefix, CAP).unwrap_err();
329 assert_eq!(err, PreambleError::Incomplete, "{prefix:?}");
330 assert!(!err.is_violation());
331 }
332 }
333
334 #[test]
335 fn bad_magic_is_a_violation() {
336 let err = parse_preamble(&[0x58, 0x01, 0x00], CAP).unwrap_err();
337 assert_eq!(err, PreambleError::BadMagic(0x58));
338 assert!(err.is_violation());
339 }
340
341 #[test]
342 fn unknown_kind_is_a_violation() {
343 let err = parse_preamble(&[MAGIC, 0x09, 0x00], CAP).unwrap_err();
351 assert_eq!(err, PreambleError::UnknownKind(9));
352 assert!(err.is_violation());
353 }
354
355 #[test]
356 fn oversized_header_is_rejected_before_allocation() {
357 let mut buf = vec![MAGIC, FrameKind::Data.to_u8()];
359 crate::varint::encode_varint(1024 * 1024, &mut buf).unwrap();
360 let err = parse_preamble(&buf, CAP).unwrap_err();
361 assert_eq!(
362 err,
363 PreambleError::HeaderTooLarge {
364 len: 1024 * 1024,
365 max: CAP
366 }
367 );
368 assert!(err.is_violation());
369 }
370
371 #[test]
372 fn a_header_exactly_at_the_cap_is_accepted() {
373 let mut buf = vec![MAGIC, FrameKind::Hello.to_u8()];
374 crate::varint::encode_varint(CAP, &mut buf).unwrap();
375 assert_eq!(parse_preamble(&buf, CAP).unwrap().0.header_len, CAP);
376 }
377
378 #[test]
379 fn non_minimal_length_encodings_are_accepted() {
380 let buf = [MAGIC, FrameKind::Error.to_u8(), 0xc0, 0, 0, 0, 0, 0, 0, 5];
382 let (p, used) = parse_preamble(&buf, CAP).unwrap();
383 assert_eq!(p.header_len, 5);
384 assert_eq!(used, 10);
385 }
386
387 #[test]
388 fn payload_bytes_after_the_preamble_are_untouched() {
389 let frame = encode_frame(FrameKind::Data, &[0xaa, 0xbb]);
390 let (p, used) = parse_preamble(&frame, CAP).unwrap();
391 assert_eq!(p.header_len, 2);
392 assert_eq!(&frame[used..], &[0xaa, 0xbb]);
393 }
394}