use std::borrow::Cow;
use crate::{
Path,
coding::{Decode, DecodeError, Encode, EncodeError},
ietf::{
Filter, GroupOrder, Location, Parameters, Properties, RequestId,
namespace::{decode_namespace, encode_namespace},
},
};
use super::Message;
use super::Version;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum PublishDoneStatus {
InternalError,
TrackEnded,
}
impl PublishDoneStatus {
pub(crate) const fn code(self, version: Version) -> u64 {
match version {
Version::Draft14
| Version::Draft15
| Version::Draft16
| Version::Draft17
| Version::Draft18
| Version::Draft19
| Version::Draft20 => match self {
Self::InternalError => 0x0,
Self::TrackEnded => 0x2,
},
}
}
}
#[derive(Clone, Debug)]
pub struct PublishDone<'a> {
pub request_id: Option<RequestId>,
pub status_code: u64,
pub stream_count: u64,
pub reason_phrase: Cow<'a, str>,
}
impl Message for PublishDone<'_> {
const ID: u64 = 0x0b;
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) {
self.request_id
.expect("request_id required for draft14-16")
.encode(w, version)?;
} else {
assert!(self.request_id.is_none(), "request_id must be None for draft17+");
}
self.status_code.encode(w, version)?;
self.stream_count.encode(w, version)?;
self.reason_phrase.encode(w, version)?;
Ok(())
}
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) {
Some(RequestId::decode(r, version)?)
} else {
None
};
let status_code = u64::decode(r, version)?;
let stream_count = u64::decode(r, version)?;
let reason_phrase = Cow::<str>::decode(r, version)?;
Ok(Self {
request_id,
status_code,
stream_count,
reason_phrase,
})
}
}
#[derive(Debug)]
pub struct Publish<'a> {
pub request_id: RequestId,
pub track_namespace: Path<'a>,
pub track_name: Cow<'a, str>,
pub track_alias: u64,
pub largest_location: Option<Location>,
pub forward: bool,
pub properties: Properties,
}
impl Message for Publish<'_> {
const ID: u64 = 0x1D;
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
self.request_id.encode(w, version)?;
if version == Version::Draft17 {
0u64.encode(w, version)?; }
encode_namespace(w, &self.track_namespace, version)?;
self.track_name.encode(w, version)?;
self.track_alias.encode(w, version)?;
match version {
Version::Draft14 => {
self.properties
.group_order
.unwrap_or(GroupOrder::Ascending)
.encode(w, version)?;
if let Some(location) = &self.largest_location {
true.encode(w, version)?;
location.encode(w, version)?;
} else {
false.encode(w, version)?;
}
self.forward.encode(w, version)?;
0u8.encode(w, version)?;
}
_ => {
let group_order = match version {
Version::Draft15 => self.properties.group_order,
_ => None,
};
encode_params!(w, version,
0x09 => self.largest_location,
0x10 => self.forward,
0x22 => group_order,
);
self.properties.encode(w, version)?;
}
}
Ok(())
}
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let request_id = RequestId::decode(r, version)?;
if version == Version::Draft17 {
let _required_request_id_delta = u64::decode(r, version)?;
}
let track_namespace = decode_namespace(r, version)?;
let track_name = Cow::<str>::decode(r, version)?;
let track_alias = u64::decode(r, version)?;
match version {
Version::Draft14 => {
let group_order = GroupOrder::decode(r, version)?.any_to_descending();
let content_exists = bool::decode(r, version)?;
let largest_location = match content_exists {
true => Some(Location::decode(r, version)?),
false => None,
};
let forward = bool::decode(r, version)?;
let _params = Parameters::decode(r, version)?;
Ok(Self {
request_id,
track_namespace,
track_name,
track_alias,
largest_location,
forward,
properties: Properties {
group_order: Some(group_order),
..Default::default()
},
})
}
_ => {
decode_params!(r, version,
0x02 => object_delivery_timeout: Option<u64>,
0x06 => subgroup_delivery_timeout: Option<u64>,
0x08 => _expires: Option<u64>,
0x09 => largest_location: Option<Location>,
0x10 => forward: Option<bool>,
0x20 => subscriber_priority: Option<u8>,
0x21 => filter: Option<Filter>,
0x22 => group_order: Option<GroupOrder>,
);
let subscription_params = [
object_delivery_timeout.is_some(),
subgroup_delivery_timeout.is_some(),
subscriber_priority.is_some(),
filter.is_some(),
];
if subscription_params.contains(&true) && !Filter::is_draft20(version) {
return Err(DecodeError::InvalidValue);
}
let mut properties = Properties::decode(r, version)?;
properties.group_order = properties.group_order.or(group_order);
let forward = forward.unwrap_or(true);
Ok(Self {
request_id,
track_namespace,
track_name,
track_alias,
largest_location,
forward,
properties,
})
}
}
}
}
#[derive(Debug)]
pub struct PublishOk {
pub request_id: Option<RequestId>,
pub forward: bool,
pub subscriber_priority: u8,
pub group_order: GroupOrder,
pub filter: Filter,
}
impl Message for PublishOk {
const ID: u64 = 0x1E;
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) {
self.request_id
.expect("request_id required for draft14-16")
.encode(w, version)?;
} else {
assert!(self.request_id.is_none(), "request_id must be None for draft17+");
}
match version {
Version::Draft14 => {
self.forward.encode(w, version)?;
self.subscriber_priority.encode(w, version)?;
self.group_order.encode(w, version)?;
self.filter.encode(w, version)?;
0u8.encode(w, version)?;
}
_ if Filter::is_draft20(version) => encode_params!(w, version,),
_ => {
encode_params!(w, version,
0x10 => self.forward,
0x20 => self.subscriber_priority,
0x21 => self.filter,
0x22 => self.group_order,
);
}
}
Ok(())
}
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) {
Some(RequestId::decode(r, version)?)
} else {
None
};
match version {
Version::Draft14 => {
let forward = bool::decode(r, version)?;
let subscriber_priority = u8::decode(r, version)?;
let group_order = GroupOrder::decode(r, version)?;
let filter = Filter::decode(r, version)?;
let _params = Parameters::decode(r, version)?;
Ok(Self {
request_id,
forward,
subscriber_priority,
group_order,
filter,
})
}
_ => {
decode_params!(r, version,
0x10 => forward: Option<bool>,
0x20 => subscriber_priority: Option<u8>,
0x21 => filter: Option<Filter>,
0x22 => group_order: Option<GroupOrder>,
);
let forward = forward.unwrap_or(true);
let subscriber_priority = subscriber_priority.unwrap_or(128);
let group_order = group_order.unwrap_or(GroupOrder::Descending);
let filter = filter.unwrap_or(Filter::Unfiltered);
Ok(Self {
request_id,
forward,
subscriber_priority,
group_order,
filter,
})
}
}
}
}
#[derive(Debug)]
pub struct PublishError<'a> {
pub request_id: RequestId,
pub error_code: u64,
pub reason_phrase: Cow<'a, str>,
}
impl Message for PublishError<'_> {
const ID: u64 = 0x1F;
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
self.request_id.encode(w, version)?;
self.error_code.encode(w, version)?;
self.reason_phrase.encode(w, version)?;
Ok(())
}
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let request_id = RequestId::decode(r, version)?;
let error_code = u64::decode(r, version)?;
let reason_phrase = Cow::<str>::decode(r, version)?;
Ok(Self {
request_id,
error_code,
reason_phrase,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn publish_accepts_the_relocated_subscription_parameters() -> Result<(), EncodeError> {
let mut body = Vec::new();
RequestId(1).encode(&mut body, Version::Draft20).unwrap();
super::super::namespace::encode_namespace(&mut body, &crate::Path::new("broadcast"), Version::Draft20).unwrap();
"video".encode(&mut body, Version::Draft20).unwrap();
1u64.encode(&mut body, Version::Draft20).unwrap();
encode_params!(&mut body, Version::Draft20,
0x20 => 128u8,
0x21 => Filter::NextObject,
);
Properties::default().encode(&mut body, Version::Draft20).unwrap();
let mut buf = bytes::Bytes::from(body);
Publish::decode_msg(&mut buf, Version::Draft20).expect("draft-20 PUBLISH parameters must parse");
Ok(())
}
#[test]
fn older_drafts_reject_the_relocated_parameters() -> Result<(), EncodeError> {
let mut body = Vec::new();
RequestId(1).encode(&mut body, Version::Draft19).unwrap();
super::super::namespace::encode_namespace(&mut body, &crate::Path::new("broadcast"), Version::Draft19).unwrap();
"video".encode(&mut body, Version::Draft19).unwrap();
1u64.encode(&mut body, Version::Draft19).unwrap();
encode_params!(&mut body, Version::Draft19, 0x20 => 128u8);
Properties::default().encode(&mut body, Version::Draft19).unwrap();
let mut buf = bytes::Bytes::from(body);
assert!(Publish::decode_msg(&mut buf, Version::Draft19).is_err());
Ok(())
}
use bytes::BytesMut;
fn encode_message<M: Message>(msg: &M, version: Version) -> Vec<u8> {
let mut buf = BytesMut::new();
msg.encode_msg(&mut buf, version).unwrap();
buf.to_vec()
}
fn decode_message<M: Message>(bytes: &[u8], version: Version) -> Result<M, DecodeError> {
let mut buf = bytes::Bytes::from(bytes.to_vec());
M::decode_msg(&mut buf, version)
}
#[test]
fn test_publish_v14_round_trip() {
let msg = Publish {
request_id: RequestId(1),
track_namespace: Path::new("test/ns"),
track_name: "video".into(),
track_alias: 42,
largest_location: Some(Location { group: 10, object: 5 }),
forward: true,
properties: Properties {
group_order: Some(GroupOrder::Descending),
..Default::default()
},
};
let encoded = encode_message(&msg, Version::Draft14);
let decoded: Publish = decode_message(&encoded, Version::Draft14).unwrap();
assert_eq!(decoded.request_id, RequestId(1));
assert_eq!(decoded.track_namespace.as_str(), "test/ns");
assert_eq!(decoded.track_name, "video");
assert_eq!(decoded.track_alias, 42);
assert_eq!(decoded.largest_location, Some(Location { group: 10, object: 5 }));
assert!(decoded.forward);
}
#[test]
fn test_publish_v15_round_trip() {
let msg = Publish {
request_id: RequestId(1),
track_namespace: Path::new("test/ns"),
track_name: "video".into(),
track_alias: 42,
largest_location: Some(Location { group: 10, object: 5 }),
forward: true,
properties: Properties {
group_order: Some(GroupOrder::Descending),
..Default::default()
},
};
let encoded = encode_message(&msg, Version::Draft15);
let decoded: Publish = decode_message(&encoded, Version::Draft15).unwrap();
assert_eq!(decoded.request_id, RequestId(1));
assert_eq!(decoded.track_namespace.as_str(), "test/ns");
assert_eq!(decoded.track_name, "video");
assert_eq!(decoded.track_alias, 42);
assert_eq!(decoded.largest_location, Some(Location { group: 10, object: 5 }));
assert!(decoded.forward);
}
#[test]
fn test_publish_ok_v14_round_trip() {
let msg = PublishOk {
request_id: Some(RequestId(7)),
forward: true,
subscriber_priority: 128,
group_order: GroupOrder::Descending,
filter: Filter::NextObject,
};
let encoded = encode_message(&msg, Version::Draft14);
let decoded: PublishOk = decode_message(&encoded, Version::Draft14).unwrap();
assert_eq!(decoded.request_id, Some(RequestId(7)));
assert!(decoded.forward);
assert_eq!(decoded.subscriber_priority, 128);
}
#[test]
fn test_publish_ok_v15_round_trip() {
let msg = PublishOk {
request_id: Some(RequestId(7)),
forward: true,
subscriber_priority: 128,
group_order: GroupOrder::Descending,
filter: Filter::NextObject,
};
let encoded = encode_message(&msg, Version::Draft15);
let decoded: PublishOk = decode_message(&encoded, Version::Draft15).unwrap();
assert_eq!(decoded.request_id, Some(RequestId(7)));
assert!(decoded.forward);
assert_eq!(decoded.subscriber_priority, 128);
}
#[test]
fn test_publish_v17_round_trip() {
let msg = Publish {
request_id: RequestId(1),
track_namespace: Path::new("test/ns"),
track_name: "video".into(),
track_alias: 42,
largest_location: Some(Location { group: 10, object: 5 }),
forward: true,
properties: Properties {
group_order: Some(GroupOrder::Descending),
..Default::default()
},
};
let encoded = encode_message(&msg, Version::Draft17);
let decoded: Publish = decode_message(&encoded, Version::Draft17).unwrap();
assert_eq!(decoded.request_id, RequestId(1));
assert_eq!(decoded.track_namespace.as_str(), "test/ns");
assert_eq!(decoded.track_name, "video");
assert_eq!(decoded.track_alias, 42);
assert_eq!(decoded.largest_location, Some(Location { group: 10, object: 5 }));
assert!(decoded.forward);
}
#[test]
fn test_publish_ok_v17_round_trip() {
let msg = PublishOk {
request_id: None,
forward: true,
subscriber_priority: 128,
group_order: GroupOrder::Descending,
filter: Filter::NextObject,
};
let encoded = encode_message(&msg, Version::Draft17);
let decoded: PublishOk = decode_message(&encoded, Version::Draft17).unwrap();
assert_eq!(decoded.request_id, None);
assert!(decoded.forward);
assert_eq!(decoded.subscriber_priority, 128);
}
#[test]
fn test_publish_done_registered_statuses_all_versions() {
for version in [
Version::Draft14,
Version::Draft15,
Version::Draft16,
Version::Draft17,
Version::Draft18,
Version::Draft19,
] {
for (status, expected) in [
(PublishDoneStatus::InternalError, 0x0),
(PublishDoneStatus::TrackEnded, 0x2),
] {
assert_eq!(status.code(version), expected);
let msg = PublishDone {
request_id: matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16)
.then_some(RequestId(7)),
status_code: status.code(version),
stream_count: 5,
reason_phrase: "done".into(),
};
let encoded = encode_message(&msg, version);
let decoded: PublishDone = decode_message(&encoded, version).unwrap();
assert_eq!(decoded.status_code, expected);
assert_eq!(decoded.stream_count, 5);
assert_eq!(decoded.reason_phrase, "done");
}
}
}
#[test]
fn test_publish_v18_round_trip() {
let msg = Publish {
request_id: RequestId(1),
track_namespace: Path::new("test/ns"),
track_name: "video".into(),
track_alias: 42,
largest_location: Some(Location { group: 10, object: 5 }),
forward: true,
properties: Properties {
group_order: Some(GroupOrder::Descending),
..Default::default()
},
};
let encoded = encode_message(&msg, Version::Draft18);
let decoded: Publish = decode_message(&encoded, Version::Draft18).unwrap();
assert_eq!(decoded.request_id, RequestId(1));
assert_eq!(decoded.track_namespace.as_str(), "test/ns");
assert_eq!(decoded.track_name, "video");
assert_eq!(decoded.track_alias, 42);
assert_eq!(decoded.largest_location, Some(Location { group: 10, object: 5 }));
assert!(decoded.forward);
}
#[test]
fn test_publish_v18_group_order_is_a_property() {
let msg = Publish {
request_id: RequestId(1),
track_namespace: Path::new("ns"),
track_name: "video".into(),
track_alias: 42,
largest_location: None,
forward: true,
properties: Properties {
timescale: None,
group_order: Some(GroupOrder::Descending),
},
};
#[rustfmt::skip]
let expected = vec![
1, 1, 2, b'n', b's', 5, b'v', b'i', b'd', b'e', b'o', 42, 1, 0x10, 0x01, 0x22, 0x02, ];
assert_eq!(encode_message(&msg, Version::Draft18), expected);
let decoded: Publish = decode_message(&expected, Version::Draft18).unwrap();
assert_eq!(decoded.properties.group_order, Some(GroupOrder::Descending));
}
#[test]
fn test_publish_v15_group_order_is_a_parameter() {
let msg = Publish {
request_id: RequestId(1),
track_namespace: Path::new("ns"),
track_name: "video".into(),
track_alias: 42,
largest_location: None,
forward: true,
properties: Properties {
timescale: None,
group_order: Some(GroupOrder::Descending),
},
};
#[rustfmt::skip]
let expected = vec![
1, 1, 2, b'n', b's', 5, b'v', b'i', b'd', b'e', b'o', 42, 2, 0x10, 0x01, 0x22, 0x02, ];
assert_eq!(encode_message(&msg, Version::Draft15), expected);
let decoded: Publish = decode_message(&expected, Version::Draft15).unwrap();
assert_eq!(decoded.properties.group_order, Some(GroupOrder::Descending));
}
#[test]
fn test_publish_v18_is_one_byte_shorter_than_v17() {
let msg = Publish {
request_id: RequestId(1),
track_namespace: Path::new("test/ns"),
track_name: "video".into(),
track_alias: 42,
largest_location: None,
forward: true,
properties: Properties {
group_order: Some(GroupOrder::Descending),
..Default::default()
},
};
let v17 = encode_message(&msg, Version::Draft17);
let v18 = encode_message(&msg, Version::Draft18);
assert_eq!(v17.len(), v18.len() + 1);
}
#[test]
fn test_publish_ok_v18_round_trip() {
let msg = PublishOk {
request_id: None,
forward: true,
subscriber_priority: 128,
group_order: GroupOrder::Descending,
filter: Filter::NextObject,
};
let encoded = encode_message(&msg, Version::Draft18);
let decoded: PublishOk = decode_message(&encoded, Version::Draft18).unwrap();
assert_eq!(decoded.request_id, None);
assert!(decoded.forward);
assert_eq!(decoded.subscriber_priority, 128);
}
#[test]
fn test_publish_done_preserves_unknown_status() {
const UNKNOWN: u64 = 0xface;
let msg = PublishDone {
request_id: None,
status_code: UNKNOWN,
stream_count: 5,
reason_phrase: "extension".into(),
};
let encoded = encode_message(&msg, Version::Draft19);
let decoded: PublishDone = decode_message(&encoded, Version::Draft19).unwrap();
assert_eq!(decoded.status_code, UNKNOWN);
}
}