use bytes::{Buf, Bytes, BytesMut};
use tracing::trace;
use crate::model::common::varint::{BufMutVarIntExt, BufVarIntExt};
use crate::model::error::ParseError;
use crate::model::extension_header::object_extension::{
ObjectExtension, deserialize_object_extensions, serialize_object_extensions,
};
use super::constant::ObjectStatus;
#[derive(Debug, Clone, PartialEq)]
pub struct SubgroupObject {
pub object_id: u64,
pub extension_headers: Option<Vec<ObjectExtension>>,
pub object_status: Option<ObjectStatus>,
pub payload: Option<Bytes>,
}
impl SubgroupObject {
pub fn serialize(
&self,
previous_object_id: Option<u64>,
has_extensions: bool,
) -> Result<Bytes, ParseError> {
let mut buf = BytesMut::new();
let object_id_delta = if let Some(id) = previous_object_id {
self.object_id - id - 1
} else {
self.object_id
};
trace!(
"SubgroupObject::serialize || object_id_delta: {} prev: {:?} object_id: {} ext_headers: {:?}",
object_id_delta, previous_object_id, self.object_id, &self.extension_headers
);
buf.put_vi(object_id_delta)?;
let extension_headers = self.extension_headers.as_deref().unwrap_or(&[]);
if !has_extensions && !extension_headers.is_empty() {
return Err(ParseError::ProtocolViolation {
context: "SubgroupObject::serialize(extension_headers)",
details: "extension headers are present but header type has no extensions".to_string(),
});
}
if has_extensions {
let ext_buf = serialize_object_extensions(extension_headers)?;
buf.put_vi(ext_buf.len())?;
buf.extend_from_slice(&ext_buf);
}
if let Some(payload) = &self.payload
&& !payload.is_empty()
{
buf.put_vi(payload.len())?;
buf.extend_from_slice(payload);
} else {
buf.put_vi(0u64)?;
buf.put_vi(self.object_status.unwrap_or(ObjectStatus::Normal))?;
}
Ok(buf.freeze())
}
pub fn deserialize(
bytes: &mut Bytes,
previous_object_id: &Option<u64>,
has_extensions: bool,
) -> Result<Self, ParseError> {
let object_id_delta = bytes.get_vi()?;
let object_id = if let Some(id) = previous_object_id {
id + object_id_delta + 1
} else {
object_id_delta
};
trace!(
"SubgroupObject::deserialize || object_id_delta: {} prev: {:?} object_id: {}",
object_id_delta, previous_object_id, object_id
);
let extension_headers = if has_extensions {
let ext_len = bytes.get_vi()?;
if ext_len > 0 {
let ext_len: usize =
ext_len
.try_into()
.map_err(|e: std::num::TryFromIntError| ParseError::CastingError {
context: "SubgroupObject::deserialize",
from_type: "u64",
to_type: "usize",
details: e.to_string(),
})?;
if bytes.remaining() < ext_len {
return Err(ParseError::NotEnoughBytes {
context: "SubgroupObject::deserialize",
needed: ext_len,
available: bytes.remaining(),
});
}
let mut header_bytes = bytes.copy_to_bytes(ext_len);
let headers = deserialize_object_extensions(&mut header_bytes).map_err(|_| {
ParseError::ProtocolViolation {
context: "SubgroupObject::deserialize",
details: "Failed to parse extension header".to_string(),
}
})?;
Some(headers)
} else {
Some(vec![])
}
} else {
None
};
let payload_len = bytes.get_vi()?;
let (object_status, payload) = if payload_len == 0 {
let status_raw = bytes.get_vi()?;
let status = ObjectStatus::try_from(status_raw)?;
(Some(status), None)
} else {
let payload_len: usize = payload_len
.try_into()
.map_err(|e: std::num::TryFromIntError| ParseError::CastingError {
context: "SubgroupObject::deserialize",
from_type: "u64",
to_type: "usize",
details: e.to_string(),
})?;
if bytes.remaining() < payload_len {
return Err(ParseError::NotEnoughBytes {
context: "SubgroupObject::deserialize",
needed: payload_len,
available: bytes.remaining(),
});
}
let payload_data = bytes.copy_to_bytes(payload_len);
(None, Some(payload_data))
};
Ok(SubgroupObject {
object_id,
extension_headers,
object_status,
payload,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::common::pair::KeyValuePair;
use bytes::Buf;
#[test]
fn test_roundtrip() {
let object_id: u64 = 10;
let prev_object_id = 9;
let extension_headers = Some(vec![
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_varint(0, 10).unwrap(),
},
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_bytes(1, Bytes::from_static(b"wololoo")).unwrap(),
},
]);
let object_status = None;
let payload = Some(Bytes::from_static(b"01239gjawkk92837aldmi"));
let subgroup_object = SubgroupObject {
object_id,
extension_headers,
payload,
object_status,
};
let mut buf = subgroup_object
.serialize(Some(prev_object_id), true)
.unwrap();
let deserialized = SubgroupObject::deserialize(&mut buf, &Some(prev_object_id), true).unwrap();
assert_eq!(deserialized, subgroup_object);
assert!(!buf.has_remaining());
}
#[test]
fn test_serializes_empty_extension_headers_as_absent() {
let object_id: u64 = 10;
let payload = Some(Bytes::from_static(&[0xab]));
let with_no_headers = SubgroupObject {
object_id,
extension_headers: None,
object_status: None,
payload: payload.clone(),
}
.serialize(None, false)
.unwrap();
let with_empty_headers = SubgroupObject {
object_id,
extension_headers: Some(vec![]),
object_status: None,
payload,
}
.serialize(None, false)
.unwrap();
assert_eq!(with_empty_headers, with_no_headers);
}
#[test]
fn test_excess_roundtrip() {
let object_id: u64 = 10;
let prev_object_id = 9;
let extension_headers = Some(vec![
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_varint(0, 10).unwrap(),
},
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_bytes(1, Bytes::from_static(b"wololoo")).unwrap(),
},
]);
let object_status = None;
let payload = Some(Bytes::from_static(b"01239gjawkk92837aldmi"));
let subgroup_object = SubgroupObject {
object_id,
extension_headers,
payload,
object_status,
};
let serialized = subgroup_object
.serialize(Some(prev_object_id), true)
.unwrap();
let mut excess = BytesMut::new();
excess.extend_from_slice(&serialized);
excess.extend_from_slice(&[9u8, 1u8, 1u8]);
let mut buf = excess.freeze();
let deserialized = SubgroupObject::deserialize(&mut buf, &Some(prev_object_id), true).unwrap();
assert_eq!(deserialized, subgroup_object);
assert_eq!(buf.chunk(), &[9u8, 1u8, 1u8]);
}
#[test]
fn test_partial_message() {
let object_id: u64 = 10;
let prev_object_id = 9;
let extension_headers = Some(vec![
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_varint(0, 10).unwrap(),
},
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_bytes(1, Bytes::from_static(b"wololoo")).unwrap(),
},
]);
let object_status = None;
let payload = Some(Bytes::from_static(b"01239gjawkk92837aldmi"));
let subgroup_object = SubgroupObject {
object_id,
extension_headers,
payload,
object_status,
};
let buf = subgroup_object
.serialize(Some(prev_object_id), true)
.unwrap();
let upper = buf.remaining() / 2;
let mut partial = buf.slice(..upper);
let deserialized = SubgroupObject::deserialize(&mut partial, &Some(prev_object_id), true);
assert!(deserialized.is_err());
}
}