mod protocol; mod physics; mod game; use axum::{ extract::{Query, State, WebSocketUpgrade}, http::StatusCode, response::IntoResponse, routing::{get, post}, Router, }; 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; use tokio::sync::Mutex; use tracing::{info, warn}; use uuid::Uuid; 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>>>>, } impl AppState { fn new() -> Self { Self { rooms: Arc::new(Mutex::new(HashMap::new())), } } 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(Mutex::new(Room::new(room_id))); rooms.insert(room_id.to_string(), room.clone()); room } } #[derive(Deserialize)] struct RoomQuery { room: String, } #[derive(Serialize)] struct NewRoomResponse { room: String, } fn generate_room_code() -> String { let s = Uuid::new_v4().to_string().replace("-", ""); s.chars().take(6).collect::().to_uppercase() } #[tokio::main] async fn main() { tracing_subscriber::fmt::init(); let state = Arc::new(AppState::new()); let app = Router::new() .route("/", get(health)) .route("/room", post(new_room)) .route("/ws", get(ws_handler)) .with_state(state); let listener = tokio::net::TcpListener::bind("0.0.0.0:3000").await.unwrap(); info!("Server listening on 0.0.0.0:3000"); axum::serve(listener, app).await.unwrap(); } async fn health() -> impl IntoResponse { "curltastic ok" } async fn new_room(State(state): State>) -> impl IntoResponse { let room_id = generate_room_code(); state.get_or_create_room(&room_id).await; (StatusCode::OK, axum::Json(NewRoomResponse { room: room_id })) } async fn ws_handler( ws: WebSocketUpgrade, Query(query): Query, State(state): State>, ) -> impl IntoResponse { ws.on_upgrade(move |socket| handle_socket(socket, state, query.room)) } async fn handle_socket( socket: WebSocket, state: Arc, room_id: String, ) { let room = state.get_or_create_room(&room_id).await; let (mut sender, receiver) = socket.split(); { let mut room_guard = room.lock().await; if matches!(room_guard.game.game_state_message(), ServerMessage::GameState { phase: Phase::Waiting, .. }) { room_guard.game.start(); } } let joined = serde_json::to_string(&ServerMessage::Joined { room: room_id.clone(), }) .unwrap(); if sender.send(Message::Text(Utf8Bytes::from(joined))).await.is_err() { return; } let tx = { room.lock().await.tx.clone() }; let send_task = spawn_forwarder(sender, tx.subscribe()); broadcast_room_state(&room, &tx).await; let recv_task = spawn_message_handler(room.clone(), tx, receiver); tokio::select! { _ = send_task => {} _ = recv_task => {} } } 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) => { let text = match serde_json::to_string(&msg) { Ok(t) => t, Err(_) => continue, }; if sender.send(Message::Text(Utf8Bytes::from(text))).await.is_err() { break; } } Err(_) => break, } } }) } async fn broadcast_room_state( room: &Arc>, tx: &tokio::sync::broadcast::Sender, ) { let room_guard = room.lock().await; let _ = tx.send(room_guard.game.game_state_message()); } fn spawn_message_handler( room: Arc>, 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 { team, broom_x, broom_y, velocity, curl, friction, }) => { let mut room_guard = room.lock().await; match room_guard .game .process_throw(team, broom_x, broom_y, velocity, curl, friction) { Ok(ThrowOutcome { trajectories, state_message, game_over, }) => { let _ = tx.send(ServerMessage::Trajectories { stones: trajectories, }); let _ = tx.send(state_message); if let Some(over) = game_over { let _ = tx.send(over); } } Err(e) => { let _ = tx.send(ServerMessage::Error { message: e }); } } } Err(e) => { warn!("Invalid message: {}", e); } } } }) }