#[cfg(feature = "websocket")]
use axum::{
Router,
extract::ws::{WebSocket, WebSocketUpgrade},
http::StatusCode,
response::{IntoResponse, Response},
};
#[cfg(all(feature = "websocket", feature = "security"))]
use axum::http::header::AUTHORIZATION;
#[cfg(feature = "websocket")]
use futures_util::SinkExt;
#[cfg(feature = "websocket")]
use futures_util::StreamExt;
#[cfg(feature = "websocket")]
use std::pin::Pin;
#[cfg(feature = "websocket")]
use std::sync::Arc;
#[cfg(feature = "websocket")]
use crate::websocket::{AppState, ConnectionManager};
#[cfg(feature = "websocket")]
use crate::websocket::{MAX_MESSAGE_SIZE, WebSocketMessage, parse_websocket_message};
#[cfg(feature = "websocket")]
use crate::core::ApiMetadata;
#[cfg(feature = "websocket")]
use crate::define_registration;
#[cfg(feature = "websocket")]
pub type BoxFuture<'a, T> = Pin<Box<dyn std::future::Future<Output = T> + Send + 'a>>;
#[cfg(feature = "websocket")]
pub trait WebSocketHandler: Send + Sync {
fn handle(&self, message: WebSocketMessage) -> BoxFuture<'static, WebSocketMessage>;
}
#[cfg(feature = "websocket")]
define_registration!(WebSocketRoute, Arc<dyn WebSocketHandler>, ApiMetadata);
#[cfg(feature = "websocket")]
pub struct DefaultWebSocketHandler;
#[cfg(feature = "websocket")]
impl WebSocketHandler for DefaultWebSocketHandler {
fn handle(&self, message: WebSocketMessage) -> BoxFuture<'static, WebSocketMessage> {
Box::pin(async move {
match message {
WebSocketMessage::Request { id, method, .. } => WebSocketMessage::Response {
id,
result: serde_json::json!({"status": "ok", "method": method}),
},
_ => message,
}
})
}
}
#[cfg(feature = "websocket")]
pub struct ValidatedWebSocketUpgrade {
ws: WebSocketUpgrade,
manager: Arc<ConnectionManager>,
handler: Option<Arc<dyn WebSocketHandler>>,
}
#[cfg(feature = "websocket")]
impl ValidatedWebSocketUpgrade {
pub fn with_handler(mut self, handler: Arc<dyn WebSocketHandler>) -> Self {
self.handler = Some(handler);
self
}
}
#[cfg(feature = "websocket")]
impl IntoResponse for ValidatedWebSocketUpgrade {
fn into_response(self) -> Response {
let handler = self
.handler
.unwrap_or_else(|| Arc::new(DefaultWebSocketHandler));
self.ws
.on_upgrade(move |socket| handle_socket(socket, self.manager.clone(), handler))
}
}
#[cfg(feature = "websocket")]
impl<S> axum::extract::FromRequest<S> for ValidatedWebSocketUpgrade
where
S: Clone + Send + Sync + 'static,
{
type Rejection = StatusCode;
async fn from_request(req: axum::extract::Request, state: &S) -> Result<Self, Self::Rejection> {
let req = req;
#[cfg(feature = "security")]
let bearer_token: Option<String> = req
.headers()
.get(AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|h| h.strip_prefix("Bearer "))
.map(String::from);
let app_state = req.extensions().get::<Arc<AppState>>().cloned();
#[cfg(feature = "security")]
if let Some(ref state_ref) = app_state {
let bearer_cfg = state_ref.config.auth.as_ref();
let api_cfg = state_ref.config.api_key_auth.as_ref();
if bearer_cfg.is_some() || api_cfg.is_some() {
let bearer_ok = bearer_cfg
.zip(bearer_token.as_deref())
.and_then(|(auth, token)| auth.validate_token(token))
.is_some();
let api_key = req.headers().get("x-api-key").and_then(|v| v.to_str().ok());
let api_ok = api_cfg
.zip(api_key)
.and_then(|(store, key)| store.validate_key(key, "unknown"))
.is_some();
if !bearer_ok && !api_ok {
return Err(StatusCode::UNAUTHORIZED);
}
}
}
let ws = axum::extract::ws::WebSocketUpgrade::from_request(req, state)
.await
.map_err(|_| StatusCode::BAD_REQUEST)?;
let manager = app_state
.map(|s| s.manager.clone())
.unwrap_or_else(|| Arc::new(ConnectionManager::new()));
Ok(Self {
ws,
manager,
handler: None,
})
}
}
#[cfg(feature = "websocket")]
pub async fn websocket_upgrade(ws: ValidatedWebSocketUpgrade) -> impl IntoResponse {
ws }
#[cfg(feature = "websocket")]
async fn handle_socket(
socket: WebSocket,
manager: Arc<ConnectionManager>,
handler: Arc<dyn WebSocketHandler>,
) {
#[cfg(feature = "context")]
{
let ctx = crate::context::current_or_new();
crate::context::scope(ctx, handle_socket_inner(socket, manager, handler)).await
}
#[cfg(not(feature = "context"))]
{
handle_socket_inner(socket, manager, handler).await;
}
}
#[cfg(feature = "websocket")]
async fn handle_socket_inner(
socket: WebSocket,
manager: Arc<ConnectionManager>,
handler: Arc<dyn WebSocketHandler>,
) {
let conn_id = uuid::Uuid::new_v4().to_string();
let (conn, mut receiver) = WebSocketConnection::new(conn_id.clone());
manager.add_connection(conn_id.clone(), conn.clone()).await;
struct ConnectionGuard {
conn_id: String,
manager: Arc<ConnectionManager>,
}
impl Drop for ConnectionGuard {
fn drop(&mut self) {
let manager = self.manager.clone();
let conn_id = self.conn_id.clone();
tokio::spawn(async move {
manager.remove_connection(&conn_id).await;
});
}
}
let _guard = ConnectionGuard {
conn_id: conn_id.clone(),
manager: manager.clone(),
};
let (mut sink, mut stream) = socket.split();
let forwarder = tokio::spawn(async move {
while let Some(msg) = receiver.recv().await {
let json = serde_json::to_string(&msg).unwrap_or_else(|_| {
serde_json::to_string(&WebSocketMessage::Error {
id: String::new(),
error: "Internal serialization error".to_string(),
})
.unwrap_or_else(|_| {
r#"{"type":"error","id":"","error":"Internal error"}"#.to_string()
})
});
if sink
.send(axum::extract::ws::Message::Text(json.into()))
.await
.is_err()
{
break;
}
}
});
while let Some(result) = stream.next().await {
match result {
Ok(msg) => {
if let Ok(text) = msg.to_text() {
if text.len() > MAX_MESSAGE_SIZE {
forwarder.abort();
return;
}
match parse_websocket_message(text) {
Ok(ws_msg) => {
let response = handler.handle(ws_msg).await;
let _ = conn.send(response).await;
}
Err(e) => {
let error_msg = WebSocketMessage::Error {
id: String::new(),
error: e,
};
let _ = conn.send(error_msg).await;
}
}
}
}
Err(_) => {
break;
}
}
}
forwarder.abort();
manager.remove_connection(&conn_id).await;
}
#[cfg(feature = "websocket")]
use crate::websocket::WebSocketConnection;
#[cfg(feature = "websocket")]
pub fn build() -> Router {
let mut router = Router::new();
let manager = Arc::new(ConnectionManager::new());
let state = Arc::new(AppState::new(manager));
for route in inventory::iter::<WebSocketRoute> {
let path = format!("/{}", route.name);
router = router.route(
&path,
axum::routing::get(move |ws: ValidatedWebSocketUpgrade| async move {
ws.with_handler((route.create_fn)())
})
.with_state(state.clone()),
);
}
router
}
#[cfg(feature = "websocket")]
#[cfg(test)]
mod tests {
use super::*;
use futures_util::FutureExt;
use std::panic::AssertUnwindSafe;
#[tokio::test]
async fn handle_socket_closes_connection_on_oversized_message() {
let app = Router::new().route("/ws", axum::routing::get(websocket_upgrade));
let server = axum_test::TestServer::builder().http_transport().build(app);
let mut ws = server.get_websocket("/ws").await.into_websocket().await;
let oversized = "x".repeat(MAX_MESSAGE_SIZE + 1);
ws.send_text(&oversized).await;
tokio::time::sleep(tokio::time::Duration::from_millis(300)).await;
let receive_result = tokio::time::timeout(
tokio::time::Duration::from_millis(500),
AssertUnwindSafe(ws.receive_text()).catch_unwind(),
)
.await;
match receive_result {
Err(_) => { }
Ok(Err(_)) => { }
Ok(Ok(text)) => {
let parse_result: Result<WebSocketMessage, _> = serde_json::from_str(&text);
assert!(
parse_result.is_err(),
"Should not receive a valid WebSocketMessage, got: {}",
text
);
}
}
}
#[tokio::test]
async fn broadcast_reaches_connected_client() {
let manager = Arc::new(ConnectionManager::new());
let state = Arc::new(AppState::new(manager.clone()));
let app = Router::new()
.route("/ws", axum::routing::get(websocket_upgrade))
.layer(axum::Extension(state));
let server = axum_test::TestServer::builder().http_transport().build(app);
let mut ws = server.get_websocket("/ws").await.into_websocket().await;
tokio::time::sleep(tokio::time::Duration::from_millis(300)).await;
assert_eq!(manager.connection_count().await, 1);
let notification = Arc::new(WebSocketMessage::Notification {
event: "ping".to_string(),
data: serde_json::json!({"hello": "world"}),
});
manager.broadcast(¬ification).await;
let received = tokio::time::timeout(tokio::time::Duration::from_secs(2), ws.receive_json())
.await
.expect("broadcast must reach the client within timeout");
match received {
WebSocketMessage::Notification { event, data } => {
assert_eq!(event, "ping");
assert_eq!(data["hello"], "world");
}
other => panic!("Expected Notification, got {:?}", other),
}
assert_eq!(
manager.connection_count().await,
1,
"connection must stay registered after a successful broadcast"
);
}
#[tokio::test]
async fn oversized_message_removes_connection_from_manager() {
let manager = Arc::new(ConnectionManager::new());
let state = Arc::new(AppState::new(manager.clone()));
let app = Router::new()
.route("/ws", axum::routing::get(websocket_upgrade))
.layer(axum::Extension(state));
let server = axum_test::TestServer::builder().http_transport().build(app);
let mut ws = server.get_websocket("/ws").await.into_websocket().await;
tokio::time::sleep(tokio::time::Duration::from_millis(300)).await;
assert_eq!(manager.connection_count().await, 1);
let oversized = "x".repeat(MAX_MESSAGE_SIZE + 1);
ws.send_text(&oversized).await;
let cleaned = tokio::time::timeout(tokio::time::Duration::from_secs(2), async {
while manager.connection_count().await > 0 {
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
}
})
.await;
assert!(
cleaned.is_ok(),
"connection entry must be removed after oversized-message close (leak)"
);
}
#[tokio::test]
async fn handle_socket_ignores_binary_messages() {
let app = Router::new().route("/ws", axum::routing::get(websocket_upgrade));
let server = axum_test::TestServer::builder().http_transport().build(app);
let mut ws = server.get_websocket("/ws").await.into_websocket().await;
ws.send_message(axum_test::WsMessage::Binary(vec![0xFF, 0xFE, 0xFD].into()))
.await;
let request = WebSocketMessage::Request {
id: "after-binary".to_string(),
method: "ping".to_string(),
params: serde_json::json!({}),
};
ws.send_json(&request).await;
let response: WebSocketMessage = ws.receive_json().await;
match response {
WebSocketMessage::Response { id, result } => {
assert_eq!(id, "after-binary");
assert_eq!(result["status"], "ok");
}
_ => panic!("Expected Response, got {:?}", response),
}
}
}