use super::constant::ControlMessageType;
use super::control_message::ControlMessageTrait;
use crate::model::common::tuple::Tuple;
use crate::model::common::varint::{BufMutVarIntExt, BufVarIntExt};
use crate::model::error::ParseError;
use crate::model::parameter::message_parameter::{
MessageParameter, deserialize_message_parameters,
};
use bytes::{BufMut, Bytes, BytesMut};
#[derive(Debug, PartialEq, Clone)]
pub struct SubscribeNamespace {
pub request_id: u64,
pub track_namespace_prefix: Tuple,
pub subscribe_options: u64,
pub parameters: Vec<MessageParameter>,
}
impl SubscribeNamespace {
pub fn new(
request_id: u64,
track_namespace_prefix: Tuple,
subscribe_options: u64,
parameters: Vec<MessageParameter>,
) -> Self {
Self {
request_id,
track_namespace_prefix,
subscribe_options,
parameters,
}
}
}
impl ControlMessageTrait for SubscribeNamespace {
fn serialize(&self) -> Result<Bytes, ParseError> {
let mut buf = BytesMut::new();
buf.put_vi(ControlMessageType::SubscribeNamespace)?;
let mut payload = BytesMut::new();
payload.put_vi(self.request_id)?;
payload.extend_from_slice(&self.track_namespace_prefix.serialize()?);
payload.put_vi(self.subscribe_options)?;
payload.put_vi(self.parameters.len())?;
for param in &self.parameters {
payload.extend_from_slice(¶m.serialize()?);
}
let payload_len: u16 = payload
.len()
.try_into()
.map_err(|e: std::num::TryFromIntError| ParseError::CastingError {
context: "Announce::serialize",
from_type: "usize",
to_type: "u16",
details: e.to_string(),
})?;
buf.put_u16(payload_len);
buf.extend_from_slice(&payload);
Ok(buf.freeze())
}
fn parse_payload(payload: &mut Bytes) -> Result<Box<Self>, ParseError> {
let request_id = payload.get_vi()?;
let track_namespace_prefix = Tuple::deserialize(payload)?;
if track_namespace_prefix.fields.len() > 32 {
return Err(ParseError::ProtocolViolation {
context: "SubscribeNamespace::parse_payload",
details: format!(
"Track namespace prefix has {} fields, maximum is 32",
track_namespace_prefix.fields.len()
),
});
}
let subscribe_options = payload.get_vi()?;
let param_count = payload.get_vi()?;
let parameters =
deserialize_message_parameters(payload, param_count, ControlMessageType::SubscribeNamespace)?;
Ok(Box::new(SubscribeNamespace {
request_id,
track_namespace_prefix,
subscribe_options,
parameters,
}))
}
fn get_type(&self) -> ControlMessageType {
ControlMessageType::SubscribeNamespace
}
}
#[cfg(test)]
mod tests {
use crate::model::parameter::authorization_token::AuthorizationToken;
use super::*;
use bytes::Buf;
#[test]
fn test_roundtrip() {
let request_id = 241421;
let track_namespace_prefix = Tuple::from_utf8_path("pre/fix/me");
let subscribe_options = 0x02u64; let parameters = vec![
MessageParameter::new_authorization_token(AuthorizationToken::new_use_value(
0,
Bytes::from_static(b"test-token"),
)),
MessageParameter::new_forward(true),
];
let subscribe_namespace = SubscribeNamespace {
request_id,
track_namespace_prefix,
subscribe_options,
parameters,
};
let mut buf = subscribe_namespace.serialize().unwrap();
let msg_type = buf.get_vi().unwrap();
assert_eq!(msg_type, ControlMessageType::SubscribeNamespace as u64);
let msg_length = buf.get_u16();
assert_eq!(msg_length as usize, buf.remaining());
let deserialized = SubscribeNamespace::parse_payload(&mut buf).unwrap();
assert_eq!(*deserialized, subscribe_namespace);
assert!(!buf.has_remaining());
}
#[test]
fn test_excess_roundtrip() {
let request_id = 241421;
let track_namespace_prefix = Tuple::from_utf8_path("pre/fix/me");
let subscribe_options = 0x01u64; let parameters = vec![
MessageParameter::new_authorization_token(AuthorizationToken::new_use_value(
0,
Bytes::from_static(b"test-token"),
)),
MessageParameter::new_forward(true),
];
let subscribe_namespace = SubscribeNamespace {
request_id,
track_namespace_prefix,
subscribe_options,
parameters,
};
let serialized = subscribe_namespace.serialize().unwrap();
let mut excess = BytesMut::new();
excess.extend_from_slice(&serialized);
excess.extend_from_slice(&[9u8, 1u8, 1u8]);
let mut buf = excess.freeze();
let msg_type = buf.get_vi().unwrap();
assert_eq!(msg_type, ControlMessageType::SubscribeNamespace as u64);
let msg_length = buf.get_u16();
assert_eq!(msg_length as usize, buf.remaining() - 3);
let deserialized = SubscribeNamespace::parse_payload(&mut buf).unwrap();
assert_eq!(*deserialized, subscribe_namespace);
assert_eq!(buf.chunk(), &[9u8, 1u8, 1u8]);
}
#[test]
fn test_partial_message() {
let request_id = 241421;
let track_namespace_prefix = Tuple::from_utf8_path("pre/fix/me");
let subscribe_options = 0x00u64; let parameters = vec![
MessageParameter::new_authorization_token(AuthorizationToken::new_use_value(
0,
Bytes::from_static(b"test-token"),
)),
MessageParameter::new_forward(true),
];
let subscribe_namespace = SubscribeNamespace {
request_id,
track_namespace_prefix,
subscribe_options,
parameters,
};
let mut buf = subscribe_namespace.serialize().unwrap();
let msg_type = buf.get_vi().unwrap();
assert_eq!(msg_type, ControlMessageType::SubscribeNamespace as u64);
let msg_length = buf.get_u16();
assert_eq!(msg_length as usize, buf.remaining());
let upper = buf.remaining() / 2;
let mut partial = buf.slice(..upper);
let deserialized = SubscribeNamespace::parse_payload(&mut partial);
assert!(deserialized.is_err());
}
}