use std::borrow::Cow;
use crate::{
Path,
coding::{Decode, DecodeError, Encode, EncodeError, Sizer},
};
use super::{Message, Version};
#[derive(Clone, Debug)]
pub struct Subscribe<'a> {
pub id: u64,
pub broadcast: Path<'a>,
pub track: Cow<'a, str>,
pub priority: u8,
pub max_age: std::time::Duration,
pub start_group: Option<u64>,
pub end_group: Option<u64>,
pub start_frame: u64,
pub end_frame: Option<u64>,
}
impl Version {
pub(crate) fn carries_max_age(self) -> bool {
!matches!(self, Version::Lite01 | Version::Lite02)
}
}
impl Message for Subscribe<'_> {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let id = u64::decode(r, version)?;
let broadcast = Path::decode(r, version)?;
let track = Cow::<str>::decode(r, version)?;
let priority = u8::decode(r, version)?;
let (max_age, start_group, end_group) = match version {
Version::Lite01 | Version::Lite02 => (std::time::Duration::ZERO, None, None),
_ => {
skip_group_order(r, version)?;
let max_age = std::time::Duration::decode(r, version)?;
let start_group = decode_start_group(r, version)?;
let end_group = Option::<u64>::decode(r, version)?;
(max_age, start_group, end_group)
}
};
let (start_frame, end_frame) = decode_frame_bounds(r, version, start_group, end_group)?;
let start_group = canonical_start_group(version, start_group, start_frame);
Ok(Self {
id,
broadcast,
track,
priority,
max_age,
start_group,
end_group,
start_frame,
end_frame,
})
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
self.id.encode(w, version)?;
self.broadcast.encode(w, version)?;
self.track.encode(w, version)?;
self.priority.encode(w, version)?;
match version {
Version::Lite01 | Version::Lite02 => {}
_ => {
pad_group_order(w, version)?;
self.max_age.encode(w, version)?;
encode_start_group(w, version, self.start_group)?;
self.end_group.encode(w, version)?;
}
}
encode_frame_bounds(
w,
version,
self.start_group,
self.start_frame,
self.end_group,
self.end_frame,
)?;
Ok(())
}
}
pub(super) fn skip_group_order<R: bytes::Buf>(r: &mut R, version: Version) -> Result<(), DecodeError> {
if version.has_group_order() {
u8::decode(r, version)?;
}
Ok(())
}
pub(super) fn pad_group_order<W: bytes::BufMut>(w: &mut W, version: Version) -> Result<(), EncodeError> {
if version.has_group_order() {
0u8.encode(w, version)?;
}
Ok(())
}
fn decode_start_group<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Option<u64>, DecodeError> {
if version.resolves_start() {
return Ok(Some(u64::decode(r, version)?));
}
Option::<u64>::decode(r, version)
}
fn canonical_start_group(version: Version, start_group: Option<u64>, start_frame: u64) -> Option<u64> {
match (start_group, start_frame) {
(Some(0), 0) if version.resolves_start() => None,
(start_group, _) => start_group,
}
}
fn encode_start_group<W: bytes::BufMut>(
w: &mut W,
version: Version,
start_group: Option<u64>,
) -> Result<(), EncodeError> {
if version.resolves_start() {
return start_group.unwrap_or(0).encode(w, version);
}
start_group.filter(|&group| group > 0).encode(w, version)
}
fn decode_frame_bounds<R: bytes::Buf>(
r: &mut R,
version: Version,
start_group: Option<u64>,
end_group: Option<u64>,
) -> Result<(u64, Option<u64>), DecodeError> {
if !version.has_frame_bounds() {
return Ok((0, None));
}
let start_frame = u64::decode(r, version)?;
let end_frame = Option::<u64>::decode(r, version)?;
if (start_frame != 0 && start_group.is_none()) || (end_frame.is_some() && end_group.is_none()) {
return Err(DecodeError::InvalidSubscribeLocation);
}
Ok((start_frame, end_frame))
}
fn encode_frame_bounds<W: bytes::BufMut>(
w: &mut W,
version: Version,
start_group: Option<u64>,
start_frame: u64,
end_group: Option<u64>,
end_frame: Option<u64>,
) -> Result<(), EncodeError> {
if (start_frame != 0 && start_group.is_none()) || (end_frame.is_some() && end_group.is_none()) {
return Err(EncodeError::InvalidState);
}
if !version.has_frame_bounds() {
if start_frame != 0 || end_frame.is_some() {
return Err(EncodeError::Version);
}
return Ok(());
}
start_frame.encode(w, version)?;
end_frame.encode(w, version)
}
#[derive(Clone, Debug)]
pub struct SubscribeOk {
pub priority: u8,
pub max_age: std::time::Duration,
pub start_group: Option<u64>,
pub end_group: Option<u64>,
}
impl Message for SubscribeOk {
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
match version {
Version::Lite01 => {
self.priority.encode(w, version)?;
}
Version::Lite02 => {}
_ => {
self.priority.encode(w, version)?;
pad_group_order(w, version)?;
self.max_age.encode(w, version)?;
self.start_group.encode(w, version)?;
self.end_group.encode(w, version)?;
}
}
Ok(())
}
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
match version {
Version::Lite01 => Ok(Self {
priority: u8::decode(r, version)?,
max_age: std::time::Duration::ZERO,
start_group: None,
end_group: None,
}),
Version::Lite02 => Ok(Self {
priority: 0,
max_age: std::time::Duration::ZERO,
start_group: None,
end_group: None,
}),
_ => {
let priority = u8::decode(r, version)?;
skip_group_order(r, version)?;
let max_age = std::time::Duration::decode(r, version)?;
let start_group = Option::<u64>::decode(r, version)?;
let end_group = Option::<u64>::decode(r, version)?;
Ok(Self {
priority,
max_age,
start_group,
end_group,
})
}
}
}
}
#[derive(Clone, Debug)]
pub struct SubscribeStart {
pub group: u64,
}
impl Message for SubscribeStart {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
if !version.has_track_stream() {
return Err(DecodeError::Version);
}
Ok(Self {
group: u64::decode(r, version)?,
})
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if !version.has_track_stream() {
return Err(EncodeError::Version);
}
self.group.encode(w, version)
}
}
#[derive(Clone, Debug)]
pub struct SubscribeEnd {
pub group: u64,
}
impl Message for SubscribeEnd {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
if !version.has_track_stream() {
return Err(DecodeError::Version);
}
Ok(Self {
group: u64::decode(r, version)?,
})
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if !version.has_track_stream() {
return Err(EncodeError::Version);
}
self.group.encode(w, version)
}
}
#[allow(dead_code)]
#[derive(Clone, Debug)]
pub struct SubscribeUpdate {
pub priority: u8,
pub max_age: std::time::Duration,
pub start_group: Option<u64>,
pub end_group: Option<u64>,
pub start_frame: u64,
pub end_frame: Option<u64>,
}
impl Message for SubscribeUpdate {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
match version {
Version::Lite01 | Version::Lite02 => {
return Err(DecodeError::Version);
}
_ => {}
}
let priority = u8::decode(r, version)?;
skip_group_order(r, version)?;
let max_age = std::time::Duration::decode(r, version)?;
let start_group = decode_start_group(r, version)?;
let end_group = match u64::decode(r, version)? {
0 => None,
group => Some(group - 1),
};
let (start_frame, end_frame) = decode_frame_bounds(r, version, start_group, end_group)?;
let start_group = canonical_start_group(version, start_group, start_frame);
Ok(Self {
priority,
max_age,
start_group,
end_group,
start_frame,
end_frame,
})
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
match version {
Version::Lite01 | Version::Lite02 => {
return Err(EncodeError::Version);
}
_ => {}
}
self.priority.encode(w, version)?;
pad_group_order(w, version)?;
self.max_age.encode(w, version)?;
encode_start_group(w, version, self.start_group)?;
match self.end_group {
Some(end_group) => end_group
.checked_add(1)
.ok_or(EncodeError::TooLarge)?
.encode(w, version)?,
None => 0u64.encode(w, version)?,
}
encode_frame_bounds(
w,
version,
self.start_group,
self.start_frame,
self.end_group,
self.end_frame,
)?;
Ok(())
}
}
#[derive(Clone, Debug)]
pub struct SubscribeDrop {
pub start: u64,
pub end: u64,
pub error: u64,
}
impl Message for SubscribeDrop {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
match version {
Version::Lite01 | Version::Lite02 => {
return Err(DecodeError::Version);
}
_ => {}
}
Ok(Self {
start: u64::decode(r, version)?,
end: u64::decode(r, version)?,
error: u64::decode(r, version)?,
})
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
match version {
Version::Lite01 | Version::Lite02 => {
return Err(EncodeError::Version);
}
_ => {}
}
self.start.encode(w, version)?;
self.end.encode(w, version)?;
self.error.encode(w, version)?;
Ok(())
}
}
#[derive(Clone, Debug)]
pub enum SubscribeResponse {
Ok(SubscribeOk),
Start(SubscribeStart),
End(SubscribeEnd),
Drop(SubscribeDrop),
}
fn encode_typed<W: bytes::BufMut, M: Message>(
w: &mut W,
typ: u64,
msg: &M,
version: Version,
) -> Result<(), EncodeError> {
typ.encode(w, version)?;
let mut sizer = Sizer::default();
msg.encode_msg(&mut sizer, version)?;
sizer.size.encode(w, version)?;
msg.encode_msg(w, version)
}
impl Encode<Version> for SubscribeResponse {
fn encode<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
match version {
Version::Lite01 | Version::Lite02 => match self {
Self::Ok(ok) => {
let mut sizer = Sizer::default();
Message::encode_msg(ok, &mut sizer, version)?;
sizer.size.encode(w, version)?;
Message::encode_msg(ok, w, version)?;
}
_ => return Err(EncodeError::Version),
},
Version::Lite03 | Version::Lite04 => match self {
Self::Ok(ok) => encode_typed(w, 0, ok, version)?,
Self::Drop(drop) => encode_typed(w, 1, drop, version)?,
_ => return Err(EncodeError::Version),
},
_ => match self {
Self::Start(start) => encode_typed(w, 0, start, version)?,
Self::End(end) => encode_typed(w, 1, end, version)?,
Self::Drop(drop) => encode_typed(w, 2, drop, version)?,
Self::Ok(_) => return Err(EncodeError::Version),
},
}
Ok(())
}
}
impl Decode<Version> for SubscribeResponse {
fn decode<B: bytes::Buf>(buf: &mut B, version: Version) -> Result<Self, DecodeError> {
match version {
Version::Lite01 | Version::Lite02 => Ok(Self::Ok(SubscribeOk::decode(buf, version)?)),
Version::Lite03 | Version::Lite04 => {
let typ = u64::decode(buf, version)?;
match typ {
0 => Ok(Self::Ok(SubscribeOk::decode(buf, version)?)),
1 => Ok(Self::Drop(SubscribeDrop::decode(buf, version)?)),
_ => Err(DecodeError::InvalidMessage(typ)),
}
}
_ => {
let typ = u64::decode(buf, version)?;
match typ {
0 => Ok(Self::Start(SubscribeStart::decode(buf, version)?)),
1 => Ok(Self::End(SubscribeEnd::decode(buf, version)?)),
2 => Ok(Self::Drop(SubscribeDrop::decode(buf, version)?)),
_ => Err(DecodeError::InvalidMessage(typ)),
}
}
}
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn subscribe_start_roundtrips_on_lite05() {
let resp = SubscribeResponse::Start(SubscribeStart { group: 42 });
let mut buf = Vec::new();
resp.encode(&mut buf, Version::Lite05).unwrap();
let mut slice = buf.as_slice();
match SubscribeResponse::decode(&mut slice, Version::Lite05).unwrap() {
SubscribeResponse::Start(start) => assert_eq!(start.group, 42),
other => panic!("expected Start, got {other:?}"),
}
}
#[test]
fn subscribe_end_roundtrips_on_lite05() {
let resp = SubscribeResponse::End(SubscribeEnd { group: 7 });
let mut buf = Vec::new();
resp.encode(&mut buf, Version::Lite05).unwrap();
let mut slice = buf.as_slice();
match SubscribeResponse::decode(&mut slice, Version::Lite05).unwrap() {
SubscribeResponse::End(end) => assert_eq!(end.group, 7),
other => panic!("expected End, got {other:?}"),
}
}
#[test]
fn subscribe_drop_is_type_2_on_lite05() {
let resp = SubscribeResponse::Drop(SubscribeDrop {
start: 1,
end: 3,
error: 0,
});
let mut buf = Vec::new();
resp.encode(&mut buf, Version::Lite05).unwrap();
assert_eq!(buf[0], 2);
let mut slice = buf.as_slice();
match SubscribeResponse::decode(&mut slice, Version::Lite05).unwrap() {
SubscribeResponse::Drop(drop) => assert_eq!((drop.start, drop.end), (1, 3)),
other => panic!("expected Drop, got {other:?}"),
}
}
#[test]
fn subscribe_drop_is_type_1_on_lite04() {
let resp = SubscribeResponse::Drop(SubscribeDrop {
start: 1,
end: 3,
error: 0,
});
let mut buf = Vec::new();
resp.encode(&mut buf, Version::Lite04).unwrap();
assert_eq!(buf[0], 1);
}
fn subscribe_sample() -> Subscribe<'static> {
Subscribe {
id: 1,
broadcast: Path::new("room").to_owned(),
track: Cow::Borrowed("video"),
priority: 3,
max_age: std::time::Duration::from_millis(250),
start_group: Some(7),
end_group: Some(9),
start_frame: 4,
end_frame: Some(2),
}
}
#[test]
fn subscribe_frame_bounds_roundtrip() {
let msg = subscribe_sample();
let mut buf = Vec::new();
msg.encode_msg(&mut buf, Version::Lite06).unwrap();
let got = Subscribe::decode_msg(&mut buf.as_slice(), Version::Lite06).unwrap();
assert_eq!((got.start_group, got.start_frame), (Some(7), 4));
assert_eq!((got.end_group, got.end_frame), (Some(9), Some(2)));
}
#[test]
fn subscribe_drops_the_retired_ordered_byte_on_lite06() {
let mut msg = subscribe_sample();
msg.start_group = None;
msg.start_frame = 0;
msg.end_frame = None;
let mut lite05 = Vec::new();
msg.encode_msg(&mut lite05, Version::Lite05).unwrap();
let mut lite06 = Vec::new();
msg.encode_msg(&mut lite06, Version::Lite06).unwrap();
let ordered_at = lite05
.iter()
.zip(&lite06)
.position(|(a, b)| a != b)
.expect("the layouts must diverge at the retired byte");
assert_eq!(lite05[ordered_at], 0, "the retired byte is written as zero");
let mut spliced = lite05.clone();
spliced.remove(ordered_at);
assert_eq!(&lite06[..spliced.len()], &spliced[..]);
assert_eq!(&lite06[spliced.len()..], &[0, 0]);
let got = Subscribe::decode_msg(&mut lite05.as_slice(), Version::Lite05).unwrap();
assert_eq!((got.start_frame, got.end_frame), (0, None));
}
#[test]
fn group_start_is_absolute_on_lite06() {
let mut msg = subscribe_sample();
msg.start_frame = 0;
msg.end_frame = None;
let mut lite05 = Vec::new();
msg.encode_msg(&mut lite05, Version::Lite05).unwrap();
let mut lite06 = Vec::new();
msg.encode_msg(&mut lite06, Version::Lite06).unwrap();
let on05 = Subscribe::decode_msg(&mut lite05.as_slice(), Version::Lite05).unwrap();
let on06 = Subscribe::decode_msg(&mut lite06.as_slice(), Version::Lite06).unwrap();
assert_eq!(on05.start_group, Some(7));
assert_eq!(on06.start_group, Some(7));
assert_ne!(lite05, lite06);
msg.start_group = None;
let mut absent = Vec::new();
msg.encode_msg(&mut absent, Version::Lite06).unwrap();
msg.start_group = Some(0);
let mut zero = Vec::new();
msg.encode_msg(&mut zero, Version::Lite06).unwrap();
assert_eq!(absent, zero);
let got = Subscribe::decode_msg(&mut zero.as_slice(), Version::Lite06).unwrap();
assert_eq!(got.start_group, None);
let mut folded = Vec::new();
msg.encode_msg(&mut folded, Version::Lite05).unwrap();
let got = Subscribe::decode_msg(&mut folded.as_slice(), Version::Lite05).unwrap();
assert_eq!(got.start_group, None);
}
#[test]
fn frame_start_may_qualify_group_zero_on_lite06() {
let mut msg = subscribe_sample();
msg.start_group = Some(0);
msg.start_frame = 4;
let mut buf = Vec::new();
msg.encode_msg(&mut buf, Version::Lite06).unwrap();
let got = Subscribe::decode_msg(&mut buf.as_slice(), Version::Lite06).unwrap();
assert_eq!((got.start_group, got.start_frame), (Some(0), 4));
}
#[test]
fn subscribe_frame_bounds_rejected_before_lite06() {
let mut buf = Vec::new();
assert!(subscribe_sample().encode_msg(&mut buf, Version::Lite05).is_err());
}
#[test]
fn subscribe_frame_bound_without_group_bound_is_invalid() {
let mut msg = subscribe_sample();
msg.start_group = None;
msg.start_frame = 4;
msg.end_group = None;
msg.end_frame = None;
let mut buf = Vec::new();
assert!(matches!(
msg.encode_msg(&mut buf, Version::Lite06),
Err(EncodeError::InvalidState)
));
msg.start_frame = 0;
msg.end_frame = Some(7);
assert!(matches!(
msg.encode_msg(&mut buf, Version::Lite06),
Err(EncodeError::InvalidState)
));
}
#[test]
fn subscribe_ok_rejected_on_lite05() {
let resp = SubscribeResponse::Ok(SubscribeOk {
priority: 1,
max_age: std::time::Duration::ZERO,
start_group: None,
end_group: None,
});
let mut buf = Vec::new();
assert!(resp.encode(&mut buf, Version::Lite05).is_err());
}
}