use bytes::{Buf, BufMut, BytesMut};
use tokio::time::{Duration, Instant, sleep_until};
use tracing::{error, info, warn};
use wtransport::error::StreamReadError;
use wtransport::{RecvStream, SendStream};
use crate::model::control::control_message::{ControlMessage, ControlMessageTrait};
use crate::model::error::{ParseError, TerminationCode};
const CONTROL_MESSAGE_TIMEOUT: Duration = Duration::from_secs(5);
const MTU_SIZE: usize = 1500;
pub struct ControlStreamHandler {
send: SendStream,
recv: RecvStream,
recv_bytes: BytesMut,
recv_buf: Box<[u8; MTU_SIZE]>,
partial_message_deadline: Option<Instant>,
}
impl ControlStreamHandler {
pub fn new(send: SendStream, recv: RecvStream) -> Self {
Self {
send,
recv,
recv_bytes: BytesMut::new(),
recv_buf: Box::new([0; MTU_SIZE]),
partial_message_deadline: None,
}
}
pub async fn send(&mut self, message: &ControlMessage) -> Result<(), TerminationCode> {
let bytes = message
.serialize()
.map_err(|_| TerminationCode::InternalError)?;
if (self.send.write_all(&bytes).await).is_err() {
warn!("Error sending message: {:?}", message);
return Err(TerminationCode::InternalError);
}
Ok(())
}
pub async fn send_impl(
&mut self,
message: &impl ControlMessageTrait,
) -> Result<(), TerminationCode> {
let bytes = message
.serialize()
.map_err(|_| TerminationCode::InternalError)?;
if (self.send.write_all(&bytes).await).is_err() {
warn!("Error sending (send_impl) message: {:?}", message);
return Err(TerminationCode::InternalError);
}
Ok(())
}
pub async fn next_message(&mut self) -> Result<ControlMessage, TerminationCode> {
loop {
if !self.recv_bytes.is_empty() {
let mut bytes = self.recv_bytes.clone().freeze();
let original_remaining = bytes.remaining();
match ControlMessage::deserialize(&mut bytes) {
Ok(msg) => {
let consumed = original_remaining - bytes.remaining();
self.recv_bytes.advance(consumed);
self.partial_message_deadline = None;
return Ok(msg);
}
Err(ParseError::ProtocolViolation { .. }) => {
return Err(TerminationCode::ProtocolViolation);
}
Err(ParseError::NotEnoughBytes { .. }) if self.partial_message_deadline.is_none() => {
self.partial_message_deadline = Some(Instant::now() + CONTROL_MESSAGE_TIMEOUT);
}
_ => {}
}
}
self.read_more_data().await?;
}
}
async fn read_more_data(&mut self) -> Result<(), TerminationCode> {
if let Some(deadline) = self.partial_message_deadline {
if Instant::now() >= deadline {
self.partial_message_deadline = None;
self.recv_bytes.clear();
return Err(TerminationCode::ControlMessageTimeout);
}
tokio::select! {
biased;
_ = sleep_until(deadline) => {
info!("Control message timeout reached");
self.partial_message_deadline = None;
self.recv_bytes.clear();
Err(TerminationCode::ControlMessageTimeout)
}
res = self.recv.read(&mut self.recv_buf[..]) => {
self.handle_read_result(res, true)
}
}
} else {
let res = self.recv.read(&mut self.recv_buf[..]).await;
self.handle_read_result(res, false)
}
}
fn handle_read_result(
&mut self,
res: Result<Option<usize>, wtransport::error::StreamReadError>,
is_partial_message: bool,
) -> Result<(), TerminationCode> {
match res {
Ok(Some(n)) => {
if n > 0 {
self.recv_bytes.put_slice(&self.recv_buf[..n]);
}
Ok(())
}
Ok(None) => {
if is_partial_message {
info!("Stream closed cleanly while waiting for partial message");
} else {
warn!("Stream closed cleanly while waiting for data");
}
Err(TerminationCode::InternalError)
}
Err(e) => {
match e {
StreamReadError::NotConnected => {
info!("Client disconnected while reading control stream");
Err(TerminationCode::NoError)
}
_ => {
if is_partial_message {
info!("Error reading from stream: {:?}", e);
} else {
error!("Error reading from stream: {:?}", e);
}
Err(TerminationCode::InternalError)
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::common::location::Location;
use crate::model::common::pair::KeyValuePair;
use crate::model::common::reason_phrase::ReasonPhrase;
use crate::model::common::tuple::{Tuple, TupleField};
use crate::model::common::varint::BufMutVarIntExt;
use crate::model::control::client_setup::ClientSetup;
use crate::model::control::constant::RequestErrorCode;
use crate::model::control::constant::{ControlMessageType, GroupOrder};
use crate::model::control::publish_namespace::PublishNamespace;
use crate::model::control::publish_namespace_cancel::PublishNamespaceCancel;
use crate::model::control::request_ok::RequestOk;
use crate::model::control::server_setup::ServerSetup;
use crate::model::control::subscribe::Subscribe;
use crate::model::control::subscribe_ok::SubscribeOk;
use crate::model::parameter::authorization_token::AuthorizationToken;
use crate::model::parameter::message_parameter::MessageParameter;
use bytes::Bytes;
use std::error::Error;
use std::sync::Arc;
use tokio::sync::Mutex;
use tokio::time::sleep;
use wtransport::endpoint::IntoConnectOptions;
use wtransport::{ClientConfig, Connection, Endpoint, Identity};
struct TestSetup {
client: Connection,
server: Connection,
}
impl TestSetup {
async fn new() -> Result<Self, Box<dyn Error>> {
let server_identity = Identity::self_signed(std::iter::once("localhost"))
.map_err(|e| format!("Failed to create server identity: {e}"))?;
let server_cert_hash = server_identity.certificate_chain().as_slice()[0].hash();
let server_config = wtransport::ServerConfig::builder()
.with_bind_address(
"127.0.0.1:0"
.parse()
.map_err(|e| format!("Failed to parse bind address: {e}"))?,
)
.with_identity(server_identity)
.build();
let server_endpoint = Endpoint::server(server_config)
.map_err(|e| format!("Failed to create server endpoint: {e}"))?;
let server_addr = server_endpoint
.local_addr()
.map_err(|e| format!("Failed to get server local address: {e}"))?;
let (tx, rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let result = async {
let incoming = server_endpoint.accept().await;
let session_request = incoming
.await
.map_err(|e| format!("Failed to await session request: {e}"))
.unwrap();
let server = session_request.accept().await.unwrap();
Ok::<_, Box<dyn Error + Send>>(server)
}
.await;
if tx.send(result).is_err() {
eprintln!("Failed to send server connection result back through the channel");
}
});
let client_config = ClientConfig::builder()
.with_bind_default()
.with_server_certificate_hashes(vec![server_cert_hash])
.build();
let client_endpoint = Endpoint::client(client_config)
.map_err(|e| format!("Failed to create client endpoint: {e}"))?;
let client = client_endpoint
.connect(
format!("https://{}:{}", server_addr.ip(), server_addr.port())
.as_str()
.into_options(),
)
.await
.map_err(|e| format!("Client connection failed: {e}"))?;
let server = rx
.await
.map_err(|_| "Server task failed to send connection back")?
.map_err(|e| format!("Server connection error: {e}"))?;
Ok(Self { client, server })
}
async fn create_control_plane(
&self,
) -> Result<(ControlStreamHandler, SendStream), Box<dyn Error>> {
let (client_send, client_recv) = match self.client.open_bi().await {
Ok(stream_fut) => match stream_fut.await {
Ok((send, recv)) => (send, recv),
Err(e) => return Err(format!("Failed to await client stream: {e}").into()),
},
Err(e) => return Err(format!("Failed to open client stream: {e}").into()),
};
let (server_send, _) = self
.server
.accept_bi()
.await
.map_err(|e| format!("Failed to accept server stream: {e}"))?;
server_send.set_priority(i32::MAX);
let plane = ControlStreamHandler::new(client_send, client_recv);
Ok((plane, server_send))
}
}
fn create_test_publish_namespace() -> PublishNamespace {
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")),
)];
PublishNamespace {
request_id,
track_namespace,
parameters,
}
}
fn create_test_announce_ok() -> RequestOk {
let request_id = 12345;
RequestOk::new(request_id, vec![])
}
fn create_test_announce_cancel() -> PublishNamespaceCancel {
let request_id = 1337;
let error_code = RequestErrorCode::InternalError;
let reason_phrase = ReasonPhrase::try_new("bomboclad".to_string()).unwrap();
PublishNamespaceCancel {
request_id,
error_code,
reason_phrase,
}
}
fn create_test_subscribe() -> Subscribe {
let request_id = 128242;
let track_namespace = Tuple::from_utf8_path("nein/nein/nein");
let track_name = TupleField::from_utf8("track_42");
let start_location = Location {
group: 81,
object: 81,
};
Subscribe::new_absolute_range(
request_id,
track_namespace,
track_name,
start_location,
100,
vec![
MessageParameter::new_subscriber_priority(31),
MessageParameter::new_group_order(GroupOrder::Original),
MessageParameter::new_forward(true),
],
)
}
fn create_test_subscribe_ok() -> SubscribeOk {
use crate::model::control::constant::GroupOrder;
use crate::model::extension_header::track_extension::TrackExtension;
use crate::model::parameter::message_parameter::MessageParameter;
SubscribeOk::new(
145136,
999,
vec![
MessageParameter::new_expires(16),
MessageParameter::new_group_order(GroupOrder::Ascending),
MessageParameter::new_largest_object(Location {
group: 34,
object: 0,
}),
MessageParameter::new_expires(100),
],
vec![TrackExtension::DeliveryTimeout { timeout_ms: 5000 }],
)
}
fn create_test_client_setup() -> ClientSetup {
let setup_parameters = vec![
KeyValuePair::try_new_varint(0, 10).unwrap(),
KeyValuePair::try_new_bytes(1, Bytes::from_static(b"Set me up!")).unwrap(),
];
ClientSetup { setup_parameters }
}
fn create_test_server_setup() -> ServerSetup {
let setup_parameters = vec![
KeyValuePair::try_new_varint(0, 10).unwrap(),
KeyValuePair::try_new_bytes(1, Bytes::from_static(b"Set me up!")).unwrap(),
];
ServerSetup { setup_parameters }
}
#[tokio::test]
async fn test_connection_setup() -> Result<(), Box<dyn Error>> {
let setup = TestSetup::new().await?;
let (client_send, _) = setup
.client
.open_bi()
.await?
.await
.map_err(|e| format!("Failed to open bidirectional stream: {e}"))?;
let mut client_send = client_send;
let (_, mut server_recv) = setup
.server
.accept_bi()
.await
.map_err(|e| format!("Failed to accept bidirectional stream: {e}"))?;
client_send.write_all(&[1, 2, 3, 4]).await?;
let mut buf = [0; 4];
server_recv.read_exact(&mut buf).await?;
assert_eq!(buf, [1, 2, 3, 4]);
Ok(())
}
#[tokio::test]
async fn test_message_timeout() -> Result<(), Box<dyn Error>> {
let setup = TestSetup::new().await?;
let (mut plane, mut server_send) = setup.create_control_plane().await?;
let mut bytes = BytesMut::new();
bytes.put_vi(ControlMessageType::PublishNamespace)?;
let bytes = bytes.freeze();
server_send.write_all(&bytes).await?;
match plane.next_message().await {
Err(TerminationCode::ControlMessageTimeout) => Ok(()),
other => panic!("Expected timeout, got {other:?}"),
}
}
#[tokio::test]
async fn test_successful_message() -> Result<(), Box<dyn Error>> {
let setup = TestSetup::new().await?;
let (mut plane, mut server_send) = setup.create_control_plane().await?;
let announce = Box::new(create_test_publish_namespace());
let msg = announce.clone();
let bytes = msg.serialize().unwrap();
server_send.write_all(&bytes).await?;
let received = plane.next_message().await.unwrap();
match received {
ControlMessage::PublishNamespace(rec_announce) => assert_eq!(rec_announce, announce),
_ => panic!("Received incorrect message type"),
}
Ok(())
}
#[tokio::test]
async fn test_partial_message_completion() -> Result<(), Box<dyn Error>> {
let setup = TestSetup::new().await?;
let (mut plane, mut server_send) = setup.create_control_plane().await?;
let announce_cancel = Box::new(create_test_announce_cancel());
let msg = ControlMessage::PublishNamespaceCancel(announce_cancel.clone()); let bytes = msg.serialize().unwrap();
let half = bytes.len() / 2;
server_send.write_all(&bytes[..half]).await?;
let remaining_bytes = bytes[half..].to_vec();
tokio::spawn(async move {
sleep(Duration::from_millis(100)).await;
server_send.write_all(&remaining_bytes).await.unwrap();
});
let received = plane.next_message().await.unwrap();
match received {
ControlMessage::PublishNamespaceCancel(rec_cancel) => assert_eq!(rec_cancel, announce_cancel),
_ => panic!("Received incorrect message type"),
}
Ok(())
}
#[tokio::test]
async fn test_multiple_messages() -> Result<(), Box<dyn Error>> {
let setup = TestSetup::new().await?;
let (mut plane, server_send) = setup.create_control_plane().await?;
let server_send = Arc::new(Mutex::new(server_send));
let announce1 = Box::new(create_test_publish_namespace());
let announce_ok1 = Box::new(create_test_announce_ok());
let subscribe1 = Box::new(create_test_subscribe());
let subscribe_ok1 = Box::new(create_test_subscribe_ok());
let announce_cancel1 = Box::new(create_test_announce_cancel());
let client_setup = Box::new(create_test_client_setup());
let server_setup = Box::new(create_test_server_setup());
let messages_to_send = vec![
ControlMessage::ClientSetup(client_setup),
ControlMessage::ServerSetup(server_setup),
ControlMessage::PublishNamespace(announce1),
ControlMessage::RequestOk(announce_ok1),
ControlMessage::Subscribe(subscribe1),
ControlMessage::SubscribeOk(subscribe_ok1),
ControlMessage::PublishNamespaceCancel(announce_cancel1),
];
let messages_clone = messages_to_send.clone();
let server_send_clone = Arc::clone(&server_send);
tokio::spawn(async move {
let mut sender = server_send_clone.lock().await;
for msg in messages_clone {
let bytes = msg.serialize().unwrap();
if sender.write_all(&bytes).await.is_err() {
eprintln!("Error sending message in test task");
return; }
sleep(Duration::from_millis(10)).await;
}
});
for expected_msg in messages_to_send {
let received_msg = plane.next_message().await.unwrap();
assert_eq!(
received_msg, expected_msg,
"Mismatch between sent and received message.\nExpected: {expected_msg:?}\nReceived: {received_msg:?}"
);
}
Ok(())
}
}