use async_trait::async_trait;
use futures_util::{SinkExt, StreamExt};
use std::sync::Arc;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{broadcast, mpsc};
use tokio_tungstenite::{
accept_async, connect_async,
tungstenite::{error, protocol::Message, Error as TungsteniteError},
MaybeTlsStream, WebSocketStream,
};
use crate::{
command::ConnectionState, connection::Connection, error::TransportError, event::TransportEvent,
packet::Packet, protocol::AdapterStats, ConnectionInfo, SessionId,
};
use crate::adapters::outbound::SEND_QUEUE_CAPACITY;
enum MessageProcessResult {
Packet(Packet),
Heartbeat,
PeerClosed,
Error(WebSocketError),
}
#[derive(Debug, thiserror::Error)]
pub enum WebSocketError {
#[error("Tungstenite error: {0}")]
Tungstenite(#[from] TungsteniteError),
#[error("IO error: {0}")]
Io(#[from] std::io::Error),
#[error("Connection closed")]
ConnectionClosed,
#[error("Invalid message type")]
InvalidMessageType,
#[error("Configuration error: {0}")]
Config(String),
}
impl From<WebSocketError> for TransportError {
fn from(error: WebSocketError) -> Self {
match error {
WebSocketError::Tungstenite(e) => {
TransportError::connection_error(format!("WebSocket protocol error: {}", e), true)
}
WebSocketError::Io(e) => {
TransportError::connection_error(format!("WebSocket IO error: {}", e), true)
}
WebSocketError::ConnectionClosed => {
TransportError::connection_error("WebSocket connection closed", false)
}
WebSocketError::InvalidMessageType => {
TransportError::protocol_error("websocket", "Invalid message type")
}
WebSocketError::Config(msg) => TransportError::config_error("websocket", msg),
}
}
}
pub struct WebSocketAdapter<C> {
state: crate::adapters::core::ConnState,
config: C,
stats: AdapterStats,
connection_info: ConnectionInfo,
send_queue: mpsc::Sender<Packet>,
event_sender: broadcast::Sender<TransportEvent>,
shutdown_sender: mpsc::UnboundedSender<()>,
event_loop_handle: Option<tokio::task::JoinHandle<()>>,
frame_policy: Arc<std::sync::atomic::AtomicU8>,
}
impl<C> WebSocketAdapter<C> {
pub fn new(config: C) -> Self {
let (event_sender, _) = broadcast::channel(8192);
let (send_queue_tx, _) = mpsc::channel(SEND_QUEUE_CAPACITY);
let (shutdown_tx, _) = mpsc::unbounded_channel();
Self {
state: crate::adapters::core::ConnState::new(
crate::adapters::core::ConnStatus::Connecting,
),
config,
stats: AdapterStats::new(),
connection_info: ConnectionInfo::default(),
send_queue: send_queue_tx,
event_sender,
shutdown_sender: shutdown_tx,
event_loop_handle: None,
frame_policy: Arc::new(std::sync::atomic::AtomicU8::new(
crate::packet::FramePolicy::Lenient as u8,
)),
}
}
pub async fn new_with_stream(
config: C,
stream: WebSocketStream<MaybeTlsStream<TcpStream>>,
event_sender: broadcast::Sender<TransportEvent>,
) -> Result<Self, WebSocketError> {
let mut connection_info = ConnectionInfo::default();
connection_info.protocol = "websocket".to_string();
connection_info.state = ConnectionState::Connected;
connection_info.established_at = std::time::SystemTime::now();
let state =
crate::adapters::core::ConnState::new(crate::adapters::core::ConnStatus::Connected);
let frame_policy = Arc::new(std::sync::atomic::AtomicU8::new(
crate::packet::FramePolicy::Lenient as u8,
));
let (send_queue_tx, send_queue_rx) = mpsc::channel(SEND_QUEUE_CAPACITY);
let (shutdown_tx, shutdown_rx) = mpsc::unbounded_channel();
let event_loop_handle = Self::start_event_loop(
stream,
state.clone(),
send_queue_rx,
shutdown_rx,
event_sender.clone(),
frame_policy.clone(),
)
.await;
Ok(Self {
state,
config,
stats: AdapterStats::new(),
connection_info,
send_queue: send_queue_tx,
event_sender,
shutdown_sender: shutdown_tx,
event_loop_handle: Some(event_loop_handle),
frame_policy,
})
}
pub fn subscribe_events(&self) -> broadcast::Receiver<TransportEvent> {
self.event_sender.subscribe()
}
async fn start_event_loop(
mut stream: WebSocketStream<MaybeTlsStream<TcpStream>>,
state: crate::adapters::core::ConnState,
mut send_queue: mpsc::Receiver<Packet>,
mut shutdown_signal: mpsc::UnboundedReceiver<()>,
event_sender: broadcast::Sender<TransportEvent>,
frame_policy: Arc<std::sync::atomic::AtomicU8>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let current_session_id = state.session_id();
tracing::debug!(
"[START] WebSocket event loop started (session: {})",
current_session_id
);
loop {
let current_session_id = state.session_id();
tokio::select! {
read_result = stream.next() => {
match read_result {
Some(Ok(message)) => {
let policy = crate::packet::FramePolicy::from(
frame_policy.load(std::sync::atomic::Ordering::Relaxed),
);
match Self::process_websocket_message(message, policy) {
MessageProcessResult::Packet(packet) => {
tracing::debug!("[RECV] WebSocket received packet: {} bytes (session: {})", packet.payload.len(), current_session_id);
let event = TransportEvent::MessageReceived(packet);
if let Err(e) = event_sender.send(event) {
tracing::warn!("[RECV] Failed to send receive event: {:?}", e);
}
}
MessageProcessResult::Heartbeat => {
continue;
}
MessageProcessResult::PeerClosed => {
let close_event = TransportEvent::ConnectionClosed { reason: crate::error::CloseReason::Normal };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[CLOSE] Failed to notify upper layer connection closed: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[CLOSE] Notified upper layer connection closed: session {}", current_session_id);
}
state.set_status(crate::adapters::core::ConnStatus::Closed);
break;
}
MessageProcessResult::Error(e) => {
tracing::error!("[ERROR] WebSocket message processing error: {:?} (session: {})", e, current_session_id);
let close_event = TransportEvent::ConnectionClosed { reason: crate::error::CloseReason::Error(format!("{:?}", e)) };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[ERROR] Failed to notify upper layer message processing error: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[ERROR] Notified upper layer message processing error: session {}", current_session_id);
}
state.set_status(crate::adapters::core::ConnStatus::Closed);
break;
}
}
}
Some(Err(e)) => {
let reason = match e {
TungsteniteError::Protocol(error::ProtocolError::ResetWithoutClosingHandshake) => {
tracing::debug!("[CLOSE] Peer actively reset WebSocket connection (session: {})", current_session_id);
crate::error::CloseReason::Normal
}
TungsteniteError::ConnectionClosed => {
tracing::debug!("[CLOSE] Peer actively closed WebSocket connection (session: {})", current_session_id);
crate::error::CloseReason::Normal
}
_ => {
tracing::error!("[ERROR] WebSocket connection error: {:?} (session: {})", e, current_session_id);
crate::error::CloseReason::Error(format!("{:?}", e))
}
};
let close_event = TransportEvent::ConnectionClosed { reason };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[CLOSE] Failed to notify upper layer connection closed: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[CLOSE] Notified upper layer connection closed: session {}", current_session_id);
}
state.set_status(crate::adapters::core::ConnStatus::Closed);
break;
}
None => {
tracing::debug!("[CLOSE] Peer actively closed WebSocket connection (session: {})", current_session_id);
let close_event = TransportEvent::ConnectionClosed { reason: crate::error::CloseReason::Normal };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[CLOSE] Failed to notify upper layer connection closed: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[CLOSE] Notified upper layer connection closed: session {}", current_session_id);
}
state.set_status(crate::adapters::core::ConnStatus::Closed);
break;
}
}
}
packet = send_queue.recv() => {
if let Some(packet) = packet {
let message = Message::Binary(packet.to_bytes());
match stream.send(message).await {
Ok(_) => {
tracing::debug!("[SEND] WebSocket send successful: {} bytes (session: {})", packet.payload.len(), current_session_id);
let event = TransportEvent::MessageSent { packet_id: packet.header.message_id };
if let Err(e) = event_sender.send(event) {
tracing::warn!("[SEND] Failed to send send event: {:?}", e);
}
}
Err(e) => {
tracing::error!("[ERROR] WebSocket send error: {:?} (session: {})", e, current_session_id);
let close_event = TransportEvent::ConnectionClosed { reason: crate::error::CloseReason::Error(format!("{:?}", e)) };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[ERROR] Failed to notify upper layer send error: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[ERROR] Notified upper layer send error: session {}", current_session_id);
}
state.set_status(crate::adapters::core::ConnStatus::Closed);
break;
}
}
}
}
_ = shutdown_signal.recv() => {
tracing::info!("[STOP] Received shutdown signal, stopping WebSocket event loop (session: {})", current_session_id);
tracing::debug!("[CLOSE] Send WebSocket Close frame for graceful shutdown");
if let Err(e) = stream.close(None).await {
tracing::warn!("[SEND] Failed to send WebSocket Close frame: {:?} (session: {})", e, current_session_id);
} else {
tracing::debug!("[SEND] WebSocket Close frame sent successfully (session: {})", current_session_id);
}
tracing::debug!("[CLOSE] Active close, not sending close event");
break;
}
}
}
tracing::debug!(
"[SUCCESS] WebSocket event loop ended (session: {})",
current_session_id
);
})
}
fn process_websocket_message(
message: Message,
frame_policy: crate::packet::FramePolicy,
) -> MessageProcessResult {
let strict = frame_policy == crate::packet::FramePolicy::Strict;
match message {
Message::Binary(data) => {
if data.len() < 16 {
if strict {
return MessageProcessResult::Error(WebSocketError::InvalidMessageType);
}
let packet = Packet::one_way(0, data.clone());
return MessageProcessResult::Packet(packet);
}
match Packet::from_bytes(&data) {
Ok(packet) => {
tracing::debug!(
"[RECV] WebSocket packet parsing successful: {} bytes",
packet.payload.len()
);
MessageProcessResult::Packet(packet)
}
Err(e) => {
if strict {
tracing::debug!(
"[RECV] WebSocket packet parse failed under strict policy: {:?}",
e
);
return MessageProcessResult::Error(WebSocketError::InvalidMessageType);
}
tracing::debug!("[RECV] WebSocket packet parsing failed: {:?}, creating basic data packet", e);
let packet = Packet::one_way(0, data.clone());
MessageProcessResult::Packet(packet)
}
}
}
Message::Text(text) => {
tracing::debug!(
"[RECV] WebSocket received text message: {} bytes",
text.len()
);
let packet = Packet::one_way(0, text.as_bytes());
MessageProcessResult::Packet(packet)
}
Message::Close(_) => {
tracing::debug!("[RECV] WebSocket received Close message");
MessageProcessResult::PeerClosed
}
Message::Ping(_) | Message::Pong(_) => {
MessageProcessResult::Heartbeat
}
Message::Frame(_) => {
tracing::warn!("[RECV] WebSocket received unsupported Frame message");
MessageProcessResult::Error(WebSocketError::InvalidMessageType)
}
}
}
}
#[async_trait]
impl<C: Send + Sync + 'static> Connection for WebSocketAdapter<C> {
async fn send(&mut self, packet: Packet) -> Result<(), TransportError> {
crate::adapters::outbound::send_bounded(
&self.send_queue,
packet,
"websocket_outbound_queue",
"WebSocket connection closed",
)
.await
}
async fn close(&mut self) -> Result<(), TransportError> {
let current_session_id = self.state.session_id();
tracing::debug!(
"[CLOSE] Close WebSocket connection (session: {})",
current_session_id
);
let _ = self.shutdown_sender.send(());
if let Some(handle) = self.event_loop_handle.take() {
let _ = handle.await;
}
self.state
.set_status(crate::adapters::core::ConnStatus::Closed);
Ok(())
}
fn session_id(&self) -> SessionId {
self.state.session_id()
}
fn set_session_id(&mut self, session_id: SessionId) {
self.state.set_session_id(session_id);
}
fn connection_info(&self) -> ConnectionInfo {
self.connection_info.clone()
}
fn is_connected(&self) -> bool {
self.state.is_connected()
}
async fn flush(&mut self) -> Result<(), TransportError> {
Ok(())
}
fn event_stream(
&self,
) -> Option<tokio::sync::broadcast::Receiver<crate::event::TransportEvent>> {
Some(self.event_sender.subscribe())
}
fn set_frame_policy(&self, policy: crate::packet::FramePolicy) {
self.frame_policy
.store(policy as u8, std::sync::atomic::Ordering::Relaxed);
}
}
pub(crate) struct WebSocketServerBuilder<C> {
config: Option<C>,
}
impl<C> WebSocketServerBuilder<C> {
pub(crate) fn new() -> Self {
Self { config: None }
}
pub(crate) fn config(mut self, config: C) -> Self {
self.config = Some(config);
self
}
pub(crate) fn bind_address(self, _addr: std::net::SocketAddr) -> Self {
self
}
pub(crate) async fn build(self) -> Result<WebSocketServer<C>, WebSocketError> {
let config = self
.config
.ok_or_else(|| WebSocketError::Config("Missing WebSocket server config".to_string()))?;
Ok(WebSocketServer {
config,
listener: None,
})
}
}
pub(crate) struct WebSocketServer<C> {
config: C,
listener: Option<TcpListener>,
}
impl<C: 'static> WebSocketServer<C> {
pub(crate) async fn accept(&mut self) -> Result<WebSocketAdapter<C>, WebSocketError>
where
C: Clone + crate::protocol::ProtocolConfig,
{
if self.listener.is_none() {
let bind_addr = if let Some(ws_config) = (&self.config as &dyn std::any::Any)
.downcast_ref::<crate::protocol::WebSocketServerConfig>(
) {
ws_config.bind_address.to_string()
} else {
"127.0.0.1:8080".parse().unwrap()
};
let listener = TcpListener::bind(&bind_addr).await?;
tracing::debug!("[START] WebSocket server listening on: {}", bind_addr);
self.listener = Some(listener);
}
if let Some(listener) = &self.listener {
let (tcp_stream, addr) = listener.accept().await?;
tracing::debug!("[ACCEPT] WebSocket server accepted connection: {}", addr);
let maybe_tls_stream = MaybeTlsStream::Plain(tcp_stream);
let ws_stream = accept_async(maybe_tls_stream).await?;
let (event_sender, _) = broadcast::channel(8192);
WebSocketAdapter::new_with_stream(self.config.clone(), ws_stream, event_sender).await
} else {
Err(WebSocketError::Config("No listener available".to_string()))
}
}
pub(crate) fn local_addr(&self) -> Result<std::net::SocketAddr, WebSocketError> {
if let Some(listener) = &self.listener {
listener.local_addr().map_err(WebSocketError::Io)
} else {
Err(WebSocketError::Config("Server not bound".to_string()))
}
}
pub(crate) async fn shutdown(&mut self) -> Result<(), WebSocketError> {
self.listener.take();
Ok(())
}
}
pub(crate) struct WebSocketClientBuilder<C> {
config: Option<C>,
}
impl<C> WebSocketClientBuilder<C> {
pub(crate) fn new() -> Self {
Self { config: None }
}
pub(crate) fn config(mut self, config: C) -> Self {
self.config = Some(config);
self
}
pub(crate) fn target_url<S: Into<String>>(self, _url: S) -> Self {
self
}
pub(crate) async fn connect(self) -> Result<WebSocketAdapter<C>, WebSocketError>
where
C: crate::protocol::ProtocolConfig,
{
let config = self
.config
.ok_or_else(|| WebSocketError::Config("Missing WebSocket client config".to_string()))?;
let url = if let Some(ws_config) =
(&config as &dyn std::any::Any).downcast_ref::<crate::protocol::WebSocketClientConfig>()
{
ws_config.target_url.clone()
} else {
"ws://127.0.0.1:8080".to_string()
};
tracing::debug!("[CONNECT] WebSocket client connecting to: {}", url);
let (ws_stream, _) = connect_async(&url).await?;
tracing::debug!("[SUCCESS] WebSocket client connected to: {}", url);
let (event_sender, _) = broadcast::channel(8192);
WebSocketAdapter::new_with_stream(config, ws_stream, event_sender).await
}
}