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}
46
47impl FrameKind {
48 pub const fn to_u8(self) -> u8 {
50 match self {
51 FrameKind::Hello => 0,
52 FrameKind::Data => 1,
53 FrameKind::Error => 2,
54 FrameKind::Subscribe => 3,
55 FrameKind::Unsubscribe => 4,
56 FrameKind::Credit => 5,
57 FrameKind::Cursor => 6,
58 }
59 }
60
61 pub const fn from_u8(code: u8) -> Option<FrameKind> {
65 match code {
66 0 => Some(FrameKind::Hello),
67 1 => Some(FrameKind::Data),
68 2 => Some(FrameKind::Error),
69 3 => Some(FrameKind::Subscribe),
70 4 => Some(FrameKind::Unsubscribe),
71 5 => Some(FrameKind::Credit),
72 6 => Some(FrameKind::Cursor),
73 _ => None,
74 }
75 }
76
77 pub const fn has_payload(self) -> bool {
79 matches!(self, FrameKind::Data)
80 }
81}
82
83impl fmt::Display for FrameKind {
84 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
85 let s = match self {
86 FrameKind::Hello => "HELLO",
87 FrameKind::Data => "DATA",
88 FrameKind::Error => "ERROR",
89 FrameKind::Subscribe => "SUBSCRIBE",
90 FrameKind::Unsubscribe => "UNSUBSCRIBE",
91 FrameKind::Credit => "CREDIT",
92 FrameKind::Cursor => "CURSOR",
93 };
94 f.write_str(s)
95 }
96}
97
98#[derive(Clone, Copy, Debug, PartialEq, Eq)]
100pub struct Preamble {
101 pub kind: FrameKind,
103 pub header_len: u64,
105}
106
107#[derive(Clone, Copy, Debug, PartialEq, Eq)]
109pub enum PreambleError {
110 Incomplete,
112 BadMagic(u8),
114 UnknownKind(u8),
116 HeaderTooLarge {
119 len: u64,
121 max: u64,
123 },
124}
125
126impl PreambleError {
127 pub const fn is_violation(self) -> bool {
130 !matches!(self, PreambleError::Incomplete)
131 }
132}
133
134impl fmt::Display for PreambleError {
135 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
136 match self {
137 PreambleError::Incomplete => f.write_str("incomplete preamble"),
138 PreambleError::BadMagic(b) => write!(f, "bad magic byte {b:#04x}, expected 0x57"),
139 PreambleError::UnknownKind(k) => write!(f, "unknown frame kind {k}"),
140 PreambleError::HeaderTooLarge { len, max } => {
141 write!(f, "header length {len} exceeds the limit of {max} bytes")
142 }
143 }
144 }
145}
146
147impl std::error::Error for PreambleError {}
148
149pub fn parse_preamble(
155 input: &[u8],
156 max_header_bytes: u64,
157) -> Result<(Preamble, usize), PreambleError> {
158 if input.len() < 2 {
159 return Err(PreambleError::Incomplete);
160 }
161 if input[0] != MAGIC {
162 return Err(PreambleError::BadMagic(input[0]));
163 }
164 let kind = FrameKind::from_u8(input[1]).ok_or(PreambleError::UnknownKind(input[1]))?;
165 let (header_len, used) = match decode_varint(&input[2..]) {
166 Ok(v) => v,
167 Err(VarintError::Truncated) => return Err(PreambleError::Incomplete),
168 Err(VarintError::OutOfRange) => unreachable!("decoding cannot overflow"),
171 };
172 if header_len > max_header_bytes {
173 return Err(PreambleError::HeaderTooLarge {
174 len: header_len,
175 max: max_header_bytes,
176 });
177 }
178 Ok((Preamble { kind, header_len }, 2 + used))
179}
180
181pub fn preamble_bytes(kind: FrameKind, header_len: u64) -> ([u8; MAX_PREAMBLE_LEN], usize) {
191 let mut bytes = [0u8; MAX_PREAMBLE_LEN];
192 bytes[0] = MAGIC;
193 bytes[1] = kind.to_u8();
194 let len = crate::varint::write_varint(header_len, &mut bytes[2..])
195 .expect("header lengths are bounded far below 2^62");
196 (bytes, 2 + len)
197}
198
199pub fn encode_preamble(kind: FrameKind, header_len: u64, out: &mut Vec<u8>) {
201 let (bytes, len) = preamble_bytes(kind, header_len);
202 out.extend_from_slice(&bytes[..len]);
203}
204
205pub fn encode_frame(kind: FrameKind, header: &[u8]) -> Vec<u8> {
207 let mut out = Vec::with_capacity(MAX_PREAMBLE_LEN + header.len());
208 encode_preamble(kind, header.len() as u64, &mut out);
209 out.extend_from_slice(header);
210 out
211}
212
213#[cfg(test)]
214mod tests {
215 use super::*;
216
217 const CAP: u64 = 16 * 1024;
218
219 #[test]
220 fn kind_codes_match_the_protocol_document() {
221 let all = [
222 (FrameKind::Hello, 0u8),
223 (FrameKind::Data, 1),
224 (FrameKind::Error, 2),
225 (FrameKind::Subscribe, 3),
226 (FrameKind::Unsubscribe, 4),
227 (FrameKind::Credit, 5),
228 (FrameKind::Cursor, 6),
229 ];
230 for (kind, code) in all {
231 assert_eq!(kind.to_u8(), code);
232 assert_eq!(FrameKind::from_u8(code), Some(kind));
233 }
234 for code in 7u8..=255 {
235 assert_eq!(FrameKind::from_u8(code), None, "kind {code}");
236 }
237 }
238
239 #[test]
240 fn only_data_carries_payload() {
241 assert!(FrameKind::Data.has_payload());
242 for k in [
243 FrameKind::Hello,
244 FrameKind::Error,
245 FrameKind::Subscribe,
246 FrameKind::Unsubscribe,
247 FrameKind::Credit,
248 FrameKind::Cursor,
249 ] {
250 assert!(!k.has_payload(), "{k}");
251 }
252 }
253
254 #[test]
255 fn roundtrip_through_encode_and_parse() {
256 for len in [0u64, 1, 63, 64, 16_383, 16_384] {
257 let mut buf = Vec::new();
258 encode_preamble(FrameKind::Data, len, &mut buf);
259 let (p, used) = parse_preamble(&buf, CAP).unwrap();
260 assert_eq!(p.kind, FrameKind::Data);
261 assert_eq!(p.header_len, len);
262 assert_eq!(used, buf.len());
263 }
264 }
265
266 #[test]
267 fn every_kind_encodes_its_documented_kind_byte() {
268 for (kind, code) in [
273 (FrameKind::Hello, 0x00u8),
274 (FrameKind::Data, 0x01),
275 (FrameKind::Error, 0x02),
276 (FrameKind::Subscribe, 0x03),
277 (FrameKind::Unsubscribe, 0x04),
278 ] {
279 for header_len in [0usize, 3, 5, 9, 11, 16, 18] {
280 let frame = encode_frame(kind, &vec![0; header_len]);
281 assert_eq!(frame[0], MAGIC, "{kind}: magic");
282 assert_eq!(frame[1], code, "{kind}: kind byte");
283 let (preamble, used) = parse_preamble(&frame, CAP).unwrap();
284 assert_eq!(preamble.kind, kind);
285 assert_eq!(preamble.header_len as usize, header_len);
286 assert_eq!(frame.len() - used, header_len, "{kind}: header follows");
287 }
288 }
289 }
290
291 #[test]
292 fn incomplete_input_is_not_a_violation() {
293 for prefix in [
294 &[][..],
295 &[MAGIC][..],
296 &[MAGIC, 1, 0x80][..],
297 &[MAGIC, 1, 0xc0, 0, 0][..],
298 ] {
299 let err = parse_preamble(prefix, CAP).unwrap_err();
300 assert_eq!(err, PreambleError::Incomplete, "{prefix:?}");
301 assert!(!err.is_violation());
302 }
303 }
304
305 #[test]
306 fn bad_magic_is_a_violation() {
307 let err = parse_preamble(&[0x58, 0x01, 0x00], CAP).unwrap_err();
308 assert_eq!(err, PreambleError::BadMagic(0x58));
309 assert!(err.is_violation());
310 }
311
312 #[test]
313 fn unknown_kind_is_a_violation() {
314 let err = parse_preamble(&[MAGIC, 0x07, 0x00], CAP).unwrap_err();
319 assert_eq!(err, PreambleError::UnknownKind(7));
320 assert!(err.is_violation());
321 }
322
323 #[test]
324 fn oversized_header_is_rejected_before_allocation() {
325 let mut buf = vec![MAGIC, FrameKind::Data.to_u8()];
327 crate::varint::encode_varint(1024 * 1024, &mut buf).unwrap();
328 let err = parse_preamble(&buf, CAP).unwrap_err();
329 assert_eq!(
330 err,
331 PreambleError::HeaderTooLarge {
332 len: 1024 * 1024,
333 max: CAP
334 }
335 );
336 assert!(err.is_violation());
337 }
338
339 #[test]
340 fn a_header_exactly_at_the_cap_is_accepted() {
341 let mut buf = vec![MAGIC, FrameKind::Hello.to_u8()];
342 crate::varint::encode_varint(CAP, &mut buf).unwrap();
343 assert_eq!(parse_preamble(&buf, CAP).unwrap().0.header_len, CAP);
344 }
345
346 #[test]
347 fn non_minimal_length_encodings_are_accepted() {
348 let buf = [MAGIC, FrameKind::Error.to_u8(), 0xc0, 0, 0, 0, 0, 0, 0, 5];
350 let (p, used) = parse_preamble(&buf, CAP).unwrap();
351 assert_eq!(p.header_len, 5);
352 assert_eq!(used, 10);
353 }
354
355 #[test]
356 fn payload_bytes_after_the_preamble_are_untouched() {
357 let frame = encode_frame(FrameKind::Data, &[0xaa, 0xbb]);
358 let (p, used) = parse_preamble(&frame, CAP).unwrap();
359 assert_eq!(p.header_len, 2);
360 assert_eq!(&frame[used..], &[0xaa, 0xbb]);
361 }
362}