diff --git a/backend/src/main.rs b/backend/src/main.rs index a7f42fa..0c2bd9c 100644 --- a/backend/src/main.rs +++ b/backend/src/main.rs @@ -13,30 +13,35 @@ use axum::extract::ws::{Message, Utf8Bytes, WebSocket}; use futures_util::{sink::SinkExt, stream::StreamExt}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; -use std::sync::{Arc, Mutex}; -use tokio::sync::Mutex as TokioMutex; +use std::sync::Arc; +use tokio::sync::Mutex; use tracing::{info, warn}; use uuid::Uuid; -use crate::game::{GamePhase, Room}; +use crate::game::{Room, ThrowOutcome}; use crate::protocol::*; +type WsSender = futures_util::stream::SplitSink; +type WsReceiver = futures_util::stream::SplitStream; + #[derive(Clone)] struct AppState { - rooms: Arc>>>>, + rooms: Arc>>>>, } impl AppState { fn new() -> Self { - Self { rooms: Arc::new(Mutex::new(HashMap::new())) } + Self { + rooms: Arc::new(Mutex::new(HashMap::new())), + } } - fn get_or_create_room(&self, room_id: &str) -> Arc> { - let mut rooms = self.rooms.lock().unwrap(); + async fn get_or_create_room(&self, room_id: &str) -> Arc> { + let mut rooms = self.rooms.lock().await; if let Some(room) = rooms.get(room_id).cloned() { return room; } - let room = Arc::new(TokioMutex::new(Room::new(room_id))); + let room = Arc::new(Mutex::new(Room::new(room_id))); rooms.insert(room_id.to_string(), room.clone()); room } @@ -81,7 +86,7 @@ async fn health() -> impl IntoResponse { async fn new_room(State(state): State>) -> impl IntoResponse { let room_id = generate_room_code(); - state.get_or_create_room(&room_id); + state.get_or_create_room(&room_id).await; (StatusCode::OK, axum::Json(NewRoomResponse { room: room_id })) } @@ -93,50 +98,78 @@ async fn ws_handler( ws.on_upgrade(move |socket| handle_socket(socket, state, query.room, query.team)) } -async fn handle_socket(socket: WebSocket, state: Arc, room_id: String, preferred_team: Option) { - let room_arc = state.get_or_create_room(&room_id); - let player_id = Uuid::new_v4().to_string(); - let player_id_for_recv = player_id.clone(); +async fn handle_socket( + socket: WebSocket, + state: Arc, + room_id: String, + preferred_team: Option, +) { + let room = state.get_or_create_room(&room_id).await; + let (mut sender, receiver) = socket.split(); - let (mut sender, mut receiver) = socket.split(); - - // Reject room-full before subscribing so the error goes only to the joining socket. - { - let room = room_arc.lock().await; - if room.game.can_start() { - let err = serde_json::to_string(&ServerMessage::Error { - message: "Room is full".to_string(), - }).unwrap(); - let _ = sender.send(Message::Text(Utf8Bytes::from(err))).await; - return; - } + if try_join_room(&room).await.is_err() { + let err = serde_json::to_string(&ServerMessage::Error { + message: "Room is full".to_string(), + }) + .unwrap(); + let _ = sender.send(Message::Text(Utf8Bytes::from(err))).await; + return; } - // Add the player to the room and tell only this socket its assigned team. - let team = { - let mut room = room_arc.lock().await; - let team = room.game.add_player(player_id.clone(), preferred_team) - .unwrap_or(Team::Red); - if room.game.can_start() { - room.game.start(); - } - team - }; + let player_id = Uuid::new_v4().to_string(); + let team = register_player(&room, &player_id, preferred_team).await; - let joined_msg = serde_json::to_string(&ServerMessage::Joined { + let joined = serde_json::to_string(&ServerMessage::Joined { room: room_id.clone(), team, - }).unwrap(); - let _ = sender.send(Message::Text(Utf8Bytes::from(joined_msg))).await; + }) + .unwrap(); + let _ = sender.send(Message::Text(Utf8Bytes::from(joined))).await; - // Subscribe to broadcast and spawn the forwarding task. - let tx = { - let room = room_arc.lock().await; - room.tx.clone() - }; - let mut rx = tx.subscribe(); + let tx = { room.lock().await.tx.clone() }; - let send_task = tokio::spawn(async move { + let send_task = spawn_forwarder(sender, tx.subscribe()); + broadcast_room_state(&room, &tx).await; + + let recv_task = spawn_message_handler(room.clone(), player_id.clone(), tx, receiver); + + tokio::select! { + _ = send_task => {} + _ = recv_task => {} + } + + remove_player(&room, &player_id).await; +} + +async fn try_join_room(room: &Arc>) -> Result<(), ()> { + let room_guard = room.lock().await; + if room_guard.game.can_start() { + return Err(()); + } + Ok(()) +} + +async fn register_player( + room: &Arc>, + player_id: &str, + preferred_team: Option, +) -> Team { + let mut room_guard = room.lock().await; + let team = room_guard + .game + .add_player(player_id.to_string(), preferred_team) + .unwrap_or(Team::Red); + if room_guard.game.can_start() { + room_guard.game.start(); + } + team +} + +fn spawn_forwarder( + mut sender: WsSender, + mut rx: tokio::sync::broadcast::Receiver, +) -> tokio::task::JoinHandle<()> { + tokio::spawn(async move { loop { match rx.recv().await { Ok(msg) => { @@ -151,49 +184,59 @@ async fn handle_socket(socket: WebSocket, state: Arc, room_id: String, Err(_) => break, } } - }); + }) +} - // Broadcast waiting/game state to everyone in the room. - { - let room = room_arc.lock().await; - if room.game.can_start() { - let state_msg = room.game.game_state_message(); - let _ = room.tx.send(state_msg); - } else { - let _ = room.tx.send(ServerMessage::Waiting { message: "Waiting for other player".to_string() }); +async fn broadcast_room_state( + room: &Arc>, + tx: &tokio::sync::broadcast::Sender, +) { + let room_guard = room.lock().await; + let msg = if room_guard.game.can_start() { + room_guard.game.game_state_message() + } else { + ServerMessage::Waiting { + message: "Waiting for other player".to_string(), } - } + }; + let _ = tx.send(msg); +} - let recv_room = room_arc.clone(); - let recv_id = player_id_for_recv; - let recv_task = tokio::spawn(async move { +fn spawn_message_handler( + room: Arc>, + player_id: String, + tx: tokio::sync::broadcast::Sender, + mut receiver: WsReceiver, +) -> tokio::task::JoinHandle<()> { + tokio::spawn(async move { while let Some(Ok(msg)) = receiver.next().await { let Message::Text(text) = msg else { continue; }; let text_ref = text.as_str(); let parsed: Result = serde_json::from_str(text_ref); match parsed { Ok(ClientMessage::Throw { broom_x, broom_y, weight, curl, friction }) => { - let room = recv_room.clone(); - let mut room = room.lock().await; - if room.game.current_team_for_player(&recv_id) != Some(team) { - let _ = room.tx.send(ServerMessage::Error { message: "Not your turn".to_string() }); - continue; - } - match room.game.handle_throw(&recv_id, broom_x, broom_y, weight, curl, friction) { - Ok(path) => { - room.tx.send(ServerMessage::Trajectory { path }).ok(); - room.game.finish_simulation(); - if let Some(scored) = room.game.take_last_end_scored() { - room.tx.send(scored).ok(); + let mut room_guard = room.lock().await; + match room_guard + .game + .process_throw(&player_id, broom_x, broom_y, weight, curl, friction) + { + Ok(ThrowOutcome { + trajectory, + end_scored, + state_message, + game_over, + }) => { + let _ = tx.send(ServerMessage::Trajectory { path: trajectory }); + if let Some(scored) = end_scored { + let _ = tx.send(scored); } - let after = room.game.game_state_message(); - room.tx.send(after).ok(); - if room.game.phase() == GamePhase::GameComplete { - room.tx.send(room.game.game_over_message()).ok(); + let _ = tx.send(state_message); + if let Some(over) = game_over { + let _ = tx.send(over); } } Err(e) => { - room.tx.send(ServerMessage::Error { message: e }).ok(); + let _ = tx.send(ServerMessage::Error { message: e }); } } } @@ -202,16 +245,10 @@ async fn handle_socket(socket: WebSocket, state: Arc, room_id: String, } } } - }); - - tokio::select! { - _ = send_task => {} - _ = recv_task => {} - } - - // On disconnect, free the team slot so refreshes and new tabs can rejoin. - { - let mut room = room_arc.lock().await; - let _ = room.game.remove_player(&player_id); - } + }) +} + +async fn remove_player(room: &Arc>, player_id: &str) { + let mut room_guard = room.lock().await; + room_guard.game.remove_player(player_id); }