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 extensions_permitted_at(status: Option<u64>, extensions_len: u64) -> bool {
match status {
None => true,
Some(code) => extensions_len == 0 || code == ObjectStatus::Normal.as_u64(),
}
}
const SUBGROUP_BASE_BIT: u8 = 0x10;
const SUBGROUP_FORM_FORBIDDEN_BITS: u8 = 0xC0;
const SUBGROUP_ID_MODE_MASK: u8 = 0x06;
const SUBGROUP_ID_MODE_RESERVED: u8 = 0b11;
fn validate_subgroup_type(raw: u64) -> Result<(), CodecError> {
if subgroup_type_is_valid(raw) {
Ok(())
} else {
Err(stream_type_error(raw))
}
}
const FETCH_STREAM_TYPE: u64 = 0x05;
fn wide_type_refusal(
buf: &mut impl Buf,
refusal: fn(u64) -> CodecError,
) -> Result<Option<CodecError>, CodecError> {
if !buf.has_remaining() {
return Err(CodecError::UnexpectedEnd);
}
if buf.chunk()[0] < 0x40 {
return Ok(None);
}
let raw = VarInt::decode(buf)?.into_inner();
Ok(Some(refusal(raw)))
}
fn datagram_type_refusal(raw: u64) -> CodecError {
match validate_datagram_type(raw) {
Ok(()) => CodecError::InvalidField,
Err(e) => e,
}
}
fn subgroup_type_is_valid(raw: u64) -> bool {
raw <= 0xFF && {
let t = raw as u8;
t & SUBGROUP_FORM_FORBIDDEN_BITS == 0
&& t & SUBGROUP_BASE_BIT != 0
&& (t & SUBGROUP_ID_MODE_MASK) >> 1 != SUBGROUP_ID_MODE_RESERVED
}
}
fn subgroup_type_is_reserved_mode(raw: u64) -> bool {
raw <= 0xFF && {
let t = raw as u8;
t & SUBGROUP_FORM_FORBIDDEN_BITS == 0
&& t & SUBGROUP_BASE_BIT != 0
&& (t & SUBGROUP_ID_MODE_MASK) >> 1 == SUBGROUP_ID_MODE_RESERVED
}
}
fn stream_type_error(raw: u64) -> CodecError {
if raw == FETCH_STREAM_TYPE || subgroup_type_is_valid(raw) {
CodecError::InvalidField
} else if subgroup_type_is_reserved_mode(raw) {
CodecError::InvalidTypeValue {
raw,
detail: "its SUBGROUP_ID_MODE is 0b11, which this draft reserves",
}
} else {
CodecError::UnknownStreamType(raw)
}
}
const DATAGRAM_END_OF_GROUP_BIT: u8 = 0x02;
const DATAGRAM_STATUS_BIT: u8 = 0x20;
const DATAGRAM_FORM_FORBIDDEN_BITS: u8 = 0xD0;
fn validate_datagram_type(raw: u64) -> Result<(), CodecError> {
if datagram_type_is_valid(raw) {
Ok(())
} else if datagram_type_is_status_end_of_group(raw) {
Err(CodecError::InvalidTypeValue {
raw,
detail: "it sets both the STATUS bit and the END_OF_GROUP bit",
})
} else {
Err(CodecError::UnknownDatagramType(raw))
}
}
fn datagram_type_is_valid(raw: u64) -> bool {
raw <= 0xFF && {
let t = raw as u8;
t & DATAGRAM_FORM_FORBIDDEN_BITS == 0
&& !(t & DATAGRAM_STATUS_BIT != 0 && t & DATAGRAM_END_OF_GROUP_BIT != 0)
}
}
fn datagram_type_is_status_end_of_group(raw: u64) -> bool {
raw <= 0xFF && {
let t = raw as u8;
t & DATAGRAM_FORM_FORBIDDEN_BITS == 0
&& t & DATAGRAM_STATUS_BIT != 0
&& t & DATAGRAM_END_OF_GROUP_BIT != 0
}
}
#[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>,
}
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 & SUBGROUP_ID_MODE_MASK) >> 1 == 1
}
pub fn has_explicit_subgroup_id(&self) -> bool {
(self.header_type & SUBGROUP_ID_MODE_MASK) >> 1 == 2
}
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 let Some(p) = self.publisher_priority {
buf.put_u8(p);
}
}
pub fn encode_checked(&self, buf: &mut impl BufMut) -> Result<(), CodecError> {
validate_subgroup_type(self.header_type as u64)?;
self.encode(buf);
Ok(())
}
pub fn decode(buf: &mut impl Buf) -> Result<Self, CodecError> {
if let Some(err) = wide_type_refusal(buf, stream_type_error)? {
return Err(err);
}
let header_type = buf.get_u8();
validate_subgroup_type(header_type as u64)?;
let track_alias = VarInt::decode(buf)?;
let group_id = VarInt::decode(buf)?;
let subgroup_id = if (header_type & SUBGROUP_ID_MODE_MASK) >> 1 == 2 {
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 {
extensions_permitted_at(
self.object_status.map(|s| s.as_u64()),
self.extension_headers.len() as u64,
)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PayloadPermission {
Permitted,
Forbidden,
}
impl PayloadPermission {
pub fn permits(self) -> bool {
matches!(self, PayloadPermission::Permitted)
}
}
#[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) => match ObjectStatus::from_u64(code)? {
ObjectStatus::Normal => Some(PayloadPermission::Permitted),
ObjectStatus::EndOfGroup | ObjectStatus::EndOfTrack => {
Some(PayloadPermission::Forbidden)
}
},
}
}
pub fn extensions_permitted(&self) -> bool {
extensions_permitted_at(self.status, self.extension_headers_len)
}
}
#[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(())
}
}
#[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_extensions(&self) -> bool {
self.datagram_type & 0x01 != 0
}
pub fn is_end_of_group(&self) -> bool {
self.datagram_type & 0x02 != 0
}
pub fn has_object_id(&self) -> bool {
self.datagram_type & 0x04 == 0
}
pub fn has_default_priority(&self) -> bool {
self.datagram_type & 0x08 != 0
}
pub fn is_status(&self) -> bool {
self.datagram_type & 0x20 != 0
}
pub fn extensions_permitted(&self) -> bool {
extensions_permitted_at(
self.object_status.map(|s| s.as_u64()),
if self.has_extensions() { self.extension_headers.len() as u64 } else { 0 },
)
}
pub fn encode_checked(&self, buf: &mut impl BufMut) -> Result<(), CodecError> {
validate_datagram_type(self.datagram_type as u64)?;
if !self.is_status() && matches!(self.object_status, Some(s) if s != ObjectStatus::Normal) {
return Err(CodecError::InvalidField);
}
if self.has_extensions() && self.extension_headers.is_empty() {
return Err(CodecError::InvalidField);
}
if !self.extensions_permitted() {
return Err(CodecError::InvalidField);
}
self.encode(buf);
Ok(())
}
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(128));
}
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> {
if let Some(err) = wide_type_refusal(buf, datagram_type_refusal)? {
return Err(err);
}
let datagram_type = buf.get_u8();
validate_datagram_type(datagram_type as u64)?;
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 {
None
} else {
if buf.remaining() < 1 {
return Err(CodecError::UnexpectedEnd);
}
Some(buf.get_u8())
};
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,
}
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)]
#[repr(u64)]
pub enum FetchEndOfRange {
NonExistent = 0x8c,
Unknown = 0x10c,
}
impl FetchEndOfRange {
pub const ALL: &[FetchEndOfRange] = &[FetchEndOfRange::NonExistent, FetchEndOfRange::Unknown];
pub fn from_u64(v: u64) -> Option<Self> {
match v {
0x8c => Some(FetchEndOfRange::NonExistent),
0x10c => Some(FetchEndOfRange::Unknown),
_ => None,
}
}
pub fn as_u64(self) -> u64 {
self as u64
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FetchSubgroupMode {
Zero,
SameAsPrior,
PriorPlusOne,
Present,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FetchObjectHeader {
pub serialization_flags: VarInt,
pub group_id: Option<VarInt>,
pub subgroup_id: Option<VarInt>,
pub object_id: Option<VarInt>,
pub publisher_priority: Option<u8>,
pub extensions: Option<Vec<u8>>,
pub payload_length: VarInt,
}
impl FetchObjectHeader {
fn flags(&self) -> u64 {
self.serialization_flags.into_inner()
}
pub fn end_of_range(&self) -> Option<FetchEndOfRange> {
FetchEndOfRange::from_u64(self.flags())
}
pub fn is_datagram(&self) -> bool {
self.end_of_range().is_none() && self.flags() & 0x40 != 0
}
pub fn subgroup_mode(&self) -> FetchSubgroupMode {
if self.end_of_range().is_some() || self.is_datagram() {
return FetchSubgroupMode::Zero;
}
match self.flags() & 0x03 {
0x00 => FetchSubgroupMode::Zero,
0x01 => FetchSubgroupMode::SameAsPrior,
0x02 => FetchSubgroupMode::PriorPlusOne,
_ => FetchSubgroupMode::Present,
}
}
pub fn has_group_id(&self) -> bool {
self.end_of_range().is_some() || self.flags() & 0x08 != 0
}
pub fn has_subgroup_id(&self) -> bool {
matches!(self.subgroup_mode(), FetchSubgroupMode::Present)
}
pub fn has_object_id(&self) -> bool {
self.end_of_range().is_some() || self.flags() & 0x04 != 0
}
pub fn has_priority(&self) -> bool {
self.end_of_range().is_none() && self.flags() & 0x10 != 0
}
pub fn has_extensions(&self) -> bool {
self.end_of_range().is_none() && self.flags() & 0x20 != 0
}
pub fn references_prior_object(&self) -> bool {
if self.end_of_range().is_some() {
return false;
}
!self.has_group_id()
|| !self.has_object_id()
|| !self.has_priority()
|| matches!(
self.subgroup_mode(),
FetchSubgroupMode::SameAsPrior | FetchSubgroupMode::PriorPlusOne
)
}
pub fn encode(&self, buf: &mut impl BufMut) -> Result<(), CodecError> {
let flags = self.flags();
if flags >= 128 && FetchEndOfRange::from_u64(flags).is_none() {
return Err(CodecError::InvalidField);
}
if self.group_id.is_some() != self.has_group_id()
|| self.subgroup_id.is_some() != self.has_subgroup_id()
|| self.object_id.is_some() != self.has_object_id()
|| self.publisher_priority.is_some() != self.has_priority()
|| self.extensions.is_some() != self.has_extensions()
{
return Err(CodecError::InvalidField);
}
self.serialization_flags.encode(buf);
if let Some(group_id) = self.group_id {
group_id.encode(buf);
}
if let Some(subgroup_id) = self.subgroup_id {
subgroup_id.encode(buf);
}
if let Some(object_id) = self.object_id {
object_id.encode(buf);
}
if let Some(priority) = self.publisher_priority {
buf.put_u8(priority);
}
if let Some(extensions) = &self.extensions {
VarInt::from_usize(extensions.len()).encode(buf);
buf.put_slice(extensions);
}
self.payload_length.encode(buf);
Ok(())
}
pub fn decode(buf: &mut impl Buf) -> Result<Self, CodecError> {
let serialization_flags = VarInt::decode(buf)?;
let flags = serialization_flags.into_inner();
if flags >= 128 && FetchEndOfRange::from_u64(flags).is_none() {
return Err(CodecError::InvalidField);
}
let probe = Self {
serialization_flags,
group_id: None,
subgroup_id: None,
object_id: None,
publisher_priority: None,
extensions: None,
payload_length: VarInt::from_usize(0),
};
let group_id = if probe.has_group_id() { Some(VarInt::decode(buf)?) } else { None };
let subgroup_id = if probe.has_subgroup_id() { Some(VarInt::decode(buf)?) } else { None };
let object_id = if probe.has_object_id() { Some(VarInt::decode(buf)?) } else { None };
let publisher_priority = if probe.has_priority() {
if buf.remaining() < 1 {
return Err(CodecError::UnexpectedEnd);
}
Some(buf.get_u8())
} else {
None
};
let extensions = if probe.has_extensions() {
let ext_len = VarInt::decode(buf)?.into_inner() as usize;
Some(crate::types::read_bytes(buf, ext_len)?)
} else {
None
};
let payload_length = VarInt::decode(buf)?;
Ok(Self {
serialization_flags,
group_id,
subgroup_id,
object_id,
publisher_priority,
extensions,
payload_length,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FetchObjectLocation {
pub group_id: u64,
pub subgroup_id: Option<u64>,
pub object_id: u64,
pub publisher_priority: Option<u8>,
pub end_of_range: Option<FetchEndOfRange>,
}
#[derive(Debug, Clone, Default)]
pub struct FetchObjectReader {
group_id: Option<u64>,
subgroup_id: Option<u64>,
object_id: Option<u64>,
publisher_priority: Option<u8>,
}
impl FetchObjectReader {
pub fn new() -> Self {
Self::default()
}
pub fn resolve(
&mut self,
header: &FetchObjectHeader,
) -> Result<FetchObjectLocation, CodecError> {
let group_id = match header.group_id {
Some(v) => v.into_inner(),
None => self.group_id.ok_or(CodecError::InvalidField)?,
};
let subgroup_id = if header.is_datagram() {
None
} else {
match header.subgroup_mode() {
FetchSubgroupMode::Zero => Some(0),
FetchSubgroupMode::SameAsPrior => {
Some(self.subgroup_id.ok_or(CodecError::InvalidField)?)
}
FetchSubgroupMode::PriorPlusOne => {
let prior = self.subgroup_id.ok_or(CodecError::InvalidField)?;
let next = prior.checked_add(1).ok_or(CodecError::InvalidField)?;
VarInt::from_u64(next).map_err(|_| CodecError::InvalidField)?;
Some(next)
}
FetchSubgroupMode::Present => {
Some(header.subgroup_id.ok_or(CodecError::InvalidField)?.into_inner())
}
}
};
let object_id = match header.object_id {
Some(v) => v.into_inner(),
None => {
let prior = self.object_id.ok_or(CodecError::InvalidField)?;
let next = prior.checked_add(1).ok_or(CodecError::InvalidField)?;
VarInt::from_u64(next).map_err(|_| CodecError::InvalidField)?;
next
}
};
if let Some(priority) = header.publisher_priority {
self.publisher_priority = Some(priority);
}
self.group_id = Some(group_id);
self.subgroup_id = subgroup_id;
self.object_id = Some(object_id);
Ok(FetchObjectLocation {
group_id,
subgroup_id,
object_id,
publisher_priority: self.publisher_priority,
end_of_range: header.end_of_range(),
})
}
}
#[derive(Debug, Clone, Default)]
pub struct FetchObjectWriter {
group_id: Option<u64>,
subgroup_id: Option<u64>,
object_id: Option<u64>,
publisher_priority: Option<u8>,
}
impl FetchObjectWriter {
pub fn new() -> Self {
Self::default()
}
pub fn header_for(
&self,
original: &FetchObjectHeader,
location: &FetchObjectLocation,
) -> Result<FetchObjectHeader, CodecError> {
if original.end_of_range().is_some() {
return Ok(FetchObjectHeader {
serialization_flags: original.serialization_flags,
group_id: Some(VarInt::from_u64(location.group_id)?),
subgroup_id: None,
object_id: Some(VarInt::from_u64(location.object_id)?),
publisher_priority: None,
extensions: None,
payload_length: original.payload_length,
});
}
let group_id = if !original.has_group_id() && self.group_id == Some(location.group_id) {
None
} else {
Some(VarInt::from_u64(location.group_id)?)
};
let object_id = if !original.has_object_id()
&& self.object_id.and_then(|p| p.checked_add(1)) == Some(location.object_id)
{
None
} else {
Some(VarInt::from_u64(location.object_id)?)
};
let (subgroup_mode, subgroup_id) = self.subgroup_field(original, location)?;
let publisher_priority = self.priority_field(original, location);
let flags = original.flags();
let mut new_flags = subgroup_mode;
if flags & 0x40 != 0 {
new_flags |= 0x40;
}
if group_id.is_some() {
new_flags |= 0x08;
}
if object_id.is_some() {
new_flags |= 0x04;
}
if publisher_priority.is_some() {
new_flags |= 0x10;
}
if original.has_extensions() {
new_flags |= 0x20;
}
Ok(FetchObjectHeader {
serialization_flags: VarInt::from_u64(new_flags)?,
group_id,
subgroup_id,
object_id,
publisher_priority,
extensions: original.extensions.clone(),
payload_length: original.payload_length,
})
}
fn subgroup_field(
&self,
original: &FetchObjectHeader,
location: &FetchObjectLocation,
) -> Result<(u64, Option<VarInt>), CodecError> {
if original.is_datagram() {
return Ok((original.flags() & 0x03, None));
}
let subgroup_id = location.subgroup_id.ok_or(CodecError::InvalidField)?;
let inherits = self.subgroup_id == Some(subgroup_id);
let successor = self.subgroup_id.is_some_and(|p| p.checked_add(1) == Some(subgroup_id));
let kept = match original.subgroup_mode() {
FetchSubgroupMode::Zero if subgroup_id == 0 => Some((0x00, None)),
FetchSubgroupMode::SameAsPrior if inherits => Some((0x01, None)),
FetchSubgroupMode::PriorPlusOne if successor => Some((0x02, None)),
FetchSubgroupMode::Present => Some((0x03, Some(subgroup_id))),
_ => None,
};
let (mode, explicit) = match kept {
Some(pair) => pair,
None if subgroup_id == 0 => (0x00, None),
None if inherits => (0x01, None),
None if successor => (0x02, None),
None => (0x03, Some(subgroup_id)),
};
Ok((mode, explicit.map(VarInt::from_u64).transpose()?))
}
fn priority_field(
&self,
original: &FetchObjectHeader,
location: &FetchObjectLocation,
) -> Option<u8> {
if original.has_priority() || self.publisher_priority != location.publisher_priority {
return location.publisher_priority;
}
None
}
pub fn write_object_header(
&mut self,
original: &FetchObjectHeader,
location: &FetchObjectLocation,
out: &mut impl BufMut,
) -> Result<FetchObjectHeader, CodecError> {
let header = self.header_for(original, location)?;
header.encode(out)?;
self.advance(&header, location);
Ok(header)
}
pub fn advance(&mut self, written: &FetchObjectHeader, location: &FetchObjectLocation) {
if let Some(priority) = written.publisher_priority {
self.publisher_priority = Some(priority);
}
self.group_id = Some(location.group_id);
self.subgroup_id = location.subgroup_id;
self.object_id = Some(location.object_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
const VECTORS: &[&str] = &[
"100100800004deadbeef",
"100100800004deadbeef0002cafe",
"3001000004deadbeef",
"11010080000004deadbeef",
"1101008000023c0104deadbeef",
"100105800004deadbeef000003",
"10010a800004deadbeef000004",
"11010080000004deadbeef000002cafe",
"1101008000023c0204deadbeef00023c0302cafe",
"1101008000023c010003",
"180105800004deadbeef",
"120103800504deadbeef",
];
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 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;
}
}
}
}
}