use axum::extract::ws::{Message, WebSocket};
use axum::extract::{Path, State, WebSocketUpgrade};
use axum::response::IntoResponse;
use futures::{SinkExt, StreamExt};
use tracing::debug;
use crate::api::state::SharedState;
use crate::engine::types::GameStatus;
use super::manager::ClientId;
use super::messages::{WsCommand, WsEvent};
pub async fn ws_handler(
ws: WebSocketUpgrade,
Path(id): Path<String>,
State(state): State<SharedState>,
) -> impl IntoResponse {
ws.on_upgrade(move |socket| handle_socket(socket, id, state))
}
async fn handle_socket(socket: WebSocket, game_id: String, state: SharedState) {
let initial_event = {
let games = state.games.read().await;
match games.get(&game_id) {
Some(game) => {
let status = game.status();
let check = matches!(status, GameStatus::Check);
let player = match game.side_to_move() {
crate::engine::types::Color::White => "white",
crate::engine::types::Color::Black => "black",
};
WsEvent::subscribed(
&game_id,
&game.to_fen(),
status.as_str(),
player,
game.move_history().len(),
check,
)
}
None => {
let (mut sink, _) = socket.split();
let err = WsEvent::error(&format!("game not found: {game_id}"));
let _ = sink.send(Message::Text(err.to_json().into())).await;
let _ = sink.close().await;
return;
}
}
};
let (client_id, mut rx) = state.ws.subscribe(&game_id).await;
let (mut sink, mut stream) = socket.split();
if sink
.send(Message::Text(initial_event.to_json().into()))
.await
.is_err()
{
cleanup(&state, &game_id, client_id).await;
return;
}
let writer_state = state.clone();
let writer_gid = game_id.clone();
let mut writer = tokio::spawn(async move {
while let Some(event) = rx.recv().await {
if sink
.send(Message::Text(event.to_json().into()))
.await
.is_err()
{
break;
}
}
let _ = sink.close().await;
cleanup(&writer_state, &writer_gid, client_id).await;
});
let reader_state = state.clone();
let reader_gid = game_id.clone();
let mut reader = tokio::spawn(async move {
while let Some(Ok(msg)) = stream.next().await {
match msg {
Message::Text(text) => {
handle_client_message(&reader_state, &reader_gid, client_id, &text).await;
}
Message::Close(_) => break,
_ => {} }
}
});
tokio::select! {
_ = &mut writer => { reader.abort(); }
_ = &mut reader => { writer.abort(); }
}
cleanup(&state, &game_id, client_id).await;
}
async fn handle_client_message(
state: &SharedState,
_current_game_id: &str,
_client_id: ClientId,
text: &str,
) {
let cmd = match serde_json::from_str::<WsCommand>(text) {
Ok(c) => c,
Err(e) => {
debug!("invalid WS command: {e}");
return;
}
};
match cmd {
WsCommand::Ping => {
let pong = WsEvent::pong();
state.ws.broadcast(_current_game_id, pong).await;
}
WsCommand::Subscribe { game_id } => {
debug!(
game_id,
"client requested subscribe (already subscribed via URL)"
);
let games = state.games.read().await;
if let Some(game) = games.get(&game_id) {
let status = game.status();
let check = matches!(status, GameStatus::Check);
let player = match game.side_to_move() {
crate::engine::types::Color::White => "white",
crate::engine::types::Color::Black => "black",
};
let evt = WsEvent::subscribed(
&game_id,
&game.to_fen(),
status.as_str(),
player,
game.move_history().len(),
check,
);
state.ws.broadcast(&game_id, evt).await;
}
}
WsCommand::Unsubscribe { game_id } => {
debug!(game_id, "client requested unsubscribe");
state.ws.unsubscribe(&game_id, _client_id).await;
}
}
}
async fn cleanup(state: &SharedState, game_id: &str, client_id: ClientId) {
state.ws.unsubscribe(game_id, client_id).await;
debug!(game_id, client_id, "WS session cleaned up");
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn handler_type_check() {
fn assert_handler<F, Fut, R>(_: F)
where
F: FnOnce(WebSocketUpgrade, Path<String>, State<SharedState>) -> Fut,
Fut: std::future::Future<Output = R>,
R: IntoResponse,
{
}
assert_handler(ws_handler);
}
}