use crate::model::common::varint::BufVarIntExt;
use crate::model::error::ParseError;
use bytes::{Buf, Bytes};
use super::{
client_setup::ClientSetup, constant::ControlMessageType, fetch::Fetch, fetch_cancel::FetchCancel,
fetch_ok::FetchOk, goaway::GoAway, max_request_id::MaxRequestId, namespace::Namespace,
namespace_done::NamespaceDone, publish::Publish, publish_done::PublishDone,
publish_namespace::PublishNamespace, publish_namespace_cancel::PublishNamespaceCancel,
publish_namespace_done::PublishNamespaceDone, publish_ok::PublishOk, request_error::RequestError,
request_ok::RequestOk, request_update::RequestUpdate, requests_blocked::RequestsBlocked,
server_setup::ServerSetup, subscribe::Subscribe, subscribe_namespace::SubscribeNamespace,
subscribe_ok::SubscribeOk, switch::Switch, track_status::TrackStatus, unsubscribe::Unsubscribe,
unsubscribe_namespace::UnsubscribeNamespace,
};
#[derive(Debug, Clone, PartialEq)]
pub enum ControlMessage {
Namespace(Box<Namespace>),
NamespaceDone(Box<NamespaceDone>),
PublishNamespace(Box<PublishNamespace>),
PublishNamespaceCancel(Box<PublishNamespaceCancel>),
RequestOk(Box<RequestOk>),
Publish(Box<Publish>),
PublishOk(Box<PublishOk>),
PublishDone(Box<PublishDone>),
ClientSetup(Box<ClientSetup>),
Fetch(Box<Fetch>),
FetchCancel(Box<FetchCancel>),
FetchOk(Box<FetchOk>),
Goaway(Box<GoAway>),
MaxRequestId(Box<MaxRequestId>),
ServerSetup(Box<ServerSetup>),
Subscribe(Box<Subscribe>),
SubscribeOk(Box<SubscribeOk>),
RequestUpdate(Box<RequestUpdate>),
RequestsBlocked(Box<RequestsBlocked>),
TrackStatus(Box<TrackStatus>),
PublishNamespaceDone(Box<PublishNamespaceDone>),
Unsubscribe(Box<Unsubscribe>),
SubscribeNamespace(Box<SubscribeNamespace>),
RequestError(Box<RequestError>),
UnsubscribeNamespace(Box<UnsubscribeNamespace>),
Switch(Box<Switch>),
}
pub trait ControlMessageTrait: std::fmt::Debug {
fn serialize(&self) -> Result<Bytes, ParseError>;
fn parse_payload(payload: &mut Bytes) -> Result<Box<Self>, ParseError>
where
Self: Sized;
fn get_type(&self) -> ControlMessageType;
}
impl ControlMessage {
pub fn deserialize(bytes: &mut Bytes) -> Result<Self, ParseError> {
let message_type = bytes.get_vi()?;
let msg_type = ControlMessageType::try_from(message_type)?;
if bytes.remaining() < 2 {
return Err(ParseError::NotEnoughBytes {
context: "ControlMessage::deserialize(payload_length)",
needed: 2,
available: 0,
});
}
let payload_length = bytes.get_u16() as usize;
if bytes.remaining() < payload_length {
return Err(ParseError::NotEnoughBytes {
context: "ControlMessage::deserialize(payload_length)",
needed: payload_length,
available: bytes.remaining(),
});
}
let mut payload = bytes.copy_to_bytes(payload_length);
let message = match msg_type {
ControlMessageType::Namespace => {
Namespace::parse_payload(&mut payload).map(ControlMessage::Namespace)
}
ControlMessageType::NamespaceDone => {
NamespaceDone::parse_payload(&mut payload).map(ControlMessage::NamespaceDone)
}
ControlMessageType::PublishNamespace => {
PublishNamespace::parse_payload(&mut payload).map(ControlMessage::PublishNamespace)
}
ControlMessageType::PublishNamespaceCancel => {
PublishNamespaceCancel::parse_payload(&mut payload)
.map(ControlMessage::PublishNamespaceCancel)
}
ControlMessageType::PublishNamespaceDone => {
PublishNamespaceDone::parse_payload(&mut payload).map(ControlMessage::PublishNamespaceDone)
}
ControlMessageType::RequestError => {
RequestError::parse_payload(&mut payload).map(ControlMessage::RequestError)
}
ControlMessageType::RequestOk => {
RequestOk::parse_payload(&mut payload).map(ControlMessage::RequestOk)
}
ControlMessageType::Publish => {
Publish::parse_payload(&mut payload).map(ControlMessage::Publish)
}
ControlMessageType::PublishOk => {
PublishOk::parse_payload(&mut payload).map(ControlMessage::PublishOk)
}
ControlMessageType::PublishDone => {
PublishDone::parse_payload(&mut payload).map(ControlMessage::PublishDone)
}
ControlMessageType::ClientSetup => {
ClientSetup::parse_payload(&mut payload).map(ControlMessage::ClientSetup)
}
ControlMessageType::Fetch => Fetch::parse_payload(&mut payload).map(ControlMessage::Fetch),
ControlMessageType::FetchCancel => {
FetchCancel::parse_payload(&mut payload).map(ControlMessage::FetchCancel)
}
ControlMessageType::FetchOk => {
FetchOk::parse_payload(&mut payload).map(ControlMessage::FetchOk)
}
ControlMessageType::GoAway => GoAway::parse_payload(&mut payload).map(ControlMessage::Goaway),
ControlMessageType::MaxRequestId => {
MaxRequestId::parse_payload(&mut payload).map(ControlMessage::MaxRequestId)
}
ControlMessageType::ServerSetup => {
ServerSetup::parse_payload(&mut payload).map(ControlMessage::ServerSetup)
}
ControlMessageType::Subscribe => {
Subscribe::parse_payload(&mut payload).map(ControlMessage::Subscribe)
}
ControlMessageType::SubscribeOk => {
SubscribeOk::parse_payload(&mut payload).map(ControlMessage::SubscribeOk)
}
ControlMessageType::RequestUpdate => {
RequestUpdate::parse_payload(&mut payload).map(ControlMessage::RequestUpdate)
}
ControlMessageType::RequestsBlocked => {
RequestsBlocked::parse_payload(&mut payload).map(ControlMessage::RequestsBlocked)
}
ControlMessageType::TrackStatus => {
TrackStatus::parse_payload(&mut payload).map(ControlMessage::TrackStatus)
}
ControlMessageType::Unsubscribe => {
Unsubscribe::parse_payload(&mut payload).map(ControlMessage::Unsubscribe)
}
ControlMessageType::SubscribeNamespace => {
SubscribeNamespace::parse_payload(&mut payload).map(ControlMessage::SubscribeNamespace)
}
ControlMessageType::UnsubscribeNamespace => {
UnsubscribeNamespace::parse_payload(&mut payload).map(ControlMessage::UnsubscribeNamespace)
}
ControlMessageType::Switch => Switch::parse_payload(&mut payload).map(ControlMessage::Switch),
}
.map_err(|err| ParseError::ProtocolViolation {
context: "ControlMessage::deserialize(payload)",
details: err.to_string(),
})?;
if payload.has_remaining() {
return Err(ParseError::ProtocolViolation {
context: "ControlMessage::deserialize(final_check)",
details: format!(
"Extra {} bytes remaining in payload after parsing",
payload.remaining()
),
});
};
Ok(message)
}
pub fn serialize(&self) -> Result<Bytes, ParseError> {
match self {
ControlMessage::Namespace(msg) => msg.serialize(),
ControlMessage::NamespaceDone(msg) => msg.serialize(),
ControlMessage::PublishNamespace(msg) => msg.serialize(),
ControlMessage::PublishNamespaceCancel(msg) => msg.serialize(),
ControlMessage::PublishNamespaceDone(msg) => msg.serialize(),
ControlMessage::RequestError(msg) => msg.serialize(),
ControlMessage::RequestOk(msg) => msg.serialize(),
ControlMessage::Publish(msg) => msg.serialize(),
ControlMessage::PublishOk(msg) => msg.serialize(),
ControlMessage::PublishDone(msg) => msg.serialize(),
ControlMessage::ClientSetup(msg) => msg.serialize(),
ControlMessage::Fetch(msg) => msg.serialize(),
ControlMessage::FetchCancel(msg) => msg.serialize(),
ControlMessage::FetchOk(msg) => msg.serialize(),
ControlMessage::Goaway(msg) => msg.serialize(),
ControlMessage::MaxRequestId(msg) => msg.serialize(),
ControlMessage::ServerSetup(msg) => msg.serialize(),
ControlMessage::Subscribe(msg) => msg.serialize(),
ControlMessage::SubscribeOk(msg) => msg.serialize(),
ControlMessage::RequestUpdate(msg) => msg.serialize(),
ControlMessage::RequestsBlocked(msg) => msg.serialize(),
ControlMessage::TrackStatus(msg) => msg.serialize(),
ControlMessage::Unsubscribe(msg) => msg.serialize(),
ControlMessage::SubscribeNamespace(msg) => msg.serialize(),
ControlMessage::UnsubscribeNamespace(msg) => msg.serialize(),
ControlMessage::Switch(msg) => msg.serialize(),
}
}
pub fn get_type(&self) -> ControlMessageType {
match self {
ControlMessage::Namespace(_) => ControlMessageType::Namespace,
ControlMessage::NamespaceDone(_) => ControlMessageType::NamespaceDone,
ControlMessage::PublishNamespace(_) => ControlMessageType::PublishNamespace,
ControlMessage::PublishNamespaceCancel(_) => ControlMessageType::PublishNamespaceCancel,
ControlMessage::PublishNamespaceDone(_) => ControlMessageType::PublishNamespaceDone,
ControlMessage::RequestError(_) => ControlMessageType::RequestError,
ControlMessage::RequestOk(_) => ControlMessageType::RequestOk,
ControlMessage::Publish(_) => ControlMessageType::Publish,
ControlMessage::PublishOk(_) => ControlMessageType::PublishOk,
ControlMessage::PublishDone(_) => ControlMessageType::PublishDone,
ControlMessage::ClientSetup(_) => ControlMessageType::ClientSetup,
ControlMessage::Fetch(_) => ControlMessageType::Fetch,
ControlMessage::FetchCancel(_) => ControlMessageType::FetchCancel,
ControlMessage::FetchOk(_) => ControlMessageType::FetchOk,
ControlMessage::Goaway(_) => ControlMessageType::GoAway,
ControlMessage::MaxRequestId(_) => ControlMessageType::MaxRequestId,
ControlMessage::ServerSetup(_) => ControlMessageType::ServerSetup,
ControlMessage::Subscribe(_) => ControlMessageType::Subscribe,
ControlMessage::SubscribeOk(_) => ControlMessageType::SubscribeOk,
ControlMessage::RequestUpdate(_) => ControlMessageType::RequestUpdate,
ControlMessage::RequestsBlocked(_) => ControlMessageType::RequestsBlocked,
ControlMessage::TrackStatus(_) => ControlMessageType::TrackStatus,
ControlMessage::Unsubscribe(_) => ControlMessageType::Unsubscribe,
ControlMessage::SubscribeNamespace(_) => ControlMessageType::SubscribeNamespace,
ControlMessage::UnsubscribeNamespace(_) => ControlMessageType::UnsubscribeNamespace,
ControlMessage::Switch(_) => ControlMessageType::Switch,
}
}
}
#[cfg(test)]
mod tests {
use crate::model::{
common::tuple::Tuple,
parameter::{authorization_token::AuthorizationToken, message_parameter::MessageParameter},
};
use super::*;
#[test]
fn test_announce_roundtrip() {
let request_id = 12345;
let track_namespace = Tuple::from_utf8_path("god/dayyum");
let parameters = vec![MessageParameter::new_authorization_token(
AuthorizationToken::new_use_value(0, Bytes::from_static(b"test-token")),
)];
let announce = PublishNamespace {
request_id,
track_namespace,
parameters,
};
let mut buf = announce.serialize().unwrap();
let deserialized = ControlMessage::deserialize(&mut buf).unwrap();
if let ControlMessage::PublishNamespace(deserialized_announce) = deserialized {
assert_eq!(*deserialized_announce, announce);
} else {
panic!("Expected ControlMessage::PublishNamespace variant");
}
assert!(!buf.has_remaining());
}
}