use super::types::ObjectStatus;
use crate::error::CodecError;
use crate::varint::VarInt;
use bytes::{Buf, BufMut};
fn skip(buf: &mut impl Buf, len: u64) -> Result<(), CodecError> {
let len = usize::try_from(len).map_err(|_| CodecError::UnexpectedEnd)?;
if buf.remaining() < len {
return Err(CodecError::UnexpectedEnd);
}
buf.advance(len);
Ok(())
}
fn decoded_status(code: u64) -> Result<ObjectStatus, CodecError> {
ObjectStatus::from_u64(code).ok_or(CodecError::InvalidField)
}
fn check_extensions_against_status_on_encode(
status: Option<u64>,
extension_headers_len: u64,
) -> Result<(), CodecError> {
if extension_headers_len != 0
&& matches!(status, Some(code) if code != ObjectStatus::Normal.as_u64())
{
return Err(CodecError::InvalidField);
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PayloadPermission {
Permitted,
Forbidden,
}
impl PayloadPermission {
pub fn for_status(status: ObjectStatus) -> Self {
match status {
ObjectStatus::Normal => PayloadPermission::Permitted,
ObjectStatus::ObjectDoesNotExist => PayloadPermission::Forbidden,
ObjectStatus::EndOfGroup => PayloadPermission::Forbidden,
ObjectStatus::EndOfTrack => PayloadPermission::Forbidden,
}
}
pub fn permits(self) -> bool {
matches!(self, PayloadPermission::Permitted)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SubgroupHeader {
pub header_type: u8,
pub track_alias: VarInt,
pub group_id: VarInt,
pub subgroup_id: VarInt,
pub publisher_priority: Option<u8>,
}
fn subgroup_type_is_assigned(ty: u64) -> bool {
let Ok(ty) = u8::try_from(ty) else {
return false;
};
(ty & 0xD0) == 0x10 && (ty & 0x06) != 0x06
}
impl SubgroupHeader {
pub fn has_extensions(&self) -> bool {
self.header_type & 0x01 != 0
}
pub fn subgroup_id_from_first_object(&self) -> bool {
self.header_type & 0x06 == 0x02
}
pub fn has_explicit_subgroup_id(&self) -> bool {
self.header_type & 0x06 == 0x04
}
pub fn has_end_of_group(&self) -> bool {
self.header_type & 0x08 != 0
}
pub fn has_priority(&self) -> bool {
self.header_type & 0x20 == 0
}
pub fn encode(&self, buf: &mut impl BufMut) {
VarInt::from_usize(self.header_type as usize).encode(buf);
self.track_alias.encode(buf);
self.group_id.encode(buf);
if self.has_explicit_subgroup_id() {
self.subgroup_id.encode(buf);
}
if self.has_priority() {
buf.put_u8(self.publisher_priority.unwrap_or(0));
}
}
pub fn encode_checked(&self, buf: &mut impl BufMut) -> Result<(), CodecError> {
if self.has_priority() != self.publisher_priority.is_some() {
return Err(CodecError::InvalidField);
}
if !subgroup_type_is_assigned(self.header_type as u64) {
return Err(CodecError::InvalidField);
}
self.encode(buf);
Ok(())
}
pub fn decode(buf: &mut impl Buf) -> Result<Self, CodecError> {
let type_value = VarInt::decode(buf)?.into_inner();
if !subgroup_type_is_assigned(type_value) {
return Err(stream_type_error(type_value));
}
let header_type = type_value as u8;
let track_alias = VarInt::decode(buf)?;
let group_id = VarInt::decode(buf)?;
let subgroup_id = if type_value as u8 & 0x06 == 0x04 {
VarInt::decode(buf)?
} else {
VarInt::from_usize(0)
};
let publisher_priority = if header_type & 0x20 == 0 {
if buf.remaining() < 1 {
return Err(CodecError::UnexpectedEnd);
}
Some(buf.get_u8())
} else {
None
};
Ok(Self { header_type, track_alias, group_id, subgroup_id, publisher_priority })
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SubgroupObject {
pub object_id: VarInt,
pub extension_headers: Vec<u8>,
pub payload_length: VarInt,
pub object_status: Option<ObjectStatus>,
pub payload: Vec<u8>,
}
impl SubgroupObject {
pub fn status(&self) -> ObjectStatus {
if self.payload_length.into_inner() == 0 {
self.object_status.unwrap_or(ObjectStatus::Normal)
} else {
ObjectStatus::Normal
}
}
pub fn extensions_permitted(&self) -> bool {
self.extension_headers.is_empty() || self.status() == ObjectStatus::Normal
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SubgroupObjectMeta {
pub object_id: u64,
pub extension_headers_len: u64,
pub payload_length: u64,
pub status: Option<u64>,
pub wire_len: u64,
}
impl SubgroupObjectMeta {
pub fn payload_permission(&self) -> Option<PayloadPermission> {
match self.status {
None => Some(PayloadPermission::Permitted),
Some(code) => ObjectStatus::from_u64(code).map(PayloadPermission::for_status),
}
}
pub fn extensions_permitted(&self) -> Option<bool> {
if self.extension_headers_len == 0 {
return Some(true);
}
match self.status {
None => Some(true),
Some(code) => ObjectStatus::from_u64(code).map(|s| s == ObjectStatus::Normal),
}
}
}
#[derive(Debug, Clone)]
pub struct SubgroupObjectReader {
extensions_present: bool,
prev_object_id: Option<u64>,
}
impl SubgroupObjectReader {
pub fn new(header: &SubgroupHeader) -> Self {
Self { extensions_present: header.has_extensions(), prev_object_id: None }
}
pub fn read_object(&mut self, buf: &mut impl Buf) -> Result<SubgroupObject, CodecError> {
let delta = VarInt::decode(buf)?.into_inner();
let object_id_val = match self.prev_object_id {
None => delta,
Some(prev) => prev
.checked_add(1)
.and_then(|v| v.checked_add(delta))
.ok_or(CodecError::InvalidField)?,
};
self.prev_object_id = Some(object_id_val);
let object_id = VarInt::from_u64(object_id_val).map_err(|_| CodecError::InvalidField)?;
let extension_headers = if self.extensions_present {
let ext_len = VarInt::decode(buf)?.into_inner() as usize;
crate::types::read_bytes(buf, ext_len)?
} else {
Vec::new()
};
let payload_length_vi = VarInt::decode(buf)?;
let payload_length_val = payload_length_vi.into_inner() as usize;
let (object_status, payload) = if payload_length_val == 0 {
let status = decoded_status(VarInt::decode(buf)?.into_inner())?;
(Some(status), Vec::new())
} else {
let payload = crate::types::read_bytes(buf, payload_length_val)?;
(None, payload)
};
Ok(SubgroupObject {
object_id,
extension_headers,
payload_length: payload_length_vi,
object_status,
payload,
})
}
pub fn read_object_meta(
&mut self,
buf: &mut impl Buf,
) -> Result<SubgroupObjectMeta, CodecError> {
let start = buf.remaining();
let delta = VarInt::decode(buf)?.into_inner();
let object_id_val = match self.prev_object_id {
None => delta,
Some(prev) => prev
.checked_add(1)
.and_then(|v| v.checked_add(delta))
.ok_or(CodecError::InvalidField)?,
};
self.prev_object_id = Some(object_id_val);
let object_id =
VarInt::from_u64(object_id_val).map_err(|_| CodecError::InvalidField)?.into_inner();
let extension_headers_len = if self.extensions_present {
let ext_len = VarInt::decode(buf)?.into_inner();
skip(buf, ext_len)?;
ext_len
} else {
0
};
let payload_length = VarInt::decode(buf)?.into_inner();
let status = if payload_length == 0 {
let code = VarInt::decode(buf)?.into_inner();
decoded_status(code)?;
Some(code)
} else {
skip(buf, payload_length)?;
None
};
Ok(SubgroupObjectMeta {
object_id,
extension_headers_len,
payload_length,
status,
wire_len: (start - buf.remaining()) as u64,
})
}
pub fn write_object(
&mut self,
object: &SubgroupObject,
buf: &mut impl BufMut,
) -> Result<(), CodecError> {
let declared = object.payload_length.into_inner();
if declared != object.payload.len() as u64 {
return Err(CodecError::InvalidField);
}
let oid = object.object_id.into_inner();
let delta = match self.prev_object_id {
None => oid,
Some(prev) => oid
.checked_sub(prev)
.and_then(|v| v.checked_sub(1))
.ok_or(CodecError::InvalidField)?,
};
VarInt::from_u64(delta).map_err(|_| CodecError::InvalidField)?.encode(buf);
if self.extensions_present {
let ext_len = object.extension_headers.len();
VarInt::from_usize(ext_len).encode(buf);
buf.put_slice(&object.extension_headers);
}
object.payload_length.encode(buf);
if object.payload_length.into_inner() == 0 {
let status = object.object_status.unwrap_or(ObjectStatus::Normal);
VarInt::from_usize(status.as_u64() as usize).encode(buf);
} else {
buf.put_slice(&object.payload);
}
self.prev_object_id = Some(oid);
Ok(())
}
}
fn datagram_type_is_assigned(ty: u64) -> bool {
let Ok(ty) = u8::try_from(ty) else {
return false;
};
ty & 0xD0 == 0 && ty & 0x22 != 0x22
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DatagramHeader {
pub datagram_type: u8,
pub track_alias: VarInt,
pub group_id: VarInt,
pub object_id: VarInt,
pub publisher_priority: Option<u8>,
pub extension_headers: Vec<u8>,
pub object_status: Option<ObjectStatus>,
}
impl DatagramHeader {
pub fn has_object_id(&self) -> bool {
self.datagram_type & 0x04 == 0
}
pub fn is_end_of_group(&self) -> bool {
self.datagram_type & 0x02 != 0
}
pub fn is_status(&self) -> bool {
self.datagram_type & 0x20 != 0
}
pub fn has_extensions(&self) -> bool {
self.datagram_type & 0x01 != 0
}
pub fn has_default_priority(&self) -> bool {
self.datagram_type & 0x08 != 0
}
pub fn encode_checked(&self, buf: &mut impl BufMut) -> Result<(), CodecError> {
if !self.is_status() && matches!(self.object_status, Some(s) if s != ObjectStatus::Normal) {
return Err(CodecError::InvalidField);
}
if !datagram_type_is_assigned(self.datagram_type as u64) {
return Err(CodecError::UnknownDatagramType(self.datagram_type as u64));
}
if self.has_extensions() {
if self.extension_headers.is_empty() {
return Err(CodecError::InvalidField);
}
} else if !self.extension_headers.is_empty() {
return Err(CodecError::InvalidField);
}
check_extensions_against_status_on_encode(
self.effective_status().map(|s| s.as_u64()),
if self.has_extensions() { self.extension_headers.len() as u64 } else { 0 },
)?;
self.encode(buf);
Ok(())
}
fn effective_status(&self) -> Option<ObjectStatus> {
if self.is_status() {
Some(self.object_status.unwrap_or(ObjectStatus::Normal))
} else {
None
}
}
pub fn extensions_permitted(&self) -> bool {
if !self.has_extensions() || self.extension_headers.is_empty() {
return true;
}
self.effective_status().unwrap_or(ObjectStatus::Normal) == ObjectStatus::Normal
}
pub fn encode(&self, buf: &mut impl BufMut) {
VarInt::from_usize(self.datagram_type as usize).encode(buf);
self.track_alias.encode(buf);
self.group_id.encode(buf);
if self.has_object_id() {
self.object_id.encode(buf);
}
if !self.has_default_priority() {
buf.put_u8(self.publisher_priority.unwrap_or(0));
}
if self.has_extensions() {
VarInt::from_usize(self.extension_headers.len()).encode(buf);
buf.put_slice(&self.extension_headers);
}
if self.is_status() {
let status = self.object_status.unwrap_or(ObjectStatus::Normal);
VarInt::from_usize(status.as_u64() as usize).encode(buf);
}
}
pub fn decode(buf: &mut impl Buf) -> Result<Self, CodecError> {
let type_value = VarInt::decode(buf)?.into_inner();
if !datagram_type_is_assigned(type_value) {
return Err(CodecError::UnknownDatagramType(type_value));
}
let datagram_type = type_value as u8;
let track_alias = VarInt::decode(buf)?;
let group_id = VarInt::decode(buf)?;
let object_id =
if datagram_type & 0x04 == 0 { VarInt::decode(buf)? } else { VarInt::from_usize(0) };
let publisher_priority = if datagram_type & 0x08 == 0 {
if buf.remaining() < 1 {
return Err(CodecError::UnexpectedEnd);
}
Some(buf.get_u8())
} else {
None
};
let extension_headers = if datagram_type & 0x01 != 0 {
let ext_len = VarInt::decode(buf)?.into_inner() as usize;
if ext_len == 0 {
return Err(CodecError::InvalidField);
}
crate::types::read_bytes(buf, ext_len)?
} else {
Vec::new()
};
let object_status = if datagram_type & 0x20 != 0 {
Some(decoded_status(VarInt::decode(buf)?.into_inner())?)
} else {
None
};
Ok(Self {
datagram_type,
track_alias,
group_id,
object_id,
publisher_priority,
extension_headers,
object_status,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FetchHeader {
pub request_id: VarInt,
}
const FETCH_STREAM_TYPE: u64 = 0x05;
fn stream_type_error(raw: u64) -> CodecError {
if raw == FETCH_STREAM_TYPE || subgroup_type_is_assigned(raw) {
CodecError::InvalidField
} else {
CodecError::UnknownStreamType(raw)
}
}
impl FetchHeader {
pub fn encode(&self, buf: &mut impl BufMut) {
VarInt::from_usize(FETCH_STREAM_TYPE as usize).encode(buf);
self.request_id.encode(buf);
}
pub fn decode(buf: &mut impl Buf) -> Result<Self, CodecError> {
let stream_type = VarInt::decode(buf)?.into_inner();
if stream_type != FETCH_STREAM_TYPE {
return Err(stream_type_error(stream_type));
}
let request_id = VarInt::decode(buf)?;
Ok(Self { request_id })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SubgroupIdEncoding {
Zero,
SameAsPrior,
PriorPlusOne,
Present,
}
impl SubgroupIdEncoding {
pub fn from_flags(flags: u8) -> Self {
match flags & 0x03 {
0x00 => SubgroupIdEncoding::Zero,
0x01 => SubgroupIdEncoding::SameAsPrior,
0x02 => SubgroupIdEncoding::PriorPlusOne,
_ => SubgroupIdEncoding::Present,
}
}
pub fn as_bits(self) -> u8 {
match self {
SubgroupIdEncoding::Zero => 0x00,
SubgroupIdEncoding::SameAsPrior => 0x01,
SubgroupIdEncoding::PriorPlusOne => 0x02,
SubgroupIdEncoding::Present => 0x03,
}
}
pub fn references_prior(self) -> bool {
matches!(self, SubgroupIdEncoding::SameAsPrior | SubgroupIdEncoding::PriorPlusOne)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FetchObjectHeader {
pub serialization_flags: u8,
pub group_id: VarInt,
pub subgroup_id: VarInt,
pub object_id: VarInt,
pub publisher_priority: u8,
pub extension_headers: Vec<u8>,
pub payload_length: VarInt,
pub object_status: Option<ObjectStatus>,
}
impl FetchObjectHeader {
pub fn subgroup_id_encoding(&self) -> SubgroupIdEncoding {
SubgroupIdEncoding::from_flags(self.serialization_flags)
}
pub fn has_object_id(&self) -> bool {
self.serialization_flags & 0x04 != 0
}
pub fn has_group_id(&self) -> bool {
self.serialization_flags & 0x08 != 0
}
pub fn has_priority(&self) -> bool {
self.serialization_flags & 0x10 != 0
}
pub fn has_extensions(&self) -> bool {
self.serialization_flags & 0x20 != 0
}
pub fn references_prior_object(&self) -> bool {
self.subgroup_id_encoding().references_prior()
|| !self.has_object_id()
|| !self.has_group_id()
|| !self.has_priority()
}
pub fn status(&self) -> ObjectStatus {
if self.payload_length.into_inner() == 0 {
self.object_status.unwrap_or(ObjectStatus::Normal)
} else {
ObjectStatus::Normal
}
}
pub fn extensions_permitted(&self) -> bool {
self.extension_headers.is_empty() || self.status() == ObjectStatus::Normal
}
pub fn encode(&self, buf: &mut impl BufMut) -> Result<(), CodecError> {
if self.serialization_flags & 0xc0 != 0 {
return Err(CodecError::InvalidField);
}
buf.put_u8(self.serialization_flags);
if self.has_group_id() {
self.group_id.encode(buf);
}
if self.subgroup_id_encoding() == SubgroupIdEncoding::Present {
self.subgroup_id.encode(buf);
}
if self.has_object_id() {
self.object_id.encode(buf);
}
if self.has_priority() {
buf.put_u8(self.publisher_priority);
}
if self.has_extensions() {
VarInt::from_usize(self.extension_headers.len()).encode(buf);
buf.put_slice(&self.extension_headers);
}
self.payload_length.encode(buf);
if self.payload_length.into_inner() == 0 {
VarInt::from_u64(self.object_status.unwrap_or(ObjectStatus::Normal).as_u64())
.map_err(|_| CodecError::InvalidField)?
.encode(buf);
}
Ok(())
}
}
#[derive(Debug, Clone, Default)]
pub struct FetchObjectWriter {
prior: Option<PriorFetchObject>,
}
impl FetchObjectWriter {
pub fn new() -> Self {
Self::default()
}
pub fn header_for(
&self,
original: &FetchObjectHeader,
) -> Result<FetchObjectHeader, CodecError> {
if original.serialization_flags & 0xc0 != 0 {
return Err(CodecError::InvalidField);
}
let group_id = original.group_id.into_inner();
let subgroup_id = original.subgroup_id.into_inner();
let object_id = original.object_id.into_inner();
let mut flags = 0u8;
if original.has_group_id() || self.prior.map(|p| p.group_id) != Some(group_id) {
flags |= 0x08;
}
if original.has_object_id()
|| self.prior.and_then(|p| p.object_id.checked_add(1)) != Some(object_id)
{
flags |= 0x04;
}
if original.has_priority()
|| self.prior.map(|p| p.publisher_priority) != Some(original.publisher_priority)
{
flags |= 0x10;
}
if original.has_extensions() {
flags |= 0x20;
}
flags |= self.subgroup_bits(original, subgroup_id).as_bits();
Ok(FetchObjectHeader {
serialization_flags: flags,
group_id: original.group_id,
subgroup_id: original.subgroup_id,
object_id: original.object_id,
publisher_priority: original.publisher_priority,
extension_headers: original.extension_headers.clone(),
payload_length: original.payload_length,
object_status: original.object_status,
})
}
fn subgroup_bits(&self, original: &FetchObjectHeader, subgroup_id: u64) -> SubgroupIdEncoding {
let inherits = self.prior.map(|p| p.subgroup_id) == Some(subgroup_id);
let successor =
self.prior.is_some_and(|p| p.subgroup_id.checked_add(1) == Some(subgroup_id));
let kept = match original.subgroup_id_encoding() {
SubgroupIdEncoding::Zero if subgroup_id == 0 => Some(SubgroupIdEncoding::Zero),
SubgroupIdEncoding::SameAsPrior if inherits => Some(SubgroupIdEncoding::SameAsPrior),
SubgroupIdEncoding::PriorPlusOne if successor => Some(SubgroupIdEncoding::PriorPlusOne),
SubgroupIdEncoding::Present => Some(SubgroupIdEncoding::Present),
_ => None,
};
match kept {
Some(encoding) => encoding,
None if subgroup_id == 0 => SubgroupIdEncoding::Zero,
None if inherits => SubgroupIdEncoding::SameAsPrior,
None if successor => SubgroupIdEncoding::PriorPlusOne,
None => SubgroupIdEncoding::Present,
}
}
pub fn write_object_header(
&mut self,
original: &FetchObjectHeader,
out: &mut impl BufMut,
) -> Result<FetchObjectHeader, CodecError> {
let header = self.header_for(original)?;
header.encode(out)?;
self.advance(&header);
Ok(header)
}
pub fn advance(&mut self, written: &FetchObjectHeader) {
self.prior = Some(PriorFetchObject {
group_id: written.group_id.into_inner(),
subgroup_id: written.subgroup_id.into_inner(),
object_id: written.object_id.into_inner(),
publisher_priority: written.publisher_priority,
});
}
}
#[derive(Debug, Clone, Copy)]
struct PriorFetchObject {
group_id: u64,
subgroup_id: u64,
object_id: u64,
publisher_priority: u8,
}
#[derive(Debug, Clone, Default)]
pub struct FetchObjectReader {
prior: Option<PriorFetchObject>,
}
impl FetchObjectReader {
pub fn new() -> Self {
Self::default()
}
pub fn read_object_header(
&mut self,
buf: &mut impl Buf,
) -> Result<FetchObjectHeader, CodecError> {
if buf.remaining() < 1 {
return Err(CodecError::UnexpectedEnd);
}
let serialization_flags = buf.get_u8();
if serialization_flags & 0xc0 != 0 {
return Err(CodecError::InvalidField);
}
let subgroup_encoding = SubgroupIdEncoding::from_flags(serialization_flags);
let has_object_id = serialization_flags & 0x04 != 0;
let has_group_id = serialization_flags & 0x08 != 0;
let has_priority = serialization_flags & 0x10 != 0;
let has_extensions = serialization_flags & 0x20 != 0;
let inherits = subgroup_encoding.references_prior()
|| !has_object_id
|| !has_group_id
|| !has_priority;
if inherits && self.prior.is_none() {
return Err(CodecError::InvalidField);
}
let prior = self.prior;
let group_id = if has_group_id {
VarInt::decode(buf)?
} else {
let prior = prior.ok_or(CodecError::InvalidField)?;
VarInt::from_u64(prior.group_id).map_err(|_| CodecError::InvalidField)?
};
let subgroup_id = match subgroup_encoding {
SubgroupIdEncoding::Zero => VarInt::from_usize(0),
SubgroupIdEncoding::SameAsPrior => {
let prior = prior.ok_or(CodecError::InvalidField)?;
VarInt::from_u64(prior.subgroup_id).map_err(|_| CodecError::InvalidField)?
}
SubgroupIdEncoding::PriorPlusOne => {
let prior = prior.ok_or(CodecError::InvalidField)?;
let next = prior.subgroup_id.checked_add(1).ok_or(CodecError::InvalidField)?;
VarInt::from_u64(next).map_err(|_| CodecError::InvalidField)?
}
SubgroupIdEncoding::Present => VarInt::decode(buf)?,
};
let object_id = if has_object_id {
VarInt::decode(buf)?
} else {
let prior = prior.ok_or(CodecError::InvalidField)?;
let next = prior.object_id.checked_add(1).ok_or(CodecError::InvalidField)?;
VarInt::from_u64(next).map_err(|_| CodecError::InvalidField)?
};
let publisher_priority = if has_priority {
if buf.remaining() < 1 {
return Err(CodecError::UnexpectedEnd);
}
buf.get_u8()
} else {
prior.ok_or(CodecError::InvalidField)?.publisher_priority
};
let extension_headers = if has_extensions {
let ext_len = VarInt::decode(buf)?.into_inner() as usize;
crate::types::read_bytes(buf, ext_len)?
} else {
Vec::new()
};
let payload_length = VarInt::decode(buf)?;
let object_status = if payload_length.into_inner() == 0 {
Some(decoded_status(VarInt::decode(buf)?.into_inner())?)
} else {
None
};
self.prior = Some(PriorFetchObject {
group_id: group_id.into_inner(),
subgroup_id: subgroup_id.into_inner(),
object_id: object_id.into_inner(),
publisher_priority,
});
Ok(FetchObjectHeader {
serialization_flags,
group_id,
subgroup_id,
object_id,
publisher_priority,
extension_headers,
payload_length,
object_status,
})
}
pub fn write_object_header(
&mut self,
header: &FetchObjectHeader,
buf: &mut impl BufMut,
) -> Result<(), CodecError> {
if header.serialization_flags & 0xc0 != 0 {
return Err(CodecError::InvalidField);
}
if header.payload_length.into_inner() != 0
&& matches!(header.object_status, Some(s) if s != ObjectStatus::Normal)
{
return Err(CodecError::InvalidField);
}
if !header.has_extensions() && !header.extension_headers.is_empty() {
return Err(CodecError::InvalidField);
}
if header.references_prior_object() && self.prior.is_none() {
return Err(CodecError::InvalidField);
}
let prior = self.prior;
if !header.has_group_id() {
let prior = prior.ok_or(CodecError::InvalidField)?;
if header.group_id.into_inner() != prior.group_id {
return Err(CodecError::InvalidField);
}
}
let subgroup_id = header.subgroup_id.into_inner();
match header.subgroup_id_encoding() {
SubgroupIdEncoding::Zero => {
if subgroup_id != 0 {
return Err(CodecError::InvalidField);
}
}
SubgroupIdEncoding::SameAsPrior => {
let prior = prior.ok_or(CodecError::InvalidField)?;
if subgroup_id != prior.subgroup_id {
return Err(CodecError::InvalidField);
}
}
SubgroupIdEncoding::PriorPlusOne => {
let prior = prior.ok_or(CodecError::InvalidField)?;
let next = prior.subgroup_id.checked_add(1).ok_or(CodecError::InvalidField)?;
if subgroup_id != next {
return Err(CodecError::InvalidField);
}
}
SubgroupIdEncoding::Present => {}
}
if !header.has_object_id() {
let prior = prior.ok_or(CodecError::InvalidField)?;
let next = prior.object_id.checked_add(1).ok_or(CodecError::InvalidField)?;
if header.object_id.into_inner() != next {
return Err(CodecError::InvalidField);
}
}
if !header.has_priority() {
let prior = prior.ok_or(CodecError::InvalidField)?;
if header.publisher_priority != prior.publisher_priority {
return Err(CodecError::InvalidField);
}
}
buf.put_u8(header.serialization_flags);
if header.has_group_id() {
header.group_id.encode(buf);
}
if header.subgroup_id_encoding() == SubgroupIdEncoding::Present {
header.subgroup_id.encode(buf);
}
if header.has_object_id() {
header.object_id.encode(buf);
}
if header.has_priority() {
buf.put_u8(header.publisher_priority);
}
if header.has_extensions() {
VarInt::from_usize(header.extension_headers.len()).encode(buf);
buf.put_slice(&header.extension_headers);
}
header.payload_length.encode(buf);
if header.payload_length.into_inner() == 0 {
let status = header.object_status.unwrap_or(ObjectStatus::Normal);
VarInt::from_usize(status.as_u64() as usize).encode(buf);
}
self.prior = Some(PriorFetchObject {
group_id: header.group_id.into_inner(),
subgroup_id,
object_id: header.object_id.into_inner(),
publisher_priority: header.publisher_priority,
});
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
const VECTORS: &[&str] = &[
"100100800004deadbeef",
"100100800004deadbeef0002cafe",
"3001000004deadbeef",
"11010080000004deadbeef",
"1101008000023c0104deadbeef",
"100105800004deadbeef000003",
"10010a800004deadbeef000004",
"11010080000004deadbeef000002cafe",
"1101008000023c0204deadbeef00023c0302cafe",
"1101008000023c010003",
"120105800004deadbeef",
];
fn vi(v: u64) -> VarInt {
VarInt::from_u64(v).unwrap()
}
fn hex(s: &str) -> Vec<u8> {
(0..s.len()).step_by(2).map(|i| u8::from_str_radix(&s[i..i + 2], 16).unwrap()).collect()
}
fn decode_all(bytes: &[u8]) -> (SubgroupHeader, Vec<SubgroupObject>) {
let mut cursor = bytes;
let header = SubgroupHeader::decode(&mut cursor)
.unwrap_or_else(|e| panic!("header decode failed: {e:?}"));
let mut reader = SubgroupObjectReader::new(&header);
let mut objects = Vec::new();
while cursor.has_remaining() {
objects.push(
reader
.read_object(&mut cursor)
.unwrap_or_else(|e| panic!("object {} decode failed: {e:?}", objects.len())),
);
}
(header, objects)
}
fn encode_all(header: &SubgroupHeader, objects: &[SubgroupObject]) -> Vec<u8> {
let mut buf = Vec::new();
header.encode(&mut buf);
let mut writer = SubgroupObjectReader::new(header);
for o in objects {
writer.write_object(o, &mut buf).unwrap_or_else(|e| panic!("write failed: {e:?}"));
}
buf
}
fn object(id: u64, extensions: Vec<u8>, payload: Vec<u8>) -> SubgroupObject {
SubgroupObject {
object_id: vi(id),
extension_headers: extensions,
payload_length: vi(payload.len() as u64),
object_status: None,
payload,
}
}
#[test]
fn two_objects_with_extensions_have_distinct_ids() {
let bytes = hex("11010080000004deadbeef000002cafe");
let (header, objects) = decode_all(&bytes);
assert!(header.has_extensions());
assert_eq!(objects.len(), 2);
assert_eq!(objects[0].object_id.into_inner(), 0);
assert_eq!(objects[1].object_id.into_inner(), 1);
assert_eq!(objects[0].payload, hex("deadbeef"));
assert_eq!(objects[1].payload, hex("cafe"));
assert!(objects.iter().all(|o| o.extension_headers.is_empty()));
}
#[test]
fn deltas_resolve_sparse_ids() {
let header = SubgroupHeader::decode(&mut &hex("100100800004deadbeef")[..]).unwrap();
let objects: Vec<_> =
[3u64, 4, 40].iter().map(|&id| object(id, vec![], vec![0xAA, id as u8])).collect();
let (_, decoded) = decode_all(&encode_all(&header, &objects));
let ids: Vec<u64> = decoded.iter().map(|o| o.object_id.into_inner()).collect();
assert_eq!(ids, vec![3, 4, 40]);
}
#[test]
fn write_rejects_non_increasing_ids() {
let header = SubgroupHeader::decode(&mut &hex("100100800004deadbeef")[..]).unwrap();
let mut writer = SubgroupObjectReader::new(&header);
let mut buf = Vec::new();
writer.write_object(&object(7, vec![], vec![0x01]), &mut buf).unwrap();
for id in [7u64, 6, 0] {
let err = writer.write_object(&object(id, vec![], vec![0x01]), &mut buf).unwrap_err();
assert!(matches!(err, CodecError::InvalidField), "id {id} gave {err:?}");
}
}
#[test]
fn eliding_an_object_renumbers_its_successor() {
let header = SubgroupHeader::decode(&mut &hex("100100800004deadbeef")[..]).unwrap();
let all: Vec<_> = (0..5u64).map(|id| object(id, vec![], vec![id as u8])).collect();
for elided in 0..5u64 {
let kept: Vec<_> =
all.iter().filter(|o| o.object_id.into_inner() != elided).cloned().collect();
let (_, decoded) = decode_all(&encode_all(&header, &kept));
let ids: Vec<u64> = decoded.iter().map(|o| o.object_id.into_inner()).collect();
let expected: Vec<u64> = (0..5u64).filter(|&i| i != elided).collect();
assert_eq!(ids, expected, "eliding object {elided}");
}
}
#[test]
fn extensions_blob_excludes_its_length_prefix() {
let bytes = hex("1101008000023c0204deadbeef00023c0302cafe");
let (_, objects) = decode_all(&bytes);
assert_eq!(objects.len(), 2);
assert_eq!(objects[0].object_id.into_inner(), 0);
assert_eq!(objects[1].object_id.into_inner(), 1);
assert_eq!(objects[0].extension_headers, hex("3c02"));
assert_eq!(objects[1].extension_headers, hex("3c03"));
assert_eq!(objects[0].payload, hex("deadbeef"));
assert_eq!(objects[1].payload, hex("cafe"));
}
#[test]
fn status_object_carries_its_extensions_block() {
let (_, objects) = decode_all(&hex("1101008000023c010003"));
assert_eq!(objects.len(), 1);
assert_eq!(objects[0].extension_headers, hex("3c01"));
assert_eq!(objects[0].payload_length.into_inner(), 0);
assert_eq!(objects[0].object_status.map(ObjectStatus::as_u64), Some(3));
assert!(objects[0].payload.is_empty());
}
fn plain_header() -> SubgroupHeader {
SubgroupHeader::decode(&mut &hex("100100800004deadbeef")[..]).unwrap()
}
fn subgroup_status_body(code: u64) -> Vec<u8> {
let mut buf = vec![0x00, 0x00];
VarInt::from_u64(code).unwrap().encode(&mut buf);
buf
}
fn status_object(status: Option<ObjectStatus>) -> SubgroupObject {
SubgroupObject {
object_id: vi(0),
extension_headers: vec![],
payload_length: vi(0),
object_status: status,
payload: vec![],
}
}
fn datagram_status_bytes(code: u64) -> Vec<u8> {
let mut buf = vec![0x20, 0x01, 0x02, 0x03, 0x80];
VarInt::from_u64(code).unwrap().encode(&mut buf);
buf
}
fn status_datagram(status: Option<ObjectStatus>) -> DatagramHeader {
DatagramHeader {
datagram_type: 0x20,
track_alias: vi(1),
group_id: vi(2),
object_id: vi(3),
publisher_priority: Some(0x80),
extension_headers: vec![],
object_status: status,
}
}
#[test]
fn assigned_statuses_round_trip() {
let header = plain_header();
for &status in ObjectStatus::ALL {
let mut bytes = Vec::new();
SubgroupObjectReader::new(&header)
.write_object(&status_object(Some(status)), &mut bytes)
.unwrap();
assert_eq!(
bytes,
subgroup_status_body(status.as_u64()),
"{status:?} on the subgroup wire"
);
let decoded = SubgroupObjectReader::new(&header)
.read_object(&mut &bytes[..])
.unwrap_or_else(|e| panic!("{status:?} was written and then refused: {e:?}"));
assert_eq!(decoded.object_status, Some(status), "{status:?} through read_object");
let meta = SubgroupObjectReader::new(&header)
.read_object_meta(&mut &bytes[..])
.unwrap_or_else(|e| panic!("{status:?} was written and then refused: {e:?}"));
assert_eq!(meta.status, Some(status.as_u64()), "{status:?} through read_object_meta");
let datagram = status_datagram(Some(status));
let mut bytes = Vec::new();
datagram.encode(&mut bytes);
assert_eq!(
bytes,
datagram_status_bytes(status.as_u64()),
"{status:?} on the datagram wire"
);
let decoded = DatagramHeader::decode(&mut &bytes[..]).unwrap_or_else(|e| {
panic!("{status:?} datagram was written and then refused: {e:?}")
});
assert_eq!(decoded, datagram, "{status:?} datagram round trip");
}
}
#[test]
fn a_status_datagram_without_a_status_encodes_normal() {
let mut bytes = Vec::new();
status_datagram(None).encode(&mut bytes);
assert_eq!(bytes, datagram_status_bytes(ObjectStatus::Normal.as_u64()));
let decoded = DatagramHeader::decode(&mut &bytes[..])
.unwrap_or_else(|e| panic!("own output refused: {e:?}"));
assert_eq!(decoded.object_status, Some(ObjectStatus::Normal));
}
#[test]
fn the_wire_accepts_exactly_what_the_type_can_hold() {
let header = plain_header();
for code in 0x00u64..=0x3f {
let assigned = ObjectStatus::ALL.iter().any(|s| s.as_u64() == code);
let body = subgroup_status_body(code);
let read = SubgroupObjectReader::new(&header).read_object(&mut &body[..]);
assert_eq!(
read.is_ok(),
assigned,
"subgroup read_object on status {code:#x}: {read:?}"
);
let meta = SubgroupObjectReader::new(&header).read_object_meta(&mut &body[..]);
assert_eq!(
meta.is_ok(),
assigned,
"subgroup read_object_meta on status {code:#x}: {meta:?}"
);
let datagram = DatagramHeader::decode(&mut &datagram_status_bytes(code)[..]);
assert_eq!(
datagram.is_ok(),
assigned,
"status datagram on status {code:#x}: {datagram:?}"
);
if assigned {
let status = ObjectStatus::from_u64(code).unwrap();
let mut bytes = Vec::new();
SubgroupObjectReader::new(&header)
.write_object(&status_object(Some(status)), &mut bytes)
.unwrap();
assert_eq!(bytes, body, "write_object on status {code:#x}");
let mut bytes = Vec::new();
status_datagram(Some(status)).encode(&mut bytes);
assert_eq!(
bytes,
datagram_status_bytes(code),
"datagram encode on status {code:#x}"
);
}
}
}
#[test]
fn vectors_re_encode_byte_identically() {
for vector in VECTORS {
let bytes = hex(vector);
let (header, objects) = decode_all(&bytes);
assert_eq!(encode_all(&header, &objects), bytes, "[{vector}] re-encode");
}
}
#[test]
fn meta_matches_read_object() {
for vector in VECTORS {
let bytes = hex(vector);
let mut cursor = &bytes[..];
let header = SubgroupHeader::decode(&mut cursor).unwrap();
let mut full_reader = SubgroupObjectReader::new(&header);
let mut meta_reader = SubgroupObjectReader::new(&header);
let mut full_cursor = cursor;
let mut meta_cursor = cursor;
while meta_cursor.has_remaining() {
let before = meta_cursor.remaining();
let object = full_reader.read_object(&mut full_cursor).unwrap();
let meta = meta_reader.read_object_meta(&mut meta_cursor).unwrap();
assert_eq!(meta.object_id, object.object_id.into_inner(), "[{vector}]");
assert_eq!(
meta.extension_headers_len,
object.extension_headers.len() as u64,
"[{vector}]"
);
assert_eq!(meta.payload_length, object.payload_length.into_inner(), "[{vector}]");
assert_eq!(
meta.status,
object.object_status.map(ObjectStatus::as_u64),
"[{vector}]"
);
assert_eq!(meta.wire_len, (before - meta_cursor.remaining()) as u64, "[{vector}]");
assert_eq!(full_cursor.remaining(), meta_cursor.remaining(), "[{vector}]");
}
}
}
#[test]
fn fetch_objects_report_extensions_beside_a_non_normal_status() {
let cases: [(&str, bool, &str); 4] = [
(
"3c000000023c010003",
false,
"a two-byte extension block beside End of Group is the violation",
),
("1c0000000003", true, "End of Group with no extensions is fine"),
(
"3c000000000003",
true,
"a present but zero-length block carries nothing, so nothing is beside the status",
),
("3c000000023c0104", true, "a normal object may carry extensions"),
];
for (vector, permitted, why) in cases {
let bytes = hex(vector);
let mut cursor = &bytes[..];
let header = FetchObjectReader::new()
.read_object_header(&mut cursor)
.unwrap_or_else(|e| panic!("[{vector}] {why}: decode failed with {e:?}"));
assert_eq!(
header.extensions_permitted(),
permitted,
"[{vector}] {why}: expected {permitted}, got {}",
header.extensions_permitted(),
);
}
}
#[test]
fn fetch_object_status_is_normal_whenever_a_payload_is_declared() {
let bytes = hex("3c000000023c0104");
let mut cursor = &bytes[..];
let header = FetchObjectReader::new().read_object_header(&mut cursor).unwrap();
assert_eq!(header.object_status, None, "no status field follows a non-zero length");
assert_eq!(header.status(), ObjectStatus::Normal);
let contradictory =
FetchObjectHeader { object_status: Some(ObjectStatus::EndOfGroup), ..header };
assert_eq!(
contradictory.status(),
ObjectStatus::Normal,
"a declared payload wins over a status the wire could not have carried",
);
assert!(contradictory.extensions_permitted());
}
#[test]
fn short_buffers_report_unexpected_end() {
let bytes = hex("1101008000023c0204deadbeef00023c0302cafe");
let mut cursor = &bytes[..];
let header = SubgroupHeader::decode(&mut cursor).unwrap();
let objects_start = bytes.len() - cursor.len();
for cut in objects_start..bytes.len() {
let mut reader = SubgroupObjectReader::new(&header);
let mut meta_reader = SubgroupObjectReader::new(&header);
let mut cursor = &bytes[objects_start..cut];
let mut meta_cursor = cursor;
while cursor.has_remaining() {
if let Err(err) = reader.read_object(&mut cursor) {
assert!(
matches!(err, CodecError::UnexpectedEnd | CodecError::VarInt(_)),
"cut {cut} gave {err:?}"
);
break;
}
}
while meta_cursor.has_remaining() {
if let Err(err) = meta_reader.read_object_meta(&mut meta_cursor) {
assert!(
matches!(err, CodecError::UnexpectedEnd | CodecError::VarInt(_)),
"cut {cut} gave {err:?}"
);
break;
}
}
}
}
}