diff --git a/crates/adapters/socketio/src/handlers.rs b/crates/adapters/socketio/src/handlers.rs index 8d0515e..eb260b9 100644 --- a/crates/adapters/socketio/src/handlers.rs +++ b/crates/adapters/socketio/src/handlers.rs @@ -3,7 +3,7 @@ use std::sync::Arc; use api_types::{self, PixelUpdatePayload}; use application::AppState; use application::canvas::place_pixel; -use domain::{BroadcastEvent, BroadcastSubscription, Color, Position}; +use domain::{BroadcastEvent, BroadcastSubscription, Color, Position, UserId}; use socketioxide::extract::{Data, SocketRef}; use tracing::info; @@ -28,8 +28,8 @@ fn send_canvas_state(socket: &SocketRef, state: &AppState) { } fn register_soldier(state: &AppState, socket: &SocketRef) { - let socket_id = socket.id.to_string(); - application::soldiers::connect::execute(state, socket_id); + let user_id = UserId::new(socket.id.to_string()); + application::soldiers::connect::execute(state, user_id); } fn spawn_broadcast_forwarder(socket: SocketRef, mut subscription: BroadcastSubscription) { @@ -71,9 +71,9 @@ fn handle_place_pixel(socket: &SocketRef, state: &AppState, payload: PixelUpdate info!("Received pixel update: {position} color={}", color.as_u32()); - let socket_id = socket.id.to_string(); + let user_id = UserId::new(socket.id.to_string()); let command = place_pixel::Command { - user_id: &socket_id, + user_id: &user_id, position, color, }; @@ -97,8 +97,8 @@ fn register_disconnect_handler(socket: &SocketRef, state: Arc) { fn handle_disconnect(socket: &SocketRef, state: &AppState) { info!("Socket disconnected: {:?}", socket.id); - let socket_id = socket.id.to_string(); - application::soldiers::disconnect::execute(state, &socket_id); + let user_id = UserId::new(socket.id.to_string()); + application::soldiers::disconnect::execute(state, &user_id); } fn emit_error(socket: &SocketRef, message: &str) { diff --git a/crates/adapters/websocket/src/handler.rs b/crates/adapters/websocket/src/handler.rs index a33d850..218f637 100644 --- a/crates/adapters/websocket/src/handler.rs +++ b/crates/adapters/websocket/src/handler.rs @@ -1,11 +1,12 @@ use std::io::Write; use std::sync::Arc; use std::sync::atomic::Ordering; +use std::time::Duration; use axum::extract::State; use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; use axum::response::IntoResponse; -use domain::{BroadcastEvent, BroadcastSubscription, Color, Position}; +use domain::{BroadcastEvent, BroadcastSubscription, Color, Position, UserId}; use flate2::Compression; use flate2::write::GzEncoder; use futures::{SinkExt, StreamExt, stream::SplitSink}; @@ -29,10 +30,12 @@ pub async fn ws_upgrade( async fn handle_connection(socket: WebSocket, state: Arc) { let (mut sender, mut receiver) = socket.split(); - let connection_id = state - .connection_counter - .fetch_add(1, Ordering::Relaxed) - .to_string(); + let connection_id = UserId::new( + state + .connection_counter + .fetch_add(1, Ordering::Relaxed) + .to_string(), + ); info!("WebSocket connected: {connection_id}"); @@ -105,11 +108,16 @@ fn gzip_compress(data: &[u8]) -> Option> { encoder.finish().ok() } +const PING_INTERVAL: Duration = Duration::from_secs(30); + async fn run_send_loop( mut sender: WsSender, mut subscription: BroadcastSubscription, mut error_receiver: mpsc::UnboundedReceiver, ) { + let mut ping_interval = tokio::time::interval(PING_INTERVAL); + ping_interval.tick().await; + loop { tokio::select! { Some(event) = subscription.recv() => { @@ -123,6 +131,11 @@ async fn run_send_loop( break; } } + _ = ping_interval.tick() => { + if sender.send(Message::Ping(Vec::new().into())).await.is_err() { + break; + } + } else => break, } } @@ -143,7 +156,7 @@ fn serialize_broadcast_event(event: &BroadcastEvent) -> Option { fn handle_client_message( state: &AppState, error_sender: &mpsc::UnboundedSender, - connection_id: &str, + connection_id: &UserId, text: &str, ) { let Ok(message) = serde_json::from_str::(text) else { @@ -160,7 +173,7 @@ fn handle_client_message( fn handle_place_pixel( state: &AppState, error_sender: &mpsc::UnboundedSender, - connection_id: &str, + connection_id: &UserId, x: u32, y: u32, color: u32, diff --git a/crates/application/src/canvas/place_pixel.rs b/crates/application/src/canvas/place_pixel.rs index 871ff35..906f24a 100644 --- a/crates/application/src/canvas/place_pixel.rs +++ b/crates/application/src/canvas/place_pixel.rs @@ -1,9 +1,9 @@ -use domain::{BroadcastEvent, Color, PixelUpdate, Position}; +use domain::{BroadcastEvent, Color, PixelUpdate, Position, UserId}; use crate::{AppState, ApplicationError}; pub struct Command<'a> { - pub user_id: &'a str, + pub user_id: &'a UserId, pub position: Position, pub color: Color, } diff --git a/crates/application/src/soldiers/connect.rs b/crates/application/src/soldiers/connect.rs index 6725374..1d356b2 100644 --- a/crates/application/src/soldiers/connect.rs +++ b/crates/application/src/soldiers/connect.rs @@ -1,8 +1,8 @@ -use domain::BroadcastEvent; +use domain::{BroadcastEvent, UserId}; use crate::AppState; -pub fn execute(state: &AppState, user_id: String) -> usize { +pub fn execute(state: &AppState, user_id: UserId) -> usize { let count = state.soldiers().add(user_id); state .broadcaster() diff --git a/crates/application/src/soldiers/disconnect.rs b/crates/application/src/soldiers/disconnect.rs index 89b3b8f..c8ba0f0 100644 --- a/crates/application/src/soldiers/disconnect.rs +++ b/crates/application/src/soldiers/disconnect.rs @@ -1,8 +1,8 @@ -use domain::BroadcastEvent; +use domain::{BroadcastEvent, UserId}; use crate::AppState; -pub fn execute(state: &AppState, user_id: &str) -> usize { +pub fn execute(state: &AppState, user_id: &UserId) -> usize { state.cooldowns().remove(user_id); let count = state.soldiers().remove(user_id); state diff --git a/crates/application/src/state.rs b/crates/application/src/state.rs index 77fae3c..0a348bc 100644 --- a/crates/application/src/state.rs +++ b/crates/application/src/state.rs @@ -7,7 +7,7 @@ use std::time::{Duration, Instant}; use domain::ports::{ BroadcastReceiverInner, BroadcastSubscription, CanvasPersistence, CanvasStore, EventBroadcaster, }; -use domain::{BroadcastEvent, Canvas, Color, DomainError, Position}; +use domain::{BroadcastEvent, Canvas, Color, DomainError, Position, UserId}; use tokio::sync::broadcast; use tracing::warn; @@ -145,7 +145,7 @@ impl CanvasStore for InMemoryCanvasStore { } pub struct CooldownTracker { - entries: Mutex>, + entries: Mutex>, cooldown: Duration, } @@ -157,7 +157,7 @@ impl CooldownTracker { } } - pub fn is_on_cooldown(&self, user_id: &str) -> bool { + pub fn is_on_cooldown(&self, user_id: &UserId) -> bool { let entries = acquire_lock(&self.entries); entries .get(user_id) @@ -165,17 +165,17 @@ impl CooldownTracker { .unwrap_or(false) } - pub fn record(&self, user_id: &str) { - acquire_lock(&self.entries).insert(user_id.to_string(), Instant::now()); + pub fn record(&self, user_id: &UserId) { + acquire_lock(&self.entries).insert(user_id.clone(), Instant::now()); } - pub fn remove(&self, user_id: &str) { + pub fn remove(&self, user_id: &UserId) { acquire_lock(&self.entries).remove(user_id); } } pub struct SoldierTracker { - connected: Mutex>, + connected: Mutex>, } impl SoldierTracker { @@ -185,13 +185,13 @@ impl SoldierTracker { } } - pub fn add(&self, user_id: String) -> usize { + pub fn add(&self, user_id: UserId) -> usize { let mut connected = acquire_lock(&self.connected); connected.insert(user_id); connected.len() } - pub fn remove(&self, user_id: &str) -> usize { + pub fn remove(&self, user_id: &UserId) -> usize { let mut connected = acquire_lock(&self.connected); connected.remove(user_id); connected.len() diff --git a/crates/application/tests/place_pixel.rs b/crates/application/tests/place_pixel.rs index 72b18f0..95bd55f 100644 --- a/crates/application/tests/place_pixel.rs +++ b/crates/application/tests/place_pixel.rs @@ -1,14 +1,18 @@ mod common; use application::canvas::place_pixel::{Command, Outcome}; -use domain::{BroadcastEvent, Color, Position}; +use domain::{BroadcastEvent, Color, Position, UserId}; + +fn uid(id: &str) -> UserId { + UserId::new(id.to_string()) +} macro_rules! place { ($state:expr, $user:expr, $x:expr, $y:expr, $color:expr) => { application::canvas::place_pixel::execute( &$state, Command { - user_id: $user, + user_id: &uid($user), position: Position::new($x, $y), color: Color::new($color), }, diff --git a/crates/application/tests/snapshots.rs b/crates/application/tests/snapshots.rs index cacd493..6db454b 100644 --- a/crates/application/tests/snapshots.rs +++ b/crates/application/tests/snapshots.rs @@ -3,7 +3,7 @@ mod common; use std::sync::Arc; use application::AppState; -use domain::Color; +use domain::{Color, UserId}; fn state_with_persistence(persistence: common::FakePersistence) -> Arc { let spy = common::SpyBroadcaster::new(); @@ -22,10 +22,11 @@ fn save_snapshot_persists_current_canvas() { let saved_ref = persistence.saved_ref(); let state = state_with_persistence(persistence); + let user = UserId::new("user".to_string()); application::canvas::place_pixel::execute( &state, application::canvas::place_pixel::Command { - user_id: "user", + user_id: &user, position: domain::Position::new(0, 0), color: Color::new(0xFF), }, diff --git a/crates/application/tests/soldiers.rs b/crates/application/tests/soldiers.rs index f28e259..3bf6fea 100644 --- a/crates/application/tests/soldiers.rs +++ b/crates/application/tests/soldiers.rs @@ -1,40 +1,41 @@ mod common; -use domain::BroadcastEvent; +use domain::{BroadcastEvent, UserId}; + +fn uid(id: &str) -> UserId { + UserId::new(id.to_string()) +} #[test] fn connect_increments_count() { let (state, _) = common::test_state(); - assert_eq!( - application::soldiers::connect::execute(&state, "a".into()), - 1 - ); - assert_eq!( - application::soldiers::connect::execute(&state, "b".into()), - 2 - ); - assert_eq!( - application::soldiers::connect::execute(&state, "c".into()), - 3 - ); + assert_eq!(application::soldiers::connect::execute(&state, uid("a")), 1); + assert_eq!(application::soldiers::connect::execute(&state, uid("b")), 2); + assert_eq!(application::soldiers::connect::execute(&state, uid("c")), 3); } #[test] fn disconnect_decrements_count() { let (state, _) = common::test_state(); - application::soldiers::connect::execute(&state, "a".into()); - application::soldiers::connect::execute(&state, "b".into()); + application::soldiers::connect::execute(&state, uid("a")); + application::soldiers::connect::execute(&state, uid("b")); - assert_eq!(application::soldiers::disconnect::execute(&state, "a"), 1); - assert_eq!(application::soldiers::disconnect::execute(&state, "b"), 0); + assert_eq!( + application::soldiers::disconnect::execute(&state, &uid("a")), + 1 + ); + assert_eq!( + application::soldiers::disconnect::execute(&state, &uid("b")), + 0 + ); } #[test] fn disconnect_unknown_user_is_harmless() { let (state, _) = common::test_state(); - application::soldiers::connect::execute(&state, "a".into()); + application::soldiers::connect::execute(&state, uid("a")); assert_eq!( - application::soldiers::disconnect::execute(&state, "unknown"), + application::soldiers::disconnect::execute(&state, &uid("unknown")), 1 ); } @@ -42,7 +43,7 @@ fn disconnect_unknown_user_is_harmless() { #[test] fn connect_publishes_soldier_count() { let (state, spy) = common::test_state(); - application::soldiers::connect::execute(&state, "a".into()); + application::soldiers::connect::execute(&state, uid("a")); let events = spy.events(); assert_eq!(events.len(), 1); @@ -52,8 +53,8 @@ fn connect_publishes_soldier_count() { #[test] fn disconnect_publishes_soldier_count() { let (state, spy) = common::test_state(); - application::soldiers::connect::execute(&state, "a".into()); - application::soldiers::disconnect::execute(&state, "a"); + application::soldiers::connect::execute(&state, uid("a")); + application::soldiers::disconnect::execute(&state, &uid("a")); let events = spy.events(); assert_eq!(events.len(), 2); @@ -63,26 +64,27 @@ fn disconnect_publishes_soldier_count() { #[test] fn disconnect_clears_cooldown() { let (state, _) = common::test_state(); - application::soldiers::connect::execute(&state, "user-1".into()); + let user = uid("user-1"); + application::soldiers::connect::execute(&state, user.clone()); application::canvas::place_pixel::execute( &state, application::canvas::place_pixel::Command { - user_id: "user-1", + user_id: &user, position: domain::Position::new(0, 0), color: domain::Color::new(0xFF), }, ) .unwrap(); - application::soldiers::disconnect::execute(&state, "user-1"); + application::soldiers::disconnect::execute(&state, &user); // Reconnect with same ID — cooldown should be gone - application::soldiers::connect::execute(&state, "user-1".into()); + application::soldiers::connect::execute(&state, user.clone()); let result = application::canvas::place_pixel::execute( &state, application::canvas::place_pixel::Command { - user_id: "user-1", + user_id: &user, position: domain::Position::new(1, 1), color: domain::Color::new(0xAA), }, diff --git a/crates/domain/src/lib.rs b/crates/domain/src/lib.rs index f4fac02..033f5c7 100644 --- a/crates/domain/src/lib.rs +++ b/crates/domain/src/lib.rs @@ -8,6 +8,6 @@ pub use canvas::Canvas; pub use errors::DomainError; pub use events::BroadcastEvent; pub use ports::BroadcastSubscription; -pub use value_objects::{Color, PixelUpdate, Position}; +pub use value_objects::{Color, PixelUpdate, Position, UserId}; pub const COOLDOWN_MESSAGE: &str = "You can only place one pixel per minute"; diff --git a/crates/domain/src/value_objects.rs b/crates/domain/src/value_objects.rs index 603c72c..973bd1d 100644 --- a/crates/domain/src/value_objects.rs +++ b/crates/domain/src/value_objects.rs @@ -1,5 +1,24 @@ use std::fmt; +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct UserId(String); + +impl UserId { + pub fn new(id: String) -> Self { + Self(id) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl fmt::Display for UserId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.0) + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(transparent)] pub struct Color(u32);