use bytes::Bytes;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{SinkExt, StreamExt};
use prost::Message as _;
use tokio::net::TcpStream;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async_with_config};
use crate::client::DEFAULT_MAX_WEBSOCKET_MESSAGE_BYTES;
use super::error::{Result, RtcError};
use super::proto::event::{
HealthCheckRequest, JoinRequest, LeaveCallRequest, SfuEvent, SfuRequest, sfu_request,
};
type WsStream = WebSocketStream<MaybeTlsStream<TcpStream>>;
pub async fn connect(endpoint: &str) -> Result<(SfuSender, SfuReceiver)> {
connect_with_limit(endpoint, DEFAULT_MAX_WEBSOCKET_MESSAGE_BYTES).await
}
pub async fn connect_with_limit(
endpoint: &str,
max_message_bytes: usize,
) -> Result<(SfuSender, SfuReceiver)> {
let config = WebSocketConfig::default()
.read_buffer_size(max_message_bytes.clamp(1, 128 * 1024))
.max_frame_size(Some(max_message_bytes))
.max_message_size(Some(max_message_bytes));
let (stream, _response) = connect_async_with_config(endpoint, Some(config), false).await?;
let (sink, source) = stream.split();
Ok((
SfuSender {
sink,
max_message_bytes,
},
SfuReceiver {
source,
max_message_bytes,
},
))
}
#[derive(Debug)]
pub struct SfuSender {
sink: SplitSink<WsStream, Message>,
max_message_bytes: usize,
}
impl SfuSender {
pub async fn send(&mut self, request: SfuRequest) -> Result<()> {
ensure_message_size(
"SFU WebSocket outbound message",
request.encoded_len(),
self.max_message_bytes,
)?;
let bytes = Bytes::from(request.encode_to_vec());
self.sink.send(Message::Binary(bytes)).await?;
Ok(())
}
pub async fn send_join(&mut self, join: JoinRequest) -> Result<()> {
self.send(SfuRequest {
request_payload: Some(sfu_request::RequestPayload::JoinRequest(join)),
})
.await
}
pub async fn send_health_check(&mut self) -> Result<()> {
self.send(SfuRequest {
request_payload: Some(sfu_request::RequestPayload::HealthCheckRequest(
HealthCheckRequest {},
)),
})
.await
}
pub async fn send_leave(
&mut self,
session_id: impl Into<String>,
reason: impl Into<String>,
) -> Result<()> {
self.send(SfuRequest {
request_payload: Some(sfu_request::RequestPayload::LeaveCallRequest(
LeaveCallRequest {
session_id: session_id.into(),
reason: reason.into(),
},
)),
})
.await
}
pub async fn close(&mut self) -> Result<()> {
self.sink.close().await?;
Ok(())
}
}
#[derive(Debug)]
pub struct SfuReceiver {
source: SplitStream<WsStream>,
max_message_bytes: usize,
}
impl SfuReceiver {
pub async fn recv(&mut self) -> Result<Option<SfuEvent>> {
while let Some(message) = self.source.next().await {
let message = message.map_err(|error| {
RtcError::from_websocket_with_boundary(error, "SFU WebSocket inbound message")
})?;
match message {
Message::Binary(bytes) => {
ensure_message_size(
"SFU WebSocket inbound message",
bytes.len(),
self.max_message_bytes,
)?;
let event = SfuEvent::decode(bytes)?;
return Ok(Some(event));
}
Message::Close(frame) => {
let reason = frame
.map(|f| format!("{} {}", f.code, f.reason))
.unwrap_or_else(|| "no close frame".to_owned());
tracing::debug!(reason = %reason, "stream.sfu.ws.closed");
return Ok(None);
}
_ => continue,
}
}
Ok(None)
}
}
fn ensure_message_size(boundary: &'static str, actual: usize, limit: usize) -> Result<()> {
if actual > limit {
return Err(RtcError::SizeLimitExceeded {
boundary,
limit,
actual,
});
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn message_limit_rejects_before_decode() {
ensure_message_size("test WebSocket message", 8, 8).expect("at limit");
assert!(matches!(
ensure_message_size("test WebSocket message", 9, 8).expect_err("oversized message"),
RtcError::SizeLimitExceeded {
boundary: "test WebSocket message",
limit: 8,
actual: 9,
}
));
}
}