#[cfg(feature = "http-server")]
use super::{AdaptiveStreamController, StreamOptions, WebSocketTransport, WsMessage};
use crate::{
Result as PjsResult,
infrastructure::bounded_channel::{self, ByteBoundedSender, byte_bounded_channel},
security::{RateLimitConfig, RateLimitGuard, WebSocketRateLimiter},
};
#[cfg(feature = "http-server")]
use axum::{
extract::{
ConnectInfo, State, WebSocketUpgrade,
ws::{Message, WebSocket},
},
http::StatusCode,
response::{IntoResponse, Response},
};
use futures::StreamExt;
use serde_json::Value;
use std::collections::HashMap;
use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::RwLock;
use tokio::sync::broadcast::error::RecvError;
use tracing::{debug, error, info, warn};
use uuid;
const OUTGOING_QUEUE_CAPACITY: usize = 1000;
const MAX_QUEUED_OUTGOING_BYTES: usize = 16 * 1024 * 1024;
const SESSION_CLEANUP_INTERVAL: Duration = Duration::from_secs(60);
const SESSION_MAX_AGE: Duration = Duration::from_secs(3600);
pub struct AxumWebSocketTransport {
controller: Arc<AdaptiveStreamController>,
active_connections: Arc<RwLock<Vec<String>>>,
outgoing_channels: Arc<RwLock<HashMap<String, ByteBoundedSender<String>>>>,
connection_sessions: Arc<RwLock<HashMap<String, Vec<String>>>>,
rate_limiter: Arc<WebSocketRateLimiter>,
}
impl AxumWebSocketTransport {
pub fn new() -> Self {
Self::with_rate_limit_config(RateLimitConfig::default())
}
pub fn with_rate_limit_config(config: RateLimitConfig) -> Self {
let controller = Arc::new(AdaptiveStreamController::new());
let weak_controller = Arc::downgrade(&controller);
tokio::spawn(async move {
let mut interval = tokio::time::interval(SESSION_CLEANUP_INTERVAL);
loop {
interval.tick().await;
let Some(controller) = weak_controller.upgrade() else {
break;
};
controller.cleanup_expired_sessions(SESSION_MAX_AGE).await;
}
});
Self {
controller,
active_connections: Arc::new(RwLock::new(Vec::new())),
outgoing_channels: Arc::new(RwLock::new(HashMap::new())),
connection_sessions: Arc::new(RwLock::new(HashMap::new())),
rate_limiter: Arc::new(WebSocketRateLimiter::new(config)),
}
}
pub async fn upgrade_handler(
ws: WebSocketUpgrade,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
State(transport): State<Arc<Self>>,
) -> Response {
let client_ip = addr.ip();
if let Err(e) = transport.rate_limiter.check_request(client_ip) {
warn!("WebSocket upgrade denied for IP {}: {}", client_ip, e);
return (StatusCode::TOO_MANY_REQUESTS, e.to_string()).into_response();
}
let max_frame_size = transport.rate_limiter.config().max_frame_size;
let ws = ws
.max_message_size(max_frame_size)
.max_frame_size(max_frame_size);
ws.on_upgrade(move |socket| transport.handle_socket(socket, client_ip))
}
pub async fn handle_socket(self: Arc<Self>, socket: WebSocket, client_ip: IpAddr) {
info!("New WebSocket connection established from {}", client_ip);
let write_timeout = self.rate_limiter.config().write_timeout;
let guard = match RateLimitGuard::new(self.rate_limiter.clone(), client_ip) {
Ok(g) => Arc::new(g),
Err(e) => {
warn!(
"WebSocket connection rejected for IP {} (rate limit): {}",
client_ip, e
);
let (mut sender, _) = socket.split();
let _ = super::send_with_write_timeout(
&mut sender,
Message::Close(Some(axum::extract::ws::CloseFrame {
code: 1008, reason: e.to_string().into(),
})),
write_timeout,
)
.await;
return;
}
};
let connection_id = uuid::Uuid::new_v4().to_string();
self.active_connections
.write()
.await
.push(connection_id.clone());
let frame_rx = self.controller.subscribe_frames();
let (outgoing_tx, mut outgoing_rx) =
byte_bounded_channel::<String>(OUTGOING_QUEUE_CAPACITY, MAX_QUEUED_OUTGOING_BYTES);
self.outgoing_channels
.write()
.await
.insert(connection_id.clone(), outgoing_tx);
let (mut sender, mut receiver) = socket.split();
let transport_clone = self.clone();
let connection_id_clone = connection_id.clone();
let guard_for_task = guard.clone();
let websocket_task = {
let mut frame_rx = frame_rx;
tokio::spawn(async move {
loop {
tokio::select! {
recv_result = frame_rx.recv() => {
match recv_result {
Ok((_session_id, message)) => {
match serde_json::to_string(&message) {
Ok(json_str) => {
if let Err(e) = super::send_with_write_timeout(&mut sender, Message::Text(json_str.into()), write_timeout).await {
error!("Failed to send message to client: {}", e);
break;
}
}
Err(e) => {
error!("Failed to serialize message: {}", e);
}
}
}
Err(RecvError::Lagged(skipped)) => {
warn!("Frame broadcast lagged; skipped {} frames", skipped);
}
Err(RecvError::Closed) => {
debug!("Frame broadcast channel closed");
break;
}
}
}
Some(envelope) = outgoing_rx.recv() => {
let (json_str, _budget_permit) = envelope.split();
if let Err(e) = super::send_with_write_timeout(&mut sender, Message::Text(json_str.into()), write_timeout).await {
error!("Failed to send outgoing message to client: {}", e);
break;
}
}
Some(msg) = receiver.next() => {
match msg {
Ok(Message::Text(text)) => {
if let Err(e) = guard_for_task.check_message(text.len()) {
warn!(
"Inbound text frame rejected for IP {} (rate limit): {}",
client_ip, e
);
let _ = super::send_with_write_timeout(
&mut sender,
Message::Close(Some(axum::extract::ws::CloseFrame {
code: 1008,
reason: e.to_string().into(),
})),
write_timeout,
).await;
break;
}
match serde_json::from_str::<WsMessage>(&text) {
Ok(ws_message) => {
if let Err(e) = transport_clone.handle_websocket_message(connection_id_clone.clone(), ws_message).await {
error!("Failed to handle message: {}", e);
}
}
Err(e) => {
warn!("Failed to parse WebSocket message: {}", e);
}
}
}
Ok(Message::Binary(data)) => {
if let Err(e) = guard_for_task.check_message(data.len()) {
warn!(
"Inbound binary frame rejected for IP {} (rate limit): {}",
client_ip, e
);
let _ = super::send_with_write_timeout(
&mut sender,
Message::Close(Some(axum::extract::ws::CloseFrame {
code: 1008,
reason: e.to_string().into(),
})),
write_timeout,
).await;
break;
}
debug!("Received binary data: {} bytes", data.len());
}
Ok(Message::Ping(data)) => {
if let Err(e) = super::send_with_write_timeout(&mut sender, Message::Pong(data), write_timeout).await {
error!("Failed to send pong: {}", e);
break;
}
}
Ok(Message::Pong(_)) => {
debug!("Received pong from client");
}
Ok(Message::Close(_)) => {
info!("Client closed WebSocket connection");
break;
}
Err(e) => {
error!("WebSocket error: {}", e);
break;
}
}
}
else => {
break;
}
}
}
drop(guard_for_task);
})
};
if let Err(e) = websocket_task.await {
error!("WebSocket task failed: {}", e);
}
self.outgoing_channels.write().await.remove(&connection_id);
let mut connections = self.active_connections.write().await;
connections.retain(|conn_id| *conn_id != connection_id);
drop(connections);
drop(guard);
if let Some(session_ids) = self
.connection_sessions
.write()
.await
.remove(&connection_id)
{
for session_id in session_ids {
self.controller.remove_session(&session_id).await;
}
}
info!("WebSocket connection closed for {}", client_ip);
}
pub fn controller(&self) -> Arc<AdaptiveStreamController> {
self.controller.clone()
}
pub async fn active_connection_count(&self) -> usize {
self.active_connections.read().await.len()
}
async fn handle_websocket_message(
&self,
connection_id: String,
message: WsMessage,
) -> PjsResult<()> {
debug!(
"Handling WebSocket message for connection {}: {:?}",
connection_id, message
);
match message {
WsMessage::FrameAck {
session_id,
frame_id,
processing_time_ms,
} => {
self.controller
.handle_frame_ack(&session_id, frame_id, processing_time_ms)
.await?;
}
WsMessage::StreamInit {
session_id: _,
data,
options,
} => {
let session_id = self.controller.create_session(data, options).await?;
self.controller.start_streaming(&session_id).await?;
self.connection_sessions
.write()
.await
.entry(connection_id.clone())
.or_default()
.push(session_id);
info!(
"Created new streaming session for connection {}",
connection_id
);
}
WsMessage::Ping { timestamp: _ } => {
debug!("Received ping from connection {}", connection_id);
}
_ => {
warn!("Unhandled message type from connection {}", connection_id);
}
}
Ok(())
}
}
impl Default for AxumWebSocketTransport {
fn default() -> Self {
Self::new()
}
}
impl WebSocketTransport for AxumWebSocketTransport {
type Connection = String;
type StartStreamFuture<'a>
= impl Future<Output = PjsResult<String>> + Send + 'a
where
Self: 'a;
type SendFrameFuture<'a>
= impl Future<Output = PjsResult<()>> + Send + 'a
where
Self: 'a;
type HandleMessageFuture<'a>
= impl Future<Output = PjsResult<()>> + Send + 'a
where
Self: 'a;
type CloseStreamFuture<'a>
= impl Future<Output = PjsResult<()>> + Send + 'a
where
Self: 'a;
fn start_stream(
&self,
_connection: Arc<Self::Connection>,
data: Value,
options: StreamOptions,
) -> Self::StartStreamFuture<'_> {
async move {
let session_id = self.controller.create_session(data, options).await?;
self.controller.start_streaming(&session_id).await?;
Ok(session_id)
}
}
fn send_frame(
&self,
connection: Arc<Self::Connection>,
message: WsMessage,
) -> Self::SendFrameFuture<'_> {
async move {
let tx = self
.outgoing_channels
.read()
.await
.get(connection.as_ref())
.cloned();
if let Some(tx) = tx {
match serde_json::to_string(&message) {
Ok(json_str) => {
let len = json_str.len();
match tx.try_send(json_str, len) {
Ok(()) => {}
Err(bounded_channel::TrySendError::BudgetExceeded(_)) => {
warn!(
"send_frame: dropping frame for connection {} (byte budget exceeded, {} bytes)",
connection.as_ref(),
len
);
}
Err(bounded_channel::TrySendError::Channel(_)) => {
warn!(
"send_frame: dropping frame for connection {} (channel full or closed)",
connection.as_ref()
);
}
}
}
Err(e) => {
warn!(
"send_frame: failed to serialize frame for connection {}: {}",
connection.as_ref(),
e
);
}
}
} else {
warn!(
"send_frame: no outgoing channel for connection {}",
connection.as_ref()
);
}
Ok(())
}
}
fn handle_message(
&self,
_connection: Arc<Self::Connection>,
message: WsMessage,
) -> Self::HandleMessageFuture<'_> {
async move {
match message {
WsMessage::StreamInit { data, options, .. } => {
info!("Initializing new stream");
let session_id = self.controller.create_session(data, options).await?;
self.controller.start_streaming(&session_id).await?;
}
WsMessage::FrameAck {
session_id,
frame_id,
processing_time_ms,
} => {
debug!(
"Received frame ack: session={}, frame={}, time={}ms",
session_id, frame_id, processing_time_ms
);
self.controller
.handle_frame_ack(&session_id, frame_id, processing_time_ms)
.await?;
}
WsMessage::Ping { timestamp } => {
debug!("Received ping with timestamp: {}", timestamp);
}
WsMessage::Error {
session_id,
error,
code,
} => {
warn!(
"Received error from client: session={:?}, error={}, code={}",
session_id, error, code
);
}
_ => {
warn!("Unhandled message type: {:?}", message);
}
}
Ok(())
}
}
fn close_stream(&self, session_id: &str) -> Self::CloseStreamFuture<'_> {
let session_id = session_id.to_string();
async move {
info!("Closing stream session: {}", session_id);
self.controller.remove_session(&session_id).await;
Ok(())
}
}
}
pub fn create_websocket_router() -> axum::Router<Arc<AxumWebSocketTransport>> {
use axum::routing::get;
axum::Router::new().route("/ws", get(AxumWebSocketTransport::upgrade_handler))
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[tokio::test]
async fn test_transport_creation() {
let transport = AxumWebSocketTransport::new();
assert!(Arc::strong_count(&transport.controller) >= 1);
}
#[tokio::test]
async fn test_stream_initialization() {
let transport = AxumWebSocketTransport::new();
let data = json!({
"critical": {"id": 1, "status": "active"},
"metadata": {"created": "2024-01-15T12:00:00Z"}
});
let session_id = transport
.controller
.create_session(data, StreamOptions::default())
.await
.unwrap();
assert!(!session_id.is_empty());
transport
.controller
.start_streaming(&session_id)
.await
.unwrap();
}
#[tokio::test]
async fn test_outgoing_channel_is_bounded() {
let (tx, mut rx) =
byte_bounded_channel::<String>(OUTGOING_QUEUE_CAPACITY, MAX_QUEUED_OUTGOING_BYTES);
for _ in 0..OUTGOING_QUEUE_CAPACITY {
tx.try_send("ping".to_string(), 4)
.expect("channel should accept sends up to its capacity");
}
let result = tx.try_send("ping".to_string(), 4);
assert!(
matches!(
result,
Err(bounded_channel::TrySendError::Channel(
tokio::sync::mpsc::error::TrySendError::Full(_)
))
),
"channel must reject sends past capacity instead of growing unbounded"
);
rx.recv().await.expect("receiver should still be open");
tx.try_send("ping".to_string(), 4)
.expect("channel should accept a send after capacity is freed");
}
#[tokio::test]
async fn test_send_frame_drops_when_byte_budget_exceeded() {
let transport = AxumWebSocketTransport::new();
let connection_id = "test-connection".to_string();
let (tx, mut rx) =
byte_bounded_channel::<String>(OUTGOING_QUEUE_CAPACITY, MAX_QUEUED_OUTGOING_BYTES);
transport
.outgoing_channels
.write()
.await
.insert(connection_id.clone(), tx);
let connection = Arc::new(connection_id);
let oversized_message = WsMessage::Error {
session_id: None,
error: "x".repeat(MAX_QUEUED_OUTGOING_BYTES + 1),
code: 0,
};
transport
.send_frame(connection, oversized_message)
.await
.expect("send_frame returns Ok even when it drops the frame");
assert!(
rx.try_recv().is_err(),
"an over-budget frame must be dropped, not queued"
);
}
#[tokio::test]
async fn test_send_frame_drops_on_full_channel_without_blocking() {
let transport = AxumWebSocketTransport::new();
let connection_id = "test-connection".to_string();
let (tx, mut rx) =
byte_bounded_channel::<String>(OUTGOING_QUEUE_CAPACITY, MAX_QUEUED_OUTGOING_BYTES);
transport
.outgoing_channels
.write()
.await
.insert(connection_id.clone(), tx);
let connection = Arc::new(connection_id);
for _ in 0..OUTGOING_QUEUE_CAPACITY {
transport
.send_frame(connection.clone(), WsMessage::Ping { timestamp: 0 })
.await
.expect("send_frame should accept sends up to channel capacity");
}
tokio::time::timeout(
std::time::Duration::from_secs(2),
transport.send_frame(connection.clone(), WsMessage::Ping { timestamp: 0 }),
)
.await
.expect("send_frame must not block when the outgoing channel is full")
.expect("send_frame must return Ok even when dropping the overflow frame");
rx.close();
let mut drained = 0;
while rx.try_recv().is_ok() {
drained += 1;
}
assert_eq!(
drained, OUTGOING_QUEUE_CAPACITY,
"the overflow frame must have been dropped, not queued"
);
}
}