diff --git a/src/web/manager/websocket.rs b/src/web/manager/websocket.rs index 4499f2c..c0af4c8 100644 --- a/src/web/manager/websocket.rs +++ b/src/web/manager/websocket.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use std::net::IpAddr; use std::sync::Arc; -use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering}; use std::time::Duration; use tokio::sync::OwnedSemaphorePermit; @@ -10,7 +10,7 @@ use tokio_util::sync::CancellationToken; use super::{ManagerError, ProfileKey, WebProcessRuntime, WebSocketBudgetLease}; /// One process-owned WebSocket carrier class used for eviction priority. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] pub(crate) enum WebSocketKind { /// One connection multiplexes every logical stream in a session. Multiplex, @@ -18,23 +18,42 @@ pub(crate) enum WebSocketKind { Lane(u32), } +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +struct WebSocketClaimKey { + session_hash: super::TokenHash, + kind: WebSocketKind, +} + +#[repr(u8)] +enum WebSocketPhase { + Claimed, + Upgraded, + Active, + Closing, +} + pub(super) struct WebSocketEntry { id: u64, owner: ProfileKey, session_id: u64, + claim: WebSocketClaimKey, client_ip: IpAddr, kind: WebSocketKind, liveness_interval_ms: u64, created_tick: u64, last_peer_tick: AtomicU64, last_progress_tick: AtomicU64, - opened: AtomicBool, + phase: AtomicU8, + closing: AtomicBool, cancel: CancellationToken, + released: CancellationToken, } #[derive(Default)] pub(super) struct WebSocketRegistry { entries: HashMap>, + claims: HashMap, + evictions_in_flight: usize, closed: bool, } @@ -64,10 +83,20 @@ impl WebSocketConnection { /// Marks successful ownership transfer from HTTP to the WebSocket codec. pub(crate) fn mark_opened(&self) { - self.entry.opened.store(true, Ordering::Release); + self.entry + .phase + .store(WebSocketPhase::Upgraded as u8, Ordering::Release); self.mark_progress(); } + /// Marks the first validated carrier binary message as active progress. + pub(crate) fn mark_active(&self) { + self.entry + .phase + .store(WebSocketPhase::Active as u8, Ordering::Release); + self.mark_peer_activity(); + } + /// Refreshes the peer-liveness deadline after any received WebSocket message. pub(crate) fn mark_peer_activity(&self) { if let Some(runtime) = self.runtime.upgrade() { @@ -90,7 +119,16 @@ impl WebSocketConnection { impl Drop for WebSocketConnection { fn drop(&mut self) { if let Some(runtime) = self.runtime.upgrade() { - runtime.websockets.lock().entries.remove(&self.entry.id); + let mut registry = runtime.websockets.lock(); + registry.entries.remove(&self.entry.id); + if registry.claims.get(&self.entry.claim) == Some(&self.entry.id) { + registry.claims.remove(&self.entry.claim); + } + if self.entry.closing.load(Ordering::Acquire) { + registry.evictions_in_flight = registry.evictions_in_flight.saturating_sub(1); + } + drop(registry); + self.entry.released.cancel(); drop(self.base_budget.take()); drop(self.slot.take()); runtime.websocket_notify.notify_waiters(); @@ -103,81 +141,126 @@ pub(super) async fn admit( runtime: &Arc, owner: ProfileKey, session_id: u64, + session_hash: super::TokenHash, client_ip: IpAddr, kind: WebSocketKind, base_bytes: usize, liveness_interval: Duration, eviction_timeout: Duration, + parent_cancellation: CancellationToken, ) -> Result { let liveness_interval_ms = liveness_interval.as_millis().min(u128::from(u64::MAX)) as u64; - if let Some(connection) = try_admit( + match try_admit( runtime, owner, session_id, + session_hash, client_ip, kind, base_bytes, liveness_interval_ms, + &parent_cancellation, ) { - return Ok(connection); + Ok(connection) => return Ok(connection), + Err(TryAdmitError::Conflict) => return Err(ManagerError::Concurrent), + Err(TryAdmitError::Closed) => return Err(ManagerError::Closed), + Err(TryAdmitError::Capacity) => {} } - let Some(victim) = select_victim(runtime, owner, session_id, client_ip, None) else { + let Some(victim) = select_victim(runtime, owner, session_id, client_ip, None, true) else { runtime.record_limit_hit(); return Err(ManagerError::Limit); }; - let released = runtime.websocket_notify.notified(); + let released = victim.released.cancelled(); victim.cancel.cancel(); let _ = tokio::time::timeout(eviction_timeout, released).await; - try_admit( + match try_admit( runtime, owner, session_id, + session_hash, client_ip, kind, base_bytes, liveness_interval_ms, - ) - .ok_or_else(|| { - runtime.record_limit_hit(); - ManagerError::Limit - }) + &parent_cancellation, + ) { + Ok(connection) => Ok(connection), + Err(TryAdmitError::Conflict) => Err(ManagerError::Concurrent), + Err(TryAdmitError::Closed) => Err(ManagerError::Closed), + Err(TryAdmitError::Capacity) => { + runtime.record_limit_hit(); + Err(ManagerError::Limit) + } + } +} + +enum TryAdmitError { + Capacity, + Conflict, + Closed, } fn try_admit( runtime: &Arc, owner: ProfileKey, session_id: u64, + session_hash: super::TokenHash, client_ip: IpAddr, kind: WebSocketKind, base_bytes: usize, liveness_interval_ms: u64, -) -> Option { + parent_cancellation: &CancellationToken, +) -> Result { + let claim = WebSocketClaimKey { session_hash, kind }; + { + let registry = runtime.websockets.lock(); + if registry.closed { + return Err(TryAdmitError::Closed); + } + if registry.claims.contains_key(&claim) { + return Err(TryAdmitError::Conflict); + } + } let slot = Arc::clone(&runtime.websocket_connections) .try_acquire_owned() - .ok()?; - let base_budget = runtime.try_websocket_base_budget(owner, base_bytes)?; - let id = runtime.websocket_next_id.fetch_add(1, Ordering::Relaxed); + .map_err(|_| TryAdmitError::Capacity)?; + let base_budget = runtime + .try_websocket_base_budget(owner, base_bytes) + .ok_or(TryAdmitError::Capacity)?; + let id = runtime + .websocket_next_id + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| { + value.checked_add(1) + }) + .map_err(|_| TryAdmitError::Capacity)?; let now = runtime.websocket_tick(); let entry = Arc::new(WebSocketEntry { id, owner, session_id, + claim, client_ip, kind, liveness_interval_ms, created_tick: now, last_peer_tick: AtomicU64::new(now), last_progress_tick: AtomicU64::new(now), - opened: AtomicBool::new(false), - cancel: CancellationToken::new(), + phase: AtomicU8::new(WebSocketPhase::Claimed as u8), + closing: AtomicBool::new(false), + cancel: parent_cancellation.child_token(), + released: CancellationToken::new(), }); let mut registry = runtime.websockets.lock(); if registry.closed { - return None; + return Err(TryAdmitError::Closed); } + if registry.claims.contains_key(&claim) { + return Err(TryAdmitError::Conflict); + } + registry.claims.insert(claim, id); registry.entries.insert(id, Arc::clone(&entry)); drop(registry); - Some(WebSocketConnection { + Ok(WebSocketConnection { runtime: Arc::downgrade(runtime), entry, slot: Some(slot), diff --git a/src/web/session.rs b/src/web/session.rs index ba4a02a..2d385a4 100644 --- a/src/web/session.rs +++ b/src/web/session.rs @@ -81,6 +81,8 @@ struct DownBatch { struct CarrierLane { instance: u64, + pending_bytes: usize, + pending_items: usize, pending_frames: VecDeque, pending_windows: HashMap, unacked: Option, @@ -96,6 +98,8 @@ impl CarrierLane { fn new(instance: u64) -> Self { Self { instance, + pending_bytes: 0, + pending_items: 0, pending_frames: VecDeque::new(), pending_windows: HashMap::new(), unacked: None, @@ -124,7 +128,6 @@ struct SessionState { last_up_sequence: u64, last_up_digest: TokenHash, carrier_lanes: HashMap, - lane_open_claims: HashSet, lane_open_waits: usize, next_lane_instance: u64, websocket_lane_reservations: HashMap, @@ -231,7 +234,6 @@ impl WebSession { last_up_sequence: 0, last_up_digest: [0; 32], carrier_lanes, - lane_open_claims: HashSet::new(), lane_open_waits: 0, next_lane_instance, websocket_lane_reservations: HashMap::new(), @@ -458,17 +460,6 @@ impl WebSession { .map(|manager| manager.budget_notify()) } - fn release_stream_reservation(&self, peer_port: u16) { - let removed = self.state.lock().active_peer_ports.remove(&peer_port); - if removed && let Some(manager) = self.manager.upgrade() { - manager.release_stream( - self.profile_key, - self.client_ip, - self.profile.public_addr, - peer_port, - ); - } - } } fn inbound_queue_cost(queue: &VecDeque) -> (usize, usize) { diff --git a/src/web/session/lanes.rs b/src/web/session/lanes.rs index 21de257..2620ab1 100644 --- a/src/web/session/lanes.rs +++ b/src/web/session/lanes.rs @@ -15,6 +15,12 @@ use crate::web::frame::{self, Frame, FrameType}; use crate::web::manager::{ManagerError, TokenHash}; impl WebSession { + /// Classifies control and pre-OPEN polls for their reserved handler pool. + pub(crate) fn lane_poll_is_auxiliary(&self, lane_id: u32) -> bool { + let state = self.state.lock(); + lane_id == 0 || !state.carrier_lanes.contains_key(&lane_id) + } + /// Applies one exactly-once uplink batch to an independent HTTPS lane. pub(crate) fn process_up_lane( self: &Arc, @@ -53,7 +59,8 @@ impl WebSession { } self.ensure_carrier_active_locked(&state)?; state.last_activity = Instant::now(); - if !state.carrier_lanes.contains_key(&lane_id) { + let new_lane = !state.carrier_lanes.contains_key(&lane_id); + if new_lane { if lane_id != 0 && frames .first() @@ -71,18 +78,20 @@ impl WebSession { self.close(); return Err(ManagerError::Protocol); } - if insert_carrier_lane(&mut state, lane_id).is_none() { - drop(state); - self.close(); - return Err(ManagerError::Protocol); + if state.carrier_lanes.len() + >= self.profile.max_streams_per_session.saturating_add(1) + { + return Err(ManagerError::Limit); } } - let lane = state + let (last_sequence, last_digest, up_active) = state .carrier_lanes - .get_mut(&lane_id) - .ok_or(ManagerError::Protocol)?; - if sequence == lane.last_up_sequence && sequence != 0 { - return if bool::from(lane.last_up_digest.ct_eq(&digest)) { + .get(&lane_id) + .map_or((0, [0; 32], false), |lane| { + (lane.last_up_sequence, lane.last_up_digest, lane.up_active) + }); + if sequence == last_sequence && sequence != 0 { + return if bool::from(last_digest.ct_eq(&digest)) { Ok(sequence) } else { drop(state); @@ -90,15 +99,14 @@ impl WebSession { Err(ManagerError::Protocol) }; } - if sequence == 0 || sequence != lane.last_up_sequence.saturating_add(1) { + if sequence == 0 || sequence != last_sequence.saturating_add(1) { drop(state); self.close(); return Err(ManagerError::Protocol); } - if lane.up_active { + if up_active { return Err(ManagerError::Concurrent); } - lane.up_active = true; if !validate_batch(&state, &frames) { drop(state); self.close(); @@ -116,6 +124,17 @@ impl WebSession { } return Err(ManagerError::Backpressure); } + if new_lane && insert_carrier_lane(&mut state, lane_id).is_none() { + self.release_locked(&mut state, reserve_bytes, reserve_items, false); + drop(state); + self.close(); + return Err(ManagerError::Protocol); + } + let Some(lane) = state.carrier_lanes.get_mut(&lane_id) else { + self.release_locked(&mut state, reserve_bytes, reserve_items, false); + return Err(ManagerError::Closed); + }; + lane.up_active = true; let mut unused_bytes = reserve_bytes; let mut unused_items = reserve_items; let applied = self.apply_batch_locked( @@ -150,6 +169,7 @@ impl WebSession { if committed { self.finish_carrier_commit(); } + self.lane_open_notify.notify_waiters(); for completion in opened { self.spawn_stream(completion, false); } @@ -168,14 +188,26 @@ impl WebSession { if !self.carrier().uses_lanes() || lane_id > frame::MAX_STREAM_ID { return Err(ManagerError::Protocol); } - let (epoch, notify) = { + if !self.wait_for_lane_open(lane_id, cursor).await? { + return Ok(PollResult { + body: Bytes::new(), + next_cursor: cursor, + lane_closed: false, + }); + } + let (instance, epoch, notify) = { let mut state = self.state.lock(); if state.closed { return Err(ManagerError::Closed); } + state.last_activity = Instant::now(); let acknowledged = { let Some(lane) = state.carrier_lanes.get_mut(&lane_id) else { - return Err(ManagerError::Protocol); + return Ok(PollResult { + body: Bytes::new(), + next_cursor: cursor, + lane_closed: true, + }); }; if let Some(unacked) = &lane.unacked { if cursor == unacked.base_cursor { @@ -201,6 +233,14 @@ impl WebSession { } }; if let Some(batch) = acknowledged { + if let Some(lane) = state.carrier_lanes.get_mut(&lane_id) { + lane.pending_bytes = lane + .pending_bytes + .saturating_sub(batch.data_bytes.saturating_add(batch.control_bytes)); + lane.pending_items = lane + .pending_items + .saturating_sub(batch.data_items.saturating_add(batch.control_items)); + } self.release_locked(&mut state, batch.data_bytes, batch.data_items, false); self.release_locked(&mut state, batch.control_bytes, batch.control_items, true); if let Some(stream) = state.streams.get_mut(&lane_id) @@ -209,13 +249,12 @@ impl WebSession { waker.wake(); } } - state.last_activity = Instant::now(); let lane = state .carrier_lanes .get_mut(&lane_id) .ok_or(ManagerError::Protocol)?; lane.down_epoch = lane.down_epoch.wrapping_add(1).max(1); - (lane.down_epoch, Arc::clone(&lane.notify)) + (lane.instance, lane.down_epoch, Arc::clone(&lane.notify)) }; notify.notify_waiters(); @@ -235,7 +274,7 @@ impl WebSession { lane_closed: true, }); }; - if lane.down_epoch != epoch { + if lane.instance != instance || lane.down_epoch != epoch { return Ok(PollResult { body: Bytes::new(), next_cursor: cursor, @@ -304,7 +343,7 @@ impl WebSession { if state .carrier_lanes .get(&lane_id) - .is_some_and(|lane| lane.down_epoch == epoch) + .is_some_and(|lane| lane.instance == instance && lane.down_epoch == epoch) { state.last_activity = Instant::now(); } @@ -317,6 +356,60 @@ impl WebSession { } } + async fn wait_for_lane_open( + &self, + lane_id: u32, + cursor: u64, + ) -> Result { + let wait = { + let mut state = self.state.lock(); + if state.closed { + return Err(ManagerError::Closed); + } + if state.carrier_lanes.contains_key(&lane_id) { + return Ok(true); + } + if cursor != 0 || lane_id == 0 { + drop(state); + self.close(); + return Err(ManagerError::Protocol); + } + if state.closed_streams.contains(&lane_id) + || state.closing_streams.contains_key(&lane_id) + { + return Ok(true); + } + if state.lane_open_waits >= self.limits.max_lane_open_waits_per_session { + return Err(ManagerError::Limit); + } + state.lane_open_waits += 1; + LaneOpenWaitGuard { session: self } + }; + let deadline = Duration::from_secs(self.timeouts.lane_open_wait_secs); + let opened = tokio::time::timeout(deadline, async { + loop { + let notified = self.lane_open_notify.notified(); + { + let state = self.state.lock(); + if state.closed + || state.carrier_lanes.contains_key(&lane_id) + || state.closed_streams.contains(&lane_id) + || state.closing_streams.contains_key(&lane_id) + { + return state.carrier_lanes.contains_key(&lane_id) + || state.closed_streams.contains(&lane_id) + || state.closing_streams.contains_key(&lane_id); + } + } + notified.await; + } + }) + .await + .unwrap_or(false); + drop(wait); + Ok(opened) + } + pub(super) fn queue_lane_frame_locked( &self, state: &mut SessionState, @@ -363,6 +456,15 @@ impl WebSession { <= self.limits.max_frame_payload_bytes }); if can_coalesce { + if state.carrier_lanes.get(&stream_id).is_none_or(|lane| { + lane.pending_bytes + > self + .limits + .pending_bytes_per_lane + .saturating_sub(payload.len()) + }) { + return false; + } if !self.reserve_locked(state, payload.len(), 0, PendingClass::Downlink) { return false; } @@ -378,6 +480,7 @@ impl WebSession { last.cost += payload.len(); let payload_len = (last.encoded.len() - frame::HEADER_BYTES) as u32; last.encoded[4..8].copy_from_slice(&payload_len.to_be_bytes()); + lane.pending_bytes += payload.len(); lane.notify.notify_waiters(); return true; } @@ -387,6 +490,14 @@ impl WebSession { } else { PendingClass::Downlink }; + if state.carrier_lanes.get(&stream_id).is_none_or(|lane| { + lane.pending_bytes + > self.limits.pending_bytes_per_lane.saturating_sub(cost) + || lane.pending_items + >= self.limits.pending_items_per_lane + }) { + return false; + } if !self.reserve_locked(state, cost, 1, class) { return false; } @@ -409,6 +520,8 @@ impl WebSession { control, cost, }); + lane.pending_bytes += cost; + lane.pending_items += 1; if frame_type == FrameType::Window { lane.pending_windows.insert(stream_id, index); } @@ -455,6 +568,18 @@ impl WebSession { } self.release_locked(state, data_bytes, data_items, false); self.release_locked(state, control_bytes, control_items, true); + self.lane_open_notify.notify_waiters(); + } +} + +struct LaneOpenWaitGuard<'a> { + session: &'a WebSession, +} + +impl Drop for LaneOpenWaitGuard<'_> { + fn drop(&mut self) { + let mut state = self.session.state.lock(); + state.lane_open_waits = state.lane_open_waits.saturating_sub(1); } } diff --git a/src/web/session/websocket.rs b/src/web/session/websocket.rs index dda7ac5..b16df8a 100644 --- a/src/web/session/websocket.rs +++ b/src/web/session/websocket.rs @@ -100,6 +100,8 @@ impl WebSession { ); return Err(ManagerError::Protocol); } + drop(state); + self.lane_open_notify.notify_waiters(); Ok(WebSocketLaneReservation { session: Arc::clone(self), lane_id, @@ -229,6 +231,7 @@ impl WebSession { let mut state = self.state.lock(); let reserved = state.websocket_lane_reservations.remove(&lane_id); if let Some(stream) = state.streams.remove(&lane_id) { + state.closing_streams.insert(lane_id, stream.instance); let (bytes, items) = inbound_queue_cost(&stream.inbound); self.release_locked(&mut state, bytes, items, false); if let Some(waker) = stream.read_waker { @@ -245,6 +248,7 @@ impl WebSession { if let Some(peer_port) = reserved { self.release_websocket_lane_reservation(lane_id, peer_port); } + self.lane_open_notify.notify_waiters(); } fn release_websocket_lane_reservation(&self, lane_id: u32, peer_port: u16) {