use bytes::{Buf, BufMut};
pub const SUBPROTOCOL_STREAM_WINDOW: u16 = 0x0B00;
pub const SUBPROTOCOL_STREAM_NACK: u16 = 0x0B01;
pub const SUBPROTOCOL_STREAM_RESET: u16 = 0x0B02;
pub const SUBPROTOCOL_STREAM_ACK: u16 = 0x0B03;
pub const STREAM_RESET_SIZE: usize = 8;
pub const STREAM_WINDOW_SIZE: usize = 24;
pub const STREAM_NACK_SIZE: usize = 24;
pub const STREAM_ACK_HEADER_SIZE: usize = 16;
pub const STREAM_ACK_RANGE_SIZE: usize = 16;
pub const MAX_ACK_RANGES: usize = 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct StreamWindow {
pub stream_id: u64,
pub total_consumed: u64,
pub ack_seq: u64,
}
#[derive(Debug, thiserror::Error)]
pub enum StreamWindowCodecError {
#[error("truncated stream subprotocol message: {got} bytes (need {need})")]
Truncated {
got: usize,
need: usize,
},
#[error("oversize stream subprotocol message: {got} bytes (need {need})")]
Oversize {
got: usize,
need: usize,
},
#[error("invalid stream-ack ranges: {reason}")]
InvalidRanges {
reason: &'static str,
},
}
#[inline]
fn require_exact_len(data: &[u8], need: usize) -> Result<(), StreamWindowCodecError> {
match data.len() {
n if n < need => Err(StreamWindowCodecError::Truncated { got: n, need }),
n if n > need => Err(StreamWindowCodecError::Oversize { got: n, need }),
_ => Ok(()),
}
}
impl StreamWindow {
#[inline]
pub fn encode(&self) -> [u8; STREAM_WINDOW_SIZE] {
let mut buf = [0u8; STREAM_WINDOW_SIZE];
(&mut buf[..8]).put_u64_le(self.stream_id);
(&mut buf[8..16]).put_u64_le(self.total_consumed);
(&mut buf[16..]).put_u64_le(self.ack_seq);
buf
}
pub fn decode(data: &[u8]) -> Result<Self, StreamWindowCodecError> {
require_exact_len(data, STREAM_WINDOW_SIZE)?;
let mut cur = std::io::Cursor::new(data);
let stream_id = cur.get_u64_le();
let total_consumed = cur.get_u64_le();
let ack_seq = cur.get_u64_le();
Ok(Self {
stream_id,
total_consumed,
ack_seq,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct StreamNack {
pub stream_id: u64,
pub next_expected: u64,
pub missing_bitmap: u64,
}
impl StreamNack {
#[inline]
pub fn encode(&self) -> [u8; STREAM_NACK_SIZE] {
let mut buf = [0u8; STREAM_NACK_SIZE];
(&mut buf[..8]).put_u64_le(self.stream_id);
(&mut buf[8..16]).put_u64_le(self.next_expected);
(&mut buf[16..]).put_u64_le(self.missing_bitmap);
buf
}
pub fn decode(data: &[u8]) -> Result<Self, StreamWindowCodecError> {
require_exact_len(data, STREAM_NACK_SIZE)?;
let mut cur = std::io::Cursor::new(data);
let stream_id = cur.get_u64_le();
let next_expected = cur.get_u64_le();
let missing_bitmap = cur.get_u64_le();
Ok(Self {
stream_id,
next_expected,
missing_bitmap,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StreamAckRanges {
pub stream_id: u64,
pub ack_seq: u64,
pub ranges: Vec<(u64, u64)>,
}
impl StreamAckRanges {
pub const MIN_SIZE: usize = STREAM_ACK_HEADER_SIZE + STREAM_ACK_RANGE_SIZE;
pub const MAX_SIZE: usize = STREAM_ACK_HEADER_SIZE + MAX_ACK_RANGES * STREAM_ACK_RANGE_SIZE;
pub fn encode(&self) -> Vec<u8> {
let mut buf =
Vec::with_capacity(STREAM_ACK_HEADER_SIZE + self.ranges.len() * STREAM_ACK_RANGE_SIZE);
buf.put_u64_le(self.stream_id);
buf.put_u64_le(self.ack_seq);
for &(start, end) in &self.ranges {
buf.put_u64_le(start);
buf.put_u64_le(end);
}
buf
}
pub fn decode(data: &[u8]) -> Result<Self, StreamWindowCodecError> {
if data.len() < Self::MIN_SIZE {
return Err(StreamWindowCodecError::Truncated {
got: data.len(),
need: Self::MIN_SIZE,
});
}
let body = data.len() - STREAM_ACK_HEADER_SIZE;
if !body.is_multiple_of(STREAM_ACK_RANGE_SIZE) {
return Err(StreamWindowCodecError::InvalidRanges {
reason: "length is not header + 16*n",
});
}
let n = body / STREAM_ACK_RANGE_SIZE;
if n > MAX_ACK_RANGES {
return Err(StreamWindowCodecError::InvalidRanges {
reason: "range count exceeds MAX_ACK_RANGES",
});
}
let mut cur = std::io::Cursor::new(data);
let stream_id = cur.get_u64_le();
let ack_seq = cur.get_u64_le();
let mut ranges = Vec::with_capacity(n);
let mut prev_start: Option<u64> = None;
for _ in 0..n {
let start = cur.get_u64_le();
let end = cur.get_u64_le();
if start >= end {
return Err(StreamWindowCodecError::InvalidRanges {
reason: "empty or inverted range",
});
}
if start <= ack_seq {
return Err(StreamWindowCodecError::InvalidRanges {
reason: "range not strictly above ack_seq",
});
}
if let Some(ps) = prev_start {
if end >= ps {
return Err(StreamWindowCodecError::InvalidRanges {
reason: "ranges not descending/merged",
});
}
}
prev_start = Some(start);
ranges.push((start, end));
}
Ok(Self {
stream_id,
ack_seq,
ranges,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct StreamReset {
pub stream_id: u64,
}
impl StreamReset {
#[inline]
pub fn encode(&self) -> [u8; STREAM_RESET_SIZE] {
self.stream_id.to_le_bytes()
}
pub fn decode(data: &[u8]) -> Result<Self, StreamWindowCodecError> {
require_exact_len(data, STREAM_RESET_SIZE)?;
let mut cur = std::io::Cursor::new(data);
Ok(Self {
stream_id: cur.get_u64_le(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_round_trip() {
let msg = StreamWindow {
stream_id: 0xDEAD_BEEF_CAFE_F00D,
total_consumed: 0x0102_0304_0506_0708,
ack_seq: 0x1122_3344_5566_7788,
};
let bytes = msg.encode();
assert_eq!(bytes.len(), STREAM_WINDOW_SIZE);
let parsed = StreamWindow::decode(&bytes).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn test_decode_truncated_rejected() {
let err = StreamWindow::decode(&[0u8; STREAM_WINDOW_SIZE - 1]).unwrap_err();
assert!(matches!(
err,
StreamWindowCodecError::Truncated {
need: STREAM_WINDOW_SIZE,
..
}
));
assert!(err.to_string().contains("need 24"), "got: {err}");
}
#[test]
fn test_decode_oversize_rejected() {
let err = StreamWindow::decode(&[0u8; STREAM_WINDOW_SIZE + 1]).unwrap_err();
assert!(matches!(
err,
StreamWindowCodecError::Oversize {
need: STREAM_WINDOW_SIZE,
..
}
));
}
#[test]
fn test_decode_empty_rejected() {
let err = StreamWindow::decode(&[]).unwrap_err();
assert!(matches!(
err,
StreamWindowCodecError::Truncated { got: 0, .. }
));
}
#[test]
fn test_endianness_is_little_endian() {
let msg = StreamWindow {
stream_id: 1,
total_consumed: 1,
ack_seq: 1,
};
let bytes = msg.encode();
assert_eq!(bytes[0], 0x01);
assert_eq!(bytes[1], 0x00);
assert_eq!(bytes[8], 0x01);
assert_eq!(bytes[9], 0x00);
assert_eq!(bytes[16], 0x01);
assert_eq!(bytes[17], 0x00);
}
#[test]
fn stream_nack_round_trip() {
let msg = StreamNack {
stream_id: 0xABCD,
next_expected: 7,
missing_bitmap: 0b1010,
};
assert_eq!(StreamNack::decode(&msg.encode()).unwrap(), msg);
}
#[test]
fn stream_reset_round_trip() {
let msg = StreamReset {
stream_id: 0x2000_0000_0000_0001,
};
assert_eq!(StreamReset::decode(&msg.encode()).unwrap(), msg);
}
fn ack(ack_seq: u64, ranges: &[(u64, u64)]) -> StreamAckRanges {
StreamAckRanges {
stream_id: 0xF00D,
ack_seq,
ranges: ranges.to_vec(),
}
}
#[test]
fn ack_ranges_round_trip_single_and_max() {
let one = ack(101, &[(102, 10_001)]);
assert_eq!(StreamAckRanges::decode(&one.encode()).unwrap(), one);
let ranges: Vec<(u64, u64)> = (0..MAX_ACK_RANGES as u64)
.map(|i| {
let hi = 10_000 - i * 100;
(hi - 10, hi)
})
.collect();
let full = ack(5, &ranges);
let bytes = full.encode();
assert_eq!(bytes.len(), StreamAckRanges::MAX_SIZE);
assert_eq!(StreamAckRanges::decode(&bytes).unwrap(), full);
}
#[test]
fn ack_ranges_rejects_rangeless_and_truncated() {
let err = StreamAckRanges::decode(&[0u8; STREAM_ACK_HEADER_SIZE]).unwrap_err();
assert!(matches!(err, StreamWindowCodecError::Truncated { .. }));
let err = StreamAckRanges::decode(&[]).unwrap_err();
assert!(matches!(err, StreamWindowCodecError::Truncated { .. }));
}
#[test]
fn ack_ranges_rejects_non_multiple_length() {
let bytes = ack(1, &[(2, 3)]).encode();
let mut long = bytes.clone();
long.push(0);
assert!(matches!(
StreamAckRanges::decode(&long).unwrap_err(),
StreamWindowCodecError::InvalidRanges { .. }
));
}
#[test]
fn ack_ranges_rejects_too_many_ranges() {
let ranges: Vec<(u64, u64)> = (0..(MAX_ACK_RANGES as u64 + 1))
.map(|i| {
let hi = 100_000 - i * 10;
(hi - 2, hi)
})
.collect();
let bytes = ack(1, &ranges).encode();
assert!(matches!(
StreamAckRanges::decode(&bytes).unwrap_err(),
StreamWindowCodecError::InvalidRanges {
reason: "range count exceeds MAX_ACK_RANGES"
}
));
}
#[test]
fn ack_ranges_rejects_bad_range_shapes() {
assert!(matches!(
StreamAckRanges::decode(&ack(1, &[(5, 5)]).encode()).unwrap_err(),
StreamWindowCodecError::InvalidRanges {
reason: "empty or inverted range"
}
));
assert!(StreamAckRanges::decode(&ack(1, &[(9, 5)]).encode()).is_err());
assert!(matches!(
StreamAckRanges::decode(&ack(10, &[(10, 12)]).encode()).unwrap_err(),
StreamWindowCodecError::InvalidRanges {
reason: "range not strictly above ack_seq"
}
));
assert!(StreamAckRanges::decode(&ack(10, &[(3, 6)]).encode()).is_err());
}
#[test]
fn ack_ranges_rejects_unmerged_or_ascending_order() {
assert!(matches!(
StreamAckRanges::decode(&ack(1, &[(2, 4), (6, 8)]).encode()).unwrap_err(),
StreamWindowCodecError::InvalidRanges {
reason: "ranges not descending/merged"
}
));
assert!(StreamAckRanges::decode(&ack(1, &[(6, 8), (4, 6)]).encode()).is_err());
assert!(StreamAckRanges::decode(&ack(1, &[(5, 9), (3, 7)]).encode()).is_err());
}
}