use axum::{
Json,
extract::{
FromRequestParts,
ws::{Message, WebSocket, WebSocketUpgrade},
},
http::StatusCode,
response::{IntoResponse, Response},
};
use tokio::sync::broadcast;
use tokio::time::Duration;
use tokio_util::sync::CancellationToken;
use crate::{
models::{ApiResponse, OutputMode},
routes::Device,
};
pub async fn ws_handler(Device(dev): Device, req: axum::extract::Request) -> Response {
if *dev.output_mode_tx.borrow() == OutputMode::Dump {
return (
StatusCode::FORBIDDEN,
Json(ApiResponse {
success: false,
message: "Server is in dump-only mode; WebSocket streaming is disabled".to_string(),
}),
)
.into_response();
}
let (mut parts, _body) = req.into_parts();
let ws = match WebSocketUpgrade::from_request_parts(&mut parts, &()).await {
Ok(ws) => ws,
Err(rejection) => return rejection.into_response(),
};
let rx = dev.csi_tx.subscribe();
let shutdown = dev.shutdown.clone();
let id = dev.id.clone();
ws.on_upgrade(move |socket| handle_socket(socket, rx, shutdown, id))
.into_response()
}
async fn handle_socket(
mut socket: WebSocket,
mut rx: broadcast::Receiver<Vec<u8>>,
shutdown: CancellationToken,
id: String,
) {
let mut sent: u64 = 0;
let mut dropped: u64 = 0;
let mut metrics = tokio::time::interval(Duration::from_secs(1));
metrics.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
metrics.tick().await;
loop {
tokio::select! {
_ = metrics.tick() => {
if sent > 0 || dropped > 0 {
tracing::debug!(
target: "csi_metrics",
"ws {id}: sent={sent}/s dropped={dropped}/s",
);
sent = 0;
dropped = 0;
}
}
_ = shutdown.cancelled() => {
let _ = socket.send(Message::Close(None)).await;
break;
}
result = rx.recv() => {
match result {
Ok(data) => {
if socket.send(Message::Binary(data.into())).await.is_err() {
break;
}
sent += 1;
}
Err(broadcast::error::RecvError::Closed) => {
break;
}
Err(broadcast::error::RecvError::Lagged(n)) => {
dropped += n;
tracing::warn!("WebSocket client for {id} lagged — dropped {n} CSI packets");
}
}
}
msg = socket.recv() => {
match msg {
Some(Ok(Message::Close(_))) | None => break,
_ => {} }
}
}
}
tracing::debug!("WebSocket client for {id} disconnected");
}