use std::io;
#[allow(unused_imports)]
use log::{debug, error, info, trace, warn};
use speedy::{Context, Error, Readable, Writable, Writer};
use enumflags2::BitFlags;
use bytes::Bytes;
use crate::{
messages::submessages::{elements::parameter_list::ParameterList, submessages::*},
structure::{
guid::EntityId,
sequence_number::{FragmentNumber, SequenceNumber},
},
};
#[derive(Debug, PartialEq, Eq, Clone)]
#[cfg_attr(test, derive(Default))]
pub struct DataFrag {
pub reader_id: EntityId,
pub writer_id: EntityId,
pub writer_sn: SequenceNumber,
pub fragment_starting_num: FragmentNumber,
pub fragments_in_submessage: u16,
pub data_size: u32,
pub fragment_size: u16,
pub inline_qos: Option<ParameterList>,
pub serialized_payload: Bytes,
}
impl DataFrag {
pub fn len_serialized(&self) -> usize {
2 + 2 + 4 + 4 + 8 + 4 + 2 + 2 + 4 + self.inline_qos.as_ref().map(|q| q.len_serialized() ).unwrap_or(0) + self.serialized_payload.len()
}
pub fn total_number_of_fragments(&self) -> FragmentNumber {
let frag_size = self.fragment_size as u32;
if frag_size < 1 {
FragmentNumber::INVALID
} else {
FragmentNumber::new(self.data_size.div_ceil(frag_size))
}
}
pub fn deserialize(buffer: &Bytes, flags: BitFlags<DATAFRAG_Flags>) -> io::Result<Self> {
let mut cursor = io::Cursor::new(&buffer);
let endianness = endianness_flag(flags.bits());
let map_speedy_err = |p: Error| io::Error::other(p);
let _extra_flags =
u16::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor).map_err(map_speedy_err)?;
let octets_to_inline_qos =
u16::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor).map_err(map_speedy_err)?;
let reader_id = EntityId::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor)
.map_err(map_speedy_err)?;
let writer_id = EntityId::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor)
.map_err(map_speedy_err)?;
let writer_sn = SequenceNumber::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor)
.map_err(map_speedy_err)?;
let fragment_starting_num =
FragmentNumber::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor)
.map_err(map_speedy_err)?;
let fragments_in_submessage =
u16::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor).map_err(map_speedy_err)?;
let fragment_size =
u16::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor).map_err(map_speedy_err)?;
let data_size =
u32::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor).map_err(map_speedy_err)?;
let expect_qos = flags.contains(DATAFRAG_Flags::InlineQos);
let rtps_v25_header_size: u16 = 28;
if octets_to_inline_qos < rtps_v25_header_size {
return Err(io::Error::other(format!(
"DataFrag has too low octetsToInlineQos = {octets_to_inline_qos}"
)));
}
if octets_to_inline_qos > rtps_v25_header_size {
let extra_octets = octets_to_inline_qos - rtps_v25_header_size;
cursor.set_position(cursor.position() + u64::from(extra_octets));
if cursor.position() > buffer.len().try_into().unwrap() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"DATAFRAG submessage octets_to_inline_qos points to byte {}, but message len={}.",
cursor.position(),
buffer.len()
),
));
}
}
let inline_qos = if expect_qos {
Some(
ParameterList::read_from_stream_unbuffered_with_ctx(endianness, &mut cursor)
.map_err(map_speedy_err)?,
)
} else {
None
};
if writer_sn < SequenceNumber::new(1) {
return Err(io::Error::other(
"DataFrag SequenceNumber < 1. Discarding as invalid.",
));
}
if fragment_size < 1 || (fragment_size as u32) > data_size {
return Err(io::Error::other(format!(
"Invalid DataFrag. fragment_size={fragment_size} data_size={data_size} Expected 1 <= \
fragment_size <= data_size."
)));
}
let serialized_payload = buffer.clone().split_off(cursor.position() as usize);
let datafrag = Self {
reader_id,
writer_id,
writer_sn,
fragment_starting_num,
fragments_in_submessage,
data_size,
fragment_size,
inline_qos,
serialized_payload,
};
let expected_total = datafrag.total_number_of_fragments();
if fragment_starting_num < FragmentNumber::new(1) || fragment_starting_num > expected_total {
return Err(io::Error::other(format!(
"DataFrag fragmentStartingNum={fragment_starting_num:?} \
expected_total={expected_total:?}. Expected 1 <= fragmentStartingNum <= expected_total. \
Discarding as invalid."
)));
}
if fragments_in_submessage < 1 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("DataFrag fragmentsInSubmessage={fragments_in_submessage} must be >= 1."),
));
}
let start_u32 = u32::from(fragment_starting_num);
let count_u32 = u32::from(fragments_in_submessage);
let total_u32 = u32::from(expected_total);
let last_frag = start_u32
.checked_add(count_u32)
.and_then(|n| n.checked_sub(1));
match last_frag {
None => {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"DataFrag fragment span overflow: fragmentStartingNum={start_u32} \
fragmentsInSubmessage={count_u32}."
),
));
}
Some(last) if last > total_u32 => {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"DataFrag fragment span exceeds total: last fragment {last} > expected_total \
{total_u32} (fragmentStartingNum={start_u32}, fragmentsInSubmessage={count_u32})."
),
));
}
_ => {}
}
let max_payload_bytes = (count_u32 as usize) * (fragment_size as usize);
if datafrag.serialized_payload.len() > max_payload_bytes {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"DataFrag serializedData length={} exceeds fragmentsInSubmessage={} x fragmentSize={}.",
datafrag.serialized_payload.len(),
fragments_in_submessage,
fragment_size
),
));
}
Ok(datafrag)
}
}
impl<C: Context> Writable<C> for DataFrag {
fn write_to<T: ?Sized + Writer<C>>(&self, writer: &mut T) -> Result<(), C::Error> {
writer.write_u16(0)?; writer.write_u16(28)?; writer.write_value(&self.reader_id)?;
writer.write_value(&self.writer_id)?;
writer.write_value(&self.writer_sn)?;
writer.write_value(&self.fragment_starting_num)?;
writer.write_value(&self.fragments_in_submessage)?;
writer.write_value(&self.fragment_size)?;
writer.write_value(&self.data_size)?;
if self.inline_qos.is_some() && !self.inline_qos.as_ref().unwrap().parameters.is_empty() {
writer.write_value(&self.inline_qos)?;
}
writer.write_bytes(&self.serialized_payload)?;
Ok(())
}
}
impl HasEntityIds for DataFrag {
fn receiver_entity_id(&self) -> EntityId {
self.reader_id
}
fn sender_entity_id(&self) -> EntityId {
self.writer_id
}
}
#[cfg(test)]
mod tests {
use bytes::Bytes;
use enumflags2::BitFlags;
use speedy::{Endianness, Writable};
use super::*;
use crate::messages::submessages::submessages::DATAFRAG_Flags;
fn serialize_body(df: &DataFrag) -> Bytes {
let mut buf = Vec::new();
df.write_to_stream_with_ctx(Endianness::LittleEndian, &mut buf)
.expect("serialize DataFrag body");
Bytes::from(buf)
}
#[test]
fn deserialize_accepts_valid_fragment_span() {
let df = DataFrag {
writer_sn: SequenceNumber::new(1),
fragment_starting_num: FragmentNumber::new(1),
fragments_in_submessage: 1,
fragment_size: 256,
data_size: 512,
serialized_payload: Bytes::from(vec![0u8; 256]),
..Default::default()
};
let expected_total = df.total_number_of_fragments();
assert_eq!(expected_total, FragmentNumber::new(2));
let start_u32 = u32::from(df.fragment_starting_num);
let count_u32 = u32::from(df.fragments_in_submessage);
let last = start_u32 + count_u32 - 1;
assert!(last <= u32::from(expected_total));
assert!(df.serialized_payload.len() <= (count_u32 as usize) * (df.fragment_size as usize));
}
#[test]
fn deserialize_rejects_fragment_span_beyond_total() {
let df = DataFrag {
fragment_starting_num: FragmentNumber::new(2),
fragments_in_submessage: 2,
fragment_size: 256,
data_size: 512,
serialized_payload: Bytes::from(vec![0u8; 256]),
..Default::default()
};
let bytes = serialize_body(&df);
assert!(DataFrag::deserialize(&bytes, BitFlags::<DATAFRAG_Flags>::empty()).is_err());
}
#[test]
fn deserialize_rejects_oversized_payload_for_fragment_run() {
let df = DataFrag {
writer_sn: SequenceNumber::new(1),
fragment_starting_num: FragmentNumber::new(1),
fragments_in_submessage: 1,
fragment_size: 256,
data_size: 512,
serialized_payload: Bytes::from(vec![0u8; 512]),
..Default::default()
};
let bytes = serialize_body(&df);
assert!(DataFrag::deserialize(&bytes, BitFlags::<DATAFRAG_Flags>::empty()).is_err());
}
#[test]
fn deserialize_rejects_zero_fragments_in_submessage() {
let df = DataFrag {
fragment_starting_num: FragmentNumber::new(1),
fragments_in_submessage: 0,
fragment_size: 256,
data_size: 512,
serialized_payload: Bytes::from(vec![0u8; 256]),
..Default::default()
};
let bytes = serialize_body(&df);
assert!(DataFrag::deserialize(&bytes, BitFlags::<DATAFRAG_Flags>::empty()).is_err());
}
}