diff --git a/src/rtc/coordinator/ws.rs b/src/rtc/coordinator/ws.rs index 30ba4a2..36efd24 100644 --- a/src/rtc/coordinator/ws.rs +++ b/src/rtc/coordinator/ws.rs @@ -442,6 +442,17 @@ mod tests { assert!(debug.contains("user")); } + #[test] + fn ws_auth_message_serializes_video_product() { + let auth = WsAuthMessage::video("jwt-token", ConnectUserDetails::new("agent")); + let json = serde_json::to_value(&auth).expect("serialize"); + assert_eq!(json["token"], "jwt-token"); + assert_eq!(json["user_details"]["id"], "agent"); + assert_eq!(json["products"][0], "video"); + // Optional user fields are omitted when unset. + assert!(json["user_details"].get("name").is_none()); + } + #[test] fn coordinator_message_limit_accepts_exact_and_rejects_oversized_input() { ensure_message_size("coordinator test message", 64, 64).expect("exact limit"); diff --git a/src/rtc/error.rs b/src/rtc/error.rs index e4734d4..fa53994 100644 --- a/src/rtc/error.rs +++ b/src/rtc/error.rs @@ -79,7 +79,7 @@ pub enum RtcError { #[error(transparent)] Join(#[from] SfuJoinError), - /// A client deadline elapsed waiting for the SFU (WS open or `JoinResponse`). + /// A client deadline elapsed, e.g. waiting for the SFU `JoinResponse`. #[error(transparent)] Timeout(#[from] SfuTimeoutError), @@ -333,9 +333,9 @@ impl SfuJoinError { } } -/// A client-side deadline elapsed waiting for the SFU (WS open or `JoinResponse`). +/// A client-side deadline elapsed. #[derive(Debug, Clone, thiserror::Error)] -#[error("sfu timeout waiting for {what} after {}ms", timeout.as_millis())] +#[error("timeout waiting for {what} after {}ms", timeout.as_millis())] pub struct SfuTimeoutError { /// What we were waiting for, e.g. `"join response"`. pub what: String, @@ -421,6 +421,28 @@ pub struct TwirpError { mod tests { use super::*; + #[test] + fn from_signal_error_maps_only_real_codes() { + // UNSPECIFIED (and absent) is success. + assert!(RtcError::from_signal_error(None).is_ok()); + assert!( + RtcError::from_signal_error(Some(models::Error { + code: models::ErrorCode::Unspecified as i32, + message: String::new(), + should_retry: false, + })) + .is_ok() + ); + // A real code becomes an error. + let err = RtcError::from_signal_error(Some(models::Error { + code: models::ErrorCode::ParticipantSignalLost as i32, + message: "boom".to_owned(), + should_retry: true, + })) + .expect_err("should be an error"); + assert!(matches!(err, RtcError::Signal { .. })); + } + #[test] fn join_error_codes_match_sfu() { assert!(is_join_error_code(ErrorCode::SfuFull as i32)); diff --git a/src/rtc/join/connection.rs b/src/rtc/join/connection.rs index fc237fe..eca997f 100644 --- a/src/rtc/join/connection.rs +++ b/src/rtc/join/connection.rs @@ -129,10 +129,10 @@ pub(super) fn register_connection_state( tracer.trace("connectionstatechange", json!(state.to_string())); if state == RTCPeerConnectionState::Connected { ever_connected.store(true, Ordering::SeqCst); - core.caps - .lock() - .unwrap_or_else(|e| e.into_inner()) - .reset_ice(); + let mut lifecycle = core.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + if lifecycle.generation == generation { + lifecycle.failure_limits.reset_ice(); + } } if state == RTCPeerConnectionState::Failed && reconnect_enabled.load(Ordering::SeqCst) { let (pub_h, sub_h) = core.pc_health().await; @@ -220,6 +220,20 @@ pub(super) async fn handle_event( return Ok(()); } use sfu_event::EventPayload as E; + // During a migration the old SFU still sends events. Its WebRTC commands + // are not for the new connection. + if !context.reconnect_enabled.load(Ordering::SeqCst) + && matches!( + payload, + E::SubscriberOffer(_) + | E::IceTrickle(_) + | E::ChangePublishOptions(_) + | E::ChangePublishQuality(_) + | E::IceRestart(_) + ) + { + return Ok(()); + } match payload { E::SubscriberOffer(offer) => { negotiate_subscriber( @@ -245,7 +259,7 @@ pub(super) async fn handle_event( } E::ParticipantJoined(ev) => { if let Some(p) = ev.participant { - core.roster_upsert(&p); + core.upsert_participant(&p); core.recompute_subscriptions_for_generation(context.generation) .await?; let _ = core.events_tx.send(CallEvent::ParticipantJoined(p)); @@ -253,7 +267,7 @@ pub(super) async fn handle_event( } E::ParticipantLeft(ev) => { if let Some(p) = ev.participant { - core.roster_remove(&p.session_id); + core.remove_participant(&p.session_id); core.recompute_subscriptions_for_generation(context.generation) .await?; let _ = core.events_tx.send(CallEvent::ParticipantLeft(p)); @@ -261,14 +275,14 @@ pub(super) async fn handle_event( } E::ParticipantUpdated(ev) => { if let Some(p) = ev.participant { - core.roster_upsert(&p); + core.upsert_participant(&p); core.recompute_subscriptions_for_generation(context.generation) .await?; let _ = core.events_tx.send(CallEvent::ParticipantUpdated(p)); } } E::TrackPublished(ev) => { - core.roster_add_track( + core.add_published_track( &ev.user_id, &ev.session_id, ev.r#type, @@ -283,7 +297,7 @@ pub(super) async fn handle_event( }); } E::TrackUnpublished(ev) => { - core.roster_remove_track(&ev.session_id, ev.r#type); + core.remove_published_track(&ev.session_id, ev.r#type); core.recompute_subscriptions_for_generation(context.generation) .await?; let _ = core.events_tx.send(CallEvent::TrackUnpublished { diff --git a/src/rtc/join/lifecycle.rs b/src/rtc/join/lifecycle.rs index 8f5f568..0c07087 100644 --- a/src/rtc/join/lifecycle.rs +++ b/src/rtc/join/lifecycle.rs @@ -36,7 +36,7 @@ impl RtcCore { *source = Some(token_source); } let user_token = match self - .while_generation(generation, self.reload_user_token()) + .while_generation(generation, self.reload_user_token(generation)) .await .and_then(|result| result) { @@ -54,7 +54,7 @@ impl RtcCore { .await .and_then(|result| result) { - self.stop_coordinator_events().await; + self.stop_coordinator_events(generation).await; self.set_state_if_current(generation, CallingState::Idle); return Err(error); } @@ -64,7 +64,7 @@ impl RtcCore { .await .and_then(|result| result); if result.is_err() { - self.stop_coordinator_events().await; + self.stop_coordinator_events(generation).await; // Restore to a non-joining terminal state so a retry is allowed. if self.state() == CallingState::Joining { self.set_state_if_current(generation, CallingState::Idle); @@ -73,7 +73,7 @@ impl RtcCore { result } - pub(super) async fn reload_user_token(&self) -> Result { + pub(super) async fn reload_user_token(&self, generation: u64) -> Result { let _refresh = self.token_refresh.lock().await; let source = self .token_source @@ -88,6 +88,10 @@ impl RtcCore { .user_id .clone(); let token = source.load_with_expiry_retry(&user_id).await?; + let lifecycle = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + if lifecycle.generation != generation { + return Err(join_cancelled()); + } *self.user_token.lock().unwrap_or_else(|e| e.into_inner()) = token.clone(); Ok(token) } @@ -108,7 +112,7 @@ impl RtcCore { Ok(token) } - pub(super) async fn refresh_expired_user_token(&self) -> Result { + pub(super) async fn refresh_expired_user_token(&self, generation: u64) -> Result { let can_refresh = self .token_source .lock() @@ -120,10 +124,10 @@ impl RtcCore { crate::error::TokenError::ExpiredByServer, )); } - self.reload_user_token().await + self.reload_user_token(generation).await } - pub(super) async fn refresh_before_full_reconnect(&self) -> Result { + pub(super) async fn refresh_before_full_reconnect(&self, generation: u64) -> Result { let refresh = self .token_source .lock() @@ -131,11 +135,13 @@ impl RtcCore { .as_ref() .is_some_and(UserTokenSource::refreshes_before_full_reconnect); if refresh { - self.reload_user_token().await + self.reload_user_token(generation).await } else { match self.current_user_token() { Ok(token) => Ok(token), - Err(error) if error.is_token_expired() => self.refresh_expired_user_token().await, + Err(error) if error.is_token_expired() => { + self.refresh_expired_user_token(generation).await + } Err(error) => Err(error), } } @@ -162,6 +168,7 @@ impl RtcCore { let mut migrating_from: Option = None; let mut edge_failures: std::collections::HashMap = std::collections::HashMap::new(); + let mut confirmed_bad_sfus: Vec = Vec::new(); let mut last_err: Option = None; let mut expired_retry_used = false; @@ -174,13 +181,7 @@ impl RtcCore { notify: data.notify.then_some(true), video: data.video.then_some(true), migrating_from: migrating_from.clone(), - migrating_from_list: { - let bad = self - .confirmed_bad_sfus - .lock() - .unwrap_or_else(|e| e.into_inner()); - bad.clone() - }, + migrating_from_list: confirmed_bad_sfus.clone(), ..Default::default() }; @@ -203,7 +204,7 @@ impl RtcCore { .err() .is_some_and(|(error, _)| error.is_token_expired()); if is_expired && !expired_retry_used { - user_token = self.refresh_expired_user_token().await?; + user_token = self.refresh_expired_user_token(generation).await?; expired_retry_used = true; continue; } @@ -216,7 +217,6 @@ impl RtcCore { return Err(join_cancelled()); } tracing::info!(cid = %self.cid(), edge = %success.edge_name, "stream.rtc.joined"); - *self.started.lock().unwrap_or_else(|e| e.into_inner()) = Some(Instant::now()); if !self.set_state_if_current(generation, CallingState::Joined) { return Err(join_cancelled()); } @@ -258,12 +258,8 @@ impl RtcCore { reconnect::JoinAttemptOutcome::Retry { delay, switch_sfu } => { if switch_sfu && let Some(edge) = edge_name { migrating_from = Some(edge.clone()); - let mut bad = self - .confirmed_bad_sfus - .lock() - .unwrap_or_else(|e| e.into_inner()); - if !bad.contains(&edge) { - bad.push(edge); + if !confirmed_bad_sfus.contains(&edge) { + confirmed_bad_sfus.push(edge); } } last_err = Some(err); @@ -590,28 +586,34 @@ impl RtcCore { })); // Spawn the WS event loop + health-check ping loop + stats loop. - let event_loop = self.spawn_runtime_task(event_loop( - receiver, - EventLoopContext { - core: self.clone(), - subscriber: subscriber.clone(), - publisher: publisher.clone(), - signal: signal.clone(), - session_id: session_id.clone(), - pending_ice: pending_ice.clone(), - generation, - ws_healthy: ws_healthy.clone(), - reconnect_enabled: reconnect_enabled.clone(), - }, - )); - let ping_loop = self.spawn_runtime_task(ping_loop( - self.clone(), - sfu_sender.clone(), + let event_loop = self.spawn_generation_task( generation, - ws_healthy.clone(), - reconnect_enabled.clone(), - )); - let stats_loop = self.spawn_runtime_task(stats::run(stats.clone())); + event_loop( + receiver, + EventLoopContext { + core: self.clone(), + subscriber: subscriber.clone(), + publisher: publisher.clone(), + signal: signal.clone(), + session_id: session_id.clone(), + pending_ice: pending_ice.clone(), + generation, + ws_healthy: ws_healthy.clone(), + reconnect_enabled: reconnect_enabled.clone(), + }, + ), + ); + let ping_loop = self.spawn_generation_task( + generation, + ping_loop( + self.clone(), + sfu_sender.clone(), + generation, + ws_healthy.clone(), + reconnect_enabled.clone(), + ), + ); + let stats_loop = self.spawn_generation_task(generation, stats::run(stats.clone())); Ok(Connection { generation, @@ -630,7 +632,7 @@ impl RtcCore { reconnect_enabled, signal_tasks: vec![event_loop, ping_loop], publisher_tasks: Vec::new(), - stats_task: stats_loop, + stats_task: Some(stats_loop), }) } } @@ -640,10 +642,7 @@ impl RtcCore { /// abort background tasks. Succeeds from any state, including `Joining` /// (JS: force to a leaving state rather than waiting for `JOINED`). pub async fn leave(&self, reason: impl Into) -> Result<()> { - self.leave_inner(reason.into()).await - } - - pub(super) async fn leave_inner(&self, reason: String) -> Result<()> { + let reason = reason.into(); let generation = self.cancel_generation(); let connection = self.connection.lock().await.take(); @@ -659,28 +658,34 @@ impl RtcCore { } connection.teardown().await; } - self.roster - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clear(); - *self - .call_state - .lock() - .unwrap_or_else(|error| error.into_inner()) = CallStateCache::default(); - self.active_subs - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clear(); - self.own_capabilities - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clear(); - *self - .reconnect_generation - .lock() - .unwrap_or_else(|e| e.into_inner()) = None; + { + // A join that started during the awaits above owns these fields. + let lifecycle = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + if lifecycle.generation == generation { + self.participants + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clear(); + *self + .call_state + .lock() + .unwrap_or_else(|error| error.into_inner()) = CallStateCache::default(); + self.active_subs + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clear(); + self.own_capabilities + .lock() + .unwrap_or_else(|e| e.into_inner()) + .clear(); + *self + .reconnect_generation + .lock() + .unwrap_or_else(|e| e.into_inner()) = None; + } + } self.set_state_if_current(generation, CallingState::Left); - self.stop_coordinator_events().await; + self.stop_coordinator_events(generation).await; Ok(()) } } @@ -723,11 +728,8 @@ impl RtcCore { let local_user_id = user_id.to_owned(); let sender = self.events_tx.clone(); let event_core = self.clone(); - let event_task = self.spawn_runtime_task(async move { + let event_task = self.spawn_generation_task(generation, async move { loop { - if !event_core.is_generation_current(generation) { - break; - } match events.recv().await { Ok(Some(event)) if event.raw.get("call_cid").and_then(|value| value.as_str()) @@ -762,14 +764,11 @@ impl RtcCore { } }); let health_core = self.clone(); - let health_task = self.spawn_runtime_task(async move { + let health_task = self.spawn_generation_task(generation, async move { let mut interval = tokio::time::interval(Duration::from_secs(20)); interval.tick().await; loop { interval.tick().await; - if !health_core.is_generation_current(generation) { - break; - } if let Err(error) = coordinator.send_health_check().await { tracing::warn!(%error, "stream.rtc.coordinator_health_failed"); health_core.clear_coordinator_connection(generation); @@ -870,15 +869,27 @@ impl RtcCore { abort_tasks(tasks).await; } - pub(super) async fn stop_coordinator_events(&self) { - *self - .coordinator_connection_id - .lock() - .unwrap_or_else(|error| error.into_inner()) = None; - self.user_token - .lock() - .unwrap_or_else(|error| error.into_inner()) - .clear(); - self.stop_coordinator_tasks().await; + /// Does nothing when `generation` is stale: the fields belong to a newer join. + pub(super) async fn stop_coordinator_events(&self, generation: u64) { + let tasks = { + let lifecycle = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + if lifecycle.generation != generation { + return; + } + *self + .coordinator_connection_id + .lock() + .unwrap_or_else(|error| error.into_inner()) = None; + self.user_token + .lock() + .unwrap_or_else(|error| error.into_inner()) + .clear(); + self.coordinator_tasks + .lock() + .unwrap_or_else(|error| error.into_inner()) + .drain(..) + .collect() + }; + abort_tasks(tasks).await; } } diff --git a/src/rtc/join/mod.rs b/src/rtc/join/mod.rs index 319195d..a14d292 100644 --- a/src/rtc/join/mod.rs +++ b/src/rtc/join/mod.rs @@ -11,7 +11,7 @@ //! - a typed [`CallEvent`] broadcast stream (participant joined/left, tracks, …); //! - the reconnect state machine (`RtcCore::run_reconnect`) driven by the pure //! decision logic in [`super::reconnect`], with dedup, the rejoin rate limiter, -//! the ICE / negotiation caps, the disconnection timeout, and the +//! the ICE / negotiation limits, the disconnection timeout, and the //! `restore_published_tracks` / `restore_subscribed_tracks` hooks. //! //! This root file holds [`RtcCore`] itself — its fields, lifecycle/generation @@ -22,7 +22,7 @@ //! - `connection` — the SFU WebSocket handshake, callbacks, event dispatch; //! - `publish` — the publish path; `publication` — its per-track state; //! - `subscriptions_runtime` — subscription negotiation and inbound tracks; -//! - `roster` — the participant roster and cached call state; +//! - `participants` — the participant state and cached call state; //! - `reconnect_runtime` — reconnect execution and media restoration. //! //! `reconnect_runtime` and `subscriptions_runtime` carry the suffix to avoid @@ -59,8 +59,8 @@ use super::proto::models::{self, PeerType, TrackType}; use super::proto::signal; use super::publish_options::ClientPublishOptions; use super::reconnect::{ - self, FailureCaps, ReconnectStrategy, SlidingWindowRateLimiter, escalate_strategy, - strategy_after_signal_close, + self, FailureLimits, ReconnectStrategy, SfuRejoinFailures, SlidingWindowRateLimiter, + escalate_strategy, strategy_after_signal_close, }; use super::sfu::signal::SignalClient; use super::sfu::ws::{self, SfuReceiver, SfuSender}; @@ -73,18 +73,18 @@ use serde_json::json; mod connection; mod lifecycle; +mod participants; mod publication; mod publish; mod reconnect_runtime; -mod roster; mod subscriptions_runtime; use connection::{ await_join_response, build_sfu_ws_url, event_loop, ping_loop, register_connection_state, register_on_track, }; +use participants::{CallStateCache, ParticipantState}; use publication::{MediaState, PublicationStatus}; -use roster::{CallStateCache, RosterEntry}; const MIGRATION_COMPLETE_TIMEOUT: Duration = Duration::from_secs(7); @@ -172,8 +172,6 @@ pub enum CallingState { ReconnectingFailed, /// Left (terminal). Left, - /// Network is offline; waiting to resume. - Offline, } /// A typed SFU event delivered on the [`Call`](crate::Call) event stream. @@ -274,6 +272,19 @@ struct Lifecycle { generation: u64, publish_options: ClientPublishOptions, generation_publish_options: ClientPublishOptions, + failure_limits: FailureLimits, + rate_limiter: SlidingWindowRateLimiter, +} + +impl Lifecycle { + /// Call with the lifecycle lock held, so events arrive in the order of the + /// state changes. + fn set_state(&mut self, next: CallingState, events: &broadcast::Sender) { + if self.state != next { + self.state = next; + let _ = events.send(CallEvent::CallingStateChanged(next)); + } + } } /// A live SFU connection bundle. Swapped out wholesale on REJOIN/MIGRATE. @@ -297,7 +308,7 @@ struct Connection { signal_tasks: Vec>, /// RTCP readers belong to the publisher PC and survive FAST reconnect. publisher_tasks: Vec>, - stats_task: JoinHandle<()>, + stats_task: Option>, } #[derive(Default)] @@ -429,13 +440,30 @@ impl Connection { self.stats.flush().await; let mut tasks = std::mem::take(&mut self.signal_tasks); tasks.append(&mut self.publisher_tasks); - tasks.push(self.stats_task); + tasks.extend(self.stats_task.take()); abort_tasks(tasks).await; let _ = self.subscriber.close().await; let _ = self.publisher.close().await; } } +impl Drop for Connection { + /// Stops the tasks when a cancelled future drops the connection before + /// `teardown`. The tasks own the PeerConnections, so this also drops them. + fn drop(&mut self) { + self.reconnect_enabled.store(false, Ordering::SeqCst); + self.stats.stop(); + for task in self + .signal_tasks + .iter() + .chain(&self.publisher_tasks) + .chain(&self.stats_task) + { + task.abort(); + } + } +} + async fn abort_tasks(tasks: Vec>) { for task in &tasks { task.abort(); @@ -467,16 +495,11 @@ pub struct RtcCore { stats_options: StdMutex, own_capabilities: StdMutex>, disconnection_timeout: StdMutex, - caps: StdMutex, - rate_limiter: StdMutex, - confirmed_bad_sfus: StdMutex>, - reconnect_edge_failures: StdMutex>, reconnect_generation: StdMutex>, reconnect_attempts: AtomicU32, next_connection_epoch: AtomicU64, migration_waiter: StdMutex)>>, join_data: StdMutex, - started: StdMutex>, /// Stable session id spanning reconnects within one join→leave lifecycle, /// reported as `SendStats.unified_session_id` so the dashboard correlates a /// participant across FAST/REJOIN/MIGRATE (JS `unifiedSessionId`). @@ -494,7 +517,7 @@ pub struct RtcCore { /// Exact per-session subscriptions, or `None` while using the coarse policy. manual_subscriptions: StdMutex>>, /// Known participants keyed by session id (correlation + subscription build). - roster: StdMutex>, + participants: StdMutex>, /// Call-level state supplied by join and incremental SFU events. call_state: StdMutex, /// Serialized publisher negotiation and retryable local publication state. @@ -537,29 +560,26 @@ impl RtcCore { generation: 0, publish_options: ClientPublishOptions::default(), generation_publish_options: ClientPublishOptions::default(), + failure_limits: FailureLimits::default(), + rate_limiter: SlidingWindowRateLimiter::rejoin_default(), }), lifecycle_changed: Notify::new(), connection: TokioMutex::new(None), stats_options: StdMutex::new(StatsOptions::default()), own_capabilities: StdMutex::new(HashSet::new()), disconnection_timeout: StdMutex::new(Duration::ZERO), - caps: StdMutex::new(FailureCaps::default()), - rate_limiter: StdMutex::new(SlidingWindowRateLimiter::rejoin_default()), - confirmed_bad_sfus: StdMutex::new(Vec::new()), - reconnect_edge_failures: StdMutex::new(HashMap::new()), reconnect_generation: StdMutex::new(None), reconnect_attempts: AtomicU32::new(0), next_connection_epoch: AtomicU64::new(0), migration_waiter: StdMutex::new(None), join_data: StdMutex::new(JoinCallData::new("")), - started: StdMutex::new(None), unified_session_id: StdMutex::new(String::new()), on_track_cb: StdMutex::new(None), sub_config: StdMutex::new(SubscriptionConfig::default()), subs_active: AtomicBool::new(false), manual_unsub: StdMutex::new(HashSet::new()), manual_subscriptions: StdMutex::new(None), - roster: StdMutex::new(HashMap::new()), + participants: StdMutex::new(HashMap::new()), call_state: StdMutex::new(CallStateCache::default()), media: TokioMutex::new(MediaState::default()), active_subs: StdMutex::new(Vec::new()), @@ -582,6 +602,19 @@ impl RtcCore { }) } + /// Spawn a runtime task that stops at its next `await` after `generation` + /// ends. A task that calls `leave` must not use this: `leave` ends the + /// generation and would cancel itself. + fn spawn_generation_task(self: &Arc, generation: u64, future: F) -> JoinHandle<()> + where + F: Future + Send + 'static, + { + let core = self.clone(); + self.spawn_runtime_task(async move { + let _ = core.while_generation(generation, future).await; + }) + } + fn cid(&self) -> String { format!("{}:{}", self.call_type, self.call_id) } @@ -628,14 +661,11 @@ impl RtcCore { } fn set_state_if_current(&self, generation: u64, next: CallingState) -> bool { - { - let mut guard = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); - if guard.generation != generation { - return false; - } - guard.state = next; + let mut guard = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + if guard.generation != generation { + return false; } - let _ = self.events_tx.send(CallEvent::CallingStateChanged(next)); + guard.set_state(next, &self.events_tx); true } @@ -694,8 +724,10 @@ impl RtcCore { match guard.state { CallingState::Idle | CallingState::Left => { guard.generation = guard.generation.wrapping_add(1); - guard.state = CallingState::Joining; + guard.set_state(CallingState::Joining, &self.events_tx); guard.generation_publish_options = guard.publish_options; + guard.failure_limits = FailureLimits::default(); + guard.rate_limiter = SlidingWindowRateLimiter::rejoin_default(); guard.generation } _ => { @@ -711,17 +743,6 @@ impl RtcCore { .lock() .unwrap_or_else(|e| e.into_inner()) = None; self.reconnect_attempts.store(0, Ordering::SeqCst); - *self.caps.lock().unwrap_or_else(|e| e.into_inner()) = FailureCaps::default(); - *self.rate_limiter.lock().unwrap_or_else(|e| e.into_inner()) = - SlidingWindowRateLimiter::rejoin_default(); - self.confirmed_bad_sfus - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clear(); - self.reconnect_edge_failures - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clear(); Ok(generation) } @@ -729,9 +750,6 @@ impl RtcCore { let generation = { let mut guard = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); guard.generation = guard.generation.wrapping_add(1); - if guard.state == CallingState::Joining { - guard.state = CallingState::Reconnecting; - } guard.generation }; self.lifecycle_changed.notify_waiters(); diff --git a/src/rtc/join/roster.rs b/src/rtc/join/participants.rs similarity index 77% rename from src/rtc/join/roster.rs rename to src/rtc/join/participants.rs index 9e7ebb5..537e7d7 100644 --- a/src/rtc/join/roster.rs +++ b/src/rtc/join/participants.rs @@ -1,5 +1,5 @@ -//! The participant roster and cached call state: the SFU's view of who is -//! in the call, what they publish, and the call-level state that join and +//! The participants and cached call state: the SFU's view of who is in the +//! call, what they publish, and the call-level state that join and //! incremental SFU events maintain. use super::*; @@ -7,7 +7,7 @@ use super::*; /// A participant known to be in the call, used to correlate inbound tracks /// (by `track_lookup_prefix`) and to build the subscription list. #[derive(Clone, Default)] -pub(super) struct RosterEntry { +pub(super) struct ParticipantState { pub(super) user_id: String, pub(super) session_id: String, pub(super) track_lookup_prefix: String, @@ -27,11 +27,11 @@ pub(super) struct CallStateCache { impl RtcCore { /// A snapshot of the participants currently known in the call (including this - /// session), built from the SFU roster. Updated as `ParticipantJoined` / - /// `ParticipantLeft` events arrive. + /// session), built from the SFU participant state. Updated as + /// `ParticipantJoined` / `ParticipantLeft` events arrive. pub fn participants(&self) -> Vec { - let roster = self.roster.lock().unwrap_or_else(|e| e.into_inner()); - roster + let participants = self.participants.lock().unwrap_or_else(|e| e.into_inner()); + participants .values() .map(|entry| { RemoteParticipant::from_proto(&entry.participant, entry.paused.iter().copied()) @@ -67,8 +67,8 @@ impl RtcCore { } pub(super) fn lookup_participant(&self, prefix: &str) -> RemoteParticipant { - let roster = self.roster.lock().unwrap_or_else(|e| e.into_inner()); - for entry in roster.values() { + let participants = self.participants.lock().unwrap_or_else(|e| e.into_inner()); + for entry in participants.values() { if !entry.track_lookup_prefix.is_empty() && entry.track_lookup_prefix == prefix { return RemoteParticipant::from_proto( &entry.participant, @@ -83,8 +83,8 @@ impl RtcCore { } } - /// Replace the roster from an authoritative SFU join response when its - /// lifecycle generation is still active. + /// Replace the participants from an authoritative SFU join response when + /// its lifecycle generation is still active. pub(super) fn apply_join_call_state_if_current( &self, generation: u64, @@ -97,7 +97,7 @@ impl RtcCore { return false; } let state = call_state.unwrap_or_default(); - let participants = state.participants.clone(); + let joined = state.participants.clone(); *self .call_state .lock() @@ -109,18 +109,20 @@ impl RtcCore { current_grants: None, }; { - let mut roster = self.roster.lock().unwrap_or_else(|e| e.into_inner()); - roster.clear(); - let me = roster.entry(session_id.to_owned()).or_default(); + let mut participants = self.participants.lock().unwrap_or_else(|e| e.into_inner()); + participants.clear(); + let me = participants.entry(session_id.to_owned()).or_default(); me.user_id = user_id.to_owned(); me.session_id = session_id.to_owned(); me.participant.user_id = user_id.to_owned(); me.participant.session_id = session_id.to_owned(); - for participant in &participants { + for participant in &joined { if participant.session_id.is_empty() { continue; } - let entry = roster.entry(participant.session_id.clone()).or_default(); + let entry = participants + .entry(participant.session_id.clone()) + .or_default(); entry.user_id.clone_from(&participant.user_id); entry.session_id.clone_from(&participant.session_id); entry.participant.clone_from(participant); @@ -135,7 +137,7 @@ impl RtcCore { .extend(participant.published_tracks.iter().copied()); } } - for participant in participants { + for participant in joined { if participant.session_id != session_id { let _ = self .events_tx @@ -145,13 +147,13 @@ impl RtcCore { true } - /// Insert/refresh a participant's roster entry from a `Participant` message. - pub(super) fn roster_upsert(&self, p: &models::Participant) { + /// Insert/refresh a participant's state from a `Participant` message. + pub(super) fn upsert_participant(&self, p: &models::Participant) { if p.session_id.is_empty() { return; } - let mut roster = self.roster.lock().unwrap_or_else(|e| e.into_inner()); - let entry = roster.entry(p.session_id.clone()).or_default(); + let mut participants = self.participants.lock().unwrap_or_else(|e| e.into_inner()); + let entry = participants.entry(p.session_id.clone()).or_default(); entry.user_id = p.user_id.clone(); entry.session_id = p.session_id.clone(); entry.participant.clone_from(p); @@ -162,8 +164,8 @@ impl RtcCore { entry.published.extend(p.published_tracks.iter().copied()); } - pub(super) fn roster_remove(&self, session_id: &str) { - self.roster + pub(super) fn remove_participant(&self, session_id: &str) { + self.participants .lock() .unwrap_or_else(|e| e.into_inner()) .remove(session_id); @@ -171,7 +173,7 @@ impl RtcCore { /// Record a newly-published track for a participant, learning the /// `track_lookup_prefix` from the optional participant hint when present. - pub(super) fn roster_add_track( + pub(super) fn add_published_track( &self, user_id: &str, session_id: &str, @@ -182,8 +184,8 @@ impl RtcCore { return; } { - let mut roster = self.roster.lock().unwrap_or_else(|e| e.into_inner()); - let entry = roster.entry(session_id.to_owned()).or_default(); + let mut participants = self.participants.lock().unwrap_or_else(|e| e.into_inner()); + let entry = participants.entry(session_id.to_owned()).or_default(); if let Some(participant) = hint { entry.participant.clone_from(participant); if !participant.track_lookup_prefix.is_empty() { @@ -210,9 +212,9 @@ impl RtcCore { } } - pub(super) fn roster_remove_track(&self, session_id: &str, track_type: i32) { - let mut roster = self.roster.lock().unwrap_or_else(|e| e.into_inner()); - if let Some(entry) = roster.get_mut(session_id) { + pub(super) fn remove_published_track(&self, session_id: &str, track_type: i32) { + let mut participants = self.participants.lock().unwrap_or_else(|e| e.into_inner()); + if let Some(entry) = participants.get_mut(session_id) { entry.published.remove(&track_type); entry .participant @@ -222,28 +224,28 @@ impl RtcCore { } pub(super) fn update_connection_quality(&self, updates: &[event::ConnectionQualityInfo]) { - let mut roster = self - .roster + let mut participants = self + .participants .lock() .unwrap_or_else(|error| error.into_inner()); for update in updates { - if let Some(entry) = roster.get_mut(&update.session_id) { + if let Some(entry) = participants.get_mut(&update.session_id) { entry.participant.connection_quality = update.connection_quality; } } } pub(super) fn update_audio_levels(&self, levels: &[event::AudioLevel]) { - let mut roster = self - .roster + let mut participants = self + .participants .lock() .unwrap_or_else(|error| error.into_inner()); - for entry in roster.values_mut() { + for entry in participants.values_mut() { entry.participant.is_speaking = false; entry.participant.audio_level = 0.0; } for level in levels { - if let Some(entry) = roster.get_mut(&level.session_id) { + if let Some(entry) = participants.get_mut(&level.session_id) { entry.participant.is_speaking = level.is_speaking; entry.participant.audio_level = level.level; } @@ -251,11 +253,11 @@ impl RtcCore { } pub(super) fn update_dominant_speaker(&self, session_id: &str) { - let mut roster = self - .roster + let mut participants = self + .participants .lock() .unwrap_or_else(|error| error.into_inner()); - for entry in roster.values_mut() { + for entry in participants.values_mut() { entry.participant.is_dominant_speaker = entry.session_id == session_id; } } @@ -275,12 +277,12 @@ impl RtcCore { } pub(super) fn update_inbound_state(&self, states: &[event::InboundVideoState]) { - let mut roster = self - .roster + let mut participants = self + .participants .lock() .unwrap_or_else(|error| error.into_inner()); for state in states { - let Some(entry) = roster.get_mut(&state.session_id) else { + let Some(entry) = participants.get_mut(&state.session_id) else { continue; }; if state.paused { diff --git a/src/rtc/join/publish.rs b/src/rtc/join/publish.rs index ffb1a8a..f4221a8 100644 --- a/src/rtc/join/publish.rs +++ b/src/rtc/join/publish.rs @@ -7,6 +7,13 @@ impl RtcCore { /// Publish a local track: add its send-only transceiver and renegotiate the /// publisher PC with the SFU (`SetPublisher`). Errors if not joined. pub async fn publish(self: &Arc, track: LocalTrack) -> Result<()> { + if let LocalTrack::Video { track_type, .. } = &track + && !matches!(track_type, TrackType::Video | TrackType::ScreenShare) + { + return Err(RtcError::IllegalState(format!( + "a video track cannot be published as {track_type:?}" + ))); + } let Some((publisher, signal, session_id, publish_options)) = self.publisher_handles().await else { return Err(RtcError::IllegalState("publish() before join()".to_owned())); @@ -115,7 +122,7 @@ impl RtcCore { .unwrap_or_else(|e| e.into_inner()) .user_id .clone(); - self.roster_add_track(&user_id, &session_id, track.track_type() as i32, None); + self.add_published_track(&user_id, &session_id, track.track_type() as i32, None); track.start_media(); signal .update_mute_states(signal::UpdateMuteStatesRequest { @@ -189,7 +196,7 @@ impl RtcCore { }) .await?; if muted { - self.roster_remove_track(&session_id, track_type as i32); + self.remove_published_track(&session_id, track_type as i32); } if let Some(removed) = media.remove(&track_id) { removed.stop(); @@ -207,6 +214,11 @@ impl RtcCore { track_type: TrackType, muted: bool, ) -> Result<()> { + if track_type == TrackType::Unspecified { + return Err(RtcError::IllegalState( + "cannot mute an unspecified track type".to_owned(), + )); + } if !muted { let capability = required_publish_capability(track_type); if !self @@ -267,7 +279,7 @@ impl RtcCore { return Err(error); } if muted { - self.roster_remove_track(&session_id, track_type as i32); + self.remove_published_track(&session_id, track_type as i32); } else { let user_id = self .join_data @@ -275,7 +287,7 @@ impl RtcCore { .unwrap_or_else(|error| error.into_inner()) .user_id .clone(); - self.roster_add_track(&user_id, &session_id, track_type as i32, None); + self.add_published_track(&user_id, &session_id, track_type as i32, None); } Ok(()) } diff --git a/src/rtc/join/reconnect_runtime.rs b/src/rtc/join/reconnect_runtime.rs index 59c187e..03cff94 100644 --- a/src/rtc/join/reconnect_runtime.rs +++ b/src/rtc/join/reconnect_runtime.rs @@ -99,9 +99,9 @@ impl RtcCore { .clone(); for track in tracks { if track.is_muted() { - self.roster_remove_track(session_id, track.track_type() as i32); + self.remove_published_track(session_id, track.track_type() as i32); } else { - self.roster_add_track(&user_id, session_id, track.track_type() as i32, None); + self.add_published_track(&user_id, session_id, track.track_type() as i32, None); } track.start_media(); } @@ -139,7 +139,7 @@ impl RtcCore { }) .await?; if muted { - self.roster_remove_track(session_id, *track_type as i32); + self.remove_published_track(session_id, *track_type as i32); } } for (track_id, _) in pending { @@ -206,7 +206,7 @@ impl RtcCore { } /// The reconnect state machine loop (JS `Call.reconnect`). Honors the rejoin - /// rate limiter, ICE / negotiation caps, the disconnection timeout, and the + /// rate limiter, ICE / negotiation limits, the disconnection timeout, and the /// restore hooks. Bounded: it stops when `JOINED`, `RECONNECTING_FAILED`, or /// `LEFT`, so it can never spin. pub(super) async fn run_reconnect( @@ -222,14 +222,7 @@ impl RtcCore { let start = Instant::now(); let mut attempt = 0; let mut was_migrating = strategy == ReconnectStrategy::Migrate; - self.reconnect_edge_failures - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clear(); - self.confirmed_bad_sfus - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clear(); + let mut sfu_failures = SfuRejoinFailures::default(); self.set_state_if_current( generation, @@ -241,12 +234,12 @@ impl RtcCore { ); if reason == reconnect::REASON_ICE_UNSUPPORTED { - let tripped = self - .caps - .lock() - .unwrap_or_else(|e| e.into_inner()) - .record_ice_never_connected(); - if tripped { + let limit_reached = { + let mut lifecycle = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + lifecycle.generation == generation + && lifecycle.failure_limits.record_ice_never_connected() + }; + if limit_reached { let _ = self.leave(reconnect::REASON_ICE_UNSUPPORTED).await; return; } @@ -277,11 +270,13 @@ impl RtcCore { // Rate limit only REJOIN/MIGRATE. if strategy.is_rate_limited() { let now_ms = elapsed_ms(); - let allowed = self - .rate_limiter - .lock() - .unwrap_or_else(|e| e.into_inner()) - .try_register(now_ms); + let allowed = { + let mut lifecycle = self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + if lifecycle.generation != generation { + return; + } + lifecycle.rate_limiter.try_register(now_ms) + }; if !allowed { let _ = self.leave(reconnect::REASON_REJOIN_LIMIT).await; return; @@ -292,7 +287,8 @@ impl RtcCore { let outcome = match self .while_generation( generation, - self.clone().reconnect_once(generation, strategy, &reason), + self.clone() + .reconnect_once(generation, strategy, &reason, &mut sfu_failures), ) .await { @@ -301,10 +297,13 @@ impl RtcCore { }; match outcome { Ok(()) => { - self.caps - .lock() - .unwrap_or_else(|e| e.into_inner()) - .reset_negotiation(); + { + let mut lifecycle = + self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + if lifecycle.generation == generation { + lifecycle.failure_limits.reset_negotiation(); + } + } self.set_state_if_current(generation, CallingState::Joined); return; } @@ -324,12 +323,13 @@ impl RtcCore { return; } if matches!(err, RtcError::Negotiation(_)) { - let tripped = self - .caps - .lock() - .unwrap_or_else(|e| e.into_inner()) - .record_negotiation_failure(); - if tripped { + let limit_reached = { + let mut lifecycle = + self.lifecycle.lock().unwrap_or_else(|e| e.into_inner()); + lifecycle.generation == generation + && lifecycle.failure_limits.record_negotiation_failure() + }; + if limit_reached { let _ = self.leave(reconnect::REASON_NEGOTIATION_FAILURES).await; return; } @@ -377,11 +377,15 @@ impl RtcCore { generation: u64, strategy: ReconnectStrategy, reason: &str, + sfu_failures: &mut SfuRejoinFailures, ) -> Result<()> { self.observe_reconnect(strategy, ReconnectFaultPoint::BeforeAttempt)?; match strategy { ReconnectStrategy::Fast => self.reconnect_fast(generation, reason).await, - ReconnectStrategy::Rejoin => self.reconnect_rejoin(generation, reason).await, + ReconnectStrategy::Rejoin => { + self.reconnect_rejoin(generation, reason, sfu_failures) + .await + } ReconnectStrategy::Migrate => self.reconnect_migrate(generation, reason).await, ReconnectStrategy::Disconnect => Err(RtcError::IllegalState( "disconnect strategy must leave the call".to_owned(), @@ -448,6 +452,7 @@ impl RtcCore { self: Arc, generation: u64, reason: &str, + sfu_failures: &mut SfuRejoinFailures, ) -> Result<()> { let data = self .join_data @@ -469,11 +474,7 @@ impl RtcCore { } }; - let confirmed_bad_sfus = self - .confirmed_bad_sfus - .lock() - .unwrap_or_else(|e| e.into_inner()) - .clone(); + let confirmed_bad_sfus = sfu_failures.confirmed().to_vec(); let request = JoinCallRequest { location: data .location @@ -495,7 +496,7 @@ impl RtcCore { .await?, ); - let mut token = self.refresh_before_full_reconnect().await?; + let mut token = self.refresh_before_full_reconnect(generation).await?; self.ensure_coordinator_events(generation, &token).await?; let reconnect_attempt = self.reconnect_attempts.load(Ordering::SeqCst); let mut result = self @@ -516,7 +517,7 @@ impl RtcCore { .err() .is_some_and(|(error, _)| error.is_token_expired()) { - token = self.refresh_expired_user_token().await?; + token = self.refresh_expired_user_token(generation).await?; result = self .clone() .join_once(JoinOnceOptions { @@ -535,7 +536,7 @@ impl RtcCore { Ok(_) => {} Err((error, edge)) => { if let Some(edge) = edge { - self.record_reconnect_edge_failure(&edge, error.is_join_error_code()); + sfu_failures.record(&edge, error.is_join_error_code()); } return Err(error); } @@ -552,30 +553,6 @@ impl RtcCore { ) } - pub(super) fn record_reconnect_edge_failure(&self, edge: &str, force_switch: bool) { - let failures = { - let mut counts = self - .reconnect_edge_failures - .lock() - .unwrap_or_else(|e| e.into_inner()); - let count = counts.entry(edge.to_owned()).or_insert(0); - *count = count.saturating_add(1); - if force_switch { - *count = (*count).max(2); - } - *count - }; - if failures >= 2 { - let mut bad = self - .confirmed_bad_sfus - .lock() - .unwrap_or_else(|e| e.into_inner()); - if !bad.iter().any(|known| known == edge) { - bad.push(edge.to_owned()); - } - } - } - pub(super) async fn reconnect_migrate( self: Arc, generation: u64, @@ -586,7 +563,7 @@ impl RtcCore { .lock() .unwrap_or_else(|e| e.into_inner()) .clone(); - let mut token = self.refresh_before_full_reconnect().await?; + let mut token = self.refresh_before_full_reconnect(generation).await?; self.ensure_coordinator_events(generation, &token).await?; let (previous_session_id, migrating_from, old_reconnect_enabled) = { let guard = self.connection.lock().await; @@ -645,7 +622,7 @@ impl RtcCore { .err() .is_some_and(|(error, _)| error.is_token_expired()) { - token = match self.refresh_expired_user_token().await { + token = match self.refresh_expired_user_token(generation).await { Ok(token) => token, Err(error) => { drop(self.take_migration_waiter(generation)); @@ -854,27 +831,33 @@ impl RtcCore { } ws_healthy.store(true, Ordering::SeqCst); reconnect_enabled.store(true, Ordering::SeqCst); - let event_task = self.spawn_runtime_task(event_loop( - receiver, - EventLoopContext { - core: self.clone(), - subscriber, - publisher, - signal, - session_id: session_id.clone(), - pending_ice, - generation, - ws_healthy: ws_healthy.clone(), - reconnect_enabled: reconnect_enabled.clone(), - }, - )); - let ping_task = self.spawn_runtime_task(ping_loop( - self.clone(), - sfu_sender, + let event_task = self.spawn_generation_task( generation, - ws_healthy, - reconnect_enabled, - )); + event_loop( + receiver, + EventLoopContext { + core: self.clone(), + subscriber, + publisher, + signal, + session_id: session_id.clone(), + pending_ice, + generation, + ws_healthy: ws_healthy.clone(), + reconnect_enabled: reconnect_enabled.clone(), + }, + ), + ); + let ping_task = self.spawn_generation_task( + generation, + ping_loop( + self.clone(), + sfu_sender, + generation, + ws_healthy, + reconnect_enabled, + ), + ); let mut guard = self.connection.lock().await; let connection = guard.as_mut().ok_or_else(|| { RtcError::IllegalState("fast reconnect connection disappeared".to_owned()) diff --git a/src/rtc/join/subscriptions_runtime.rs b/src/rtc/join/subscriptions_runtime.rs index 7b24fa5..1be0f62 100644 --- a/src/rtc/join/subscriptions_runtime.rs +++ b/src/rtc/join/subscriptions_runtime.rs @@ -13,8 +13,8 @@ impl RtcCore { } /// Set the subscription policy and (re)send `UpdateSubscriptions`. Activates - /// the reactive subscriber: subscriptions are recomputed on every roster - /// change from here on. + /// the reactive subscriber: subscriptions are recomputed on every + /// participant change from here on. pub async fn update_subscriptions(&self, config: SubscriptionConfig) -> Result<()> { *self.sub_config.lock().unwrap_or_else(|e| e.into_inner()) = config; *self @@ -59,7 +59,7 @@ impl RtcCore { Ok(()) } - /// Rebuild the desired subscription list from the roster + policy and send it + /// Rebuild the desired subscription list from the participants + policy and send it /// to the SFU if it changed since the last send on this connection. pub(super) async fn recompute_subscriptions(&self) -> Result<()> { self.recompute_subscriptions_for_generation(self.generation()) @@ -99,10 +99,10 @@ impl RtcCore { let mut tracks: Vec = Vec::new(); { - let roster = self.roster.lock().unwrap_or_else(|e| e.into_inner()); + let participants = self.participants.lock().unwrap_or_else(|e| e.into_inner()); if let Some(targets) = targets { for target in targets { - let Some(entry) = roster.get(&target.session_id) else { + let Some(entry) = participants.get(&target.session_id) else { continue; }; if entry.session_id == session_id @@ -123,7 +123,7 @@ impl RtcCore { }); } } else { - for entry in roster.values() { + for entry in participants.values() { if entry.session_id == session_id { continue; } @@ -160,7 +160,7 @@ impl RtcCore { left.session_id == right.session_id && left.track_type == right.track_type }); - // Skip an identical resend (roster churn that doesn't change the set). + // Skip an identical resend (participant churn that doesn't change the set). if *self.active_subs.lock().unwrap_or_else(|e| e.into_inner()) == tracks { return Ok(()); } @@ -234,7 +234,7 @@ impl RtcCore { let on_drop = Box::new(move || { if let Some(core) = weak.upgrade() { let task_core = core.clone(); - std::mem::drop(core.spawn_runtime_task(async move { + std::mem::drop(core.spawn_generation_task(generation, async move { task_core .on_remote_track_dropped(generation, connection_epoch, key) .await; diff --git a/src/rtc/join/tests.rs b/src/rtc/join/tests.rs index 85d8c23..6300f69 100644 --- a/src/rtc/join/tests.rs +++ b/src/rtc/join/tests.rs @@ -41,17 +41,24 @@ fn prepare_joined_core(core: &Arc, user_id: &str) -> u64 { generation } -fn refresh_server() -> ( +/// A local HTTP server that answers one request with the JSON `body`. +fn one_shot_http_server( + body: &str, +) -> ( String, std::sync::mpsc::Receiver, thread::JoinHandle<()>, ) { - let listener = TcpListener::bind("127.0.0.1:0").expect("bind refresh server"); + let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTP server"); listener .set_nonblocking(true) - .expect("set refresh server nonblocking"); - let address = listener.local_addr().expect("refresh server address"); + .expect("set HTTP server nonblocking"); + let address = listener.local_addr().expect("HTTP server address"); let (request_tx, request_rx) = std::sync::mpsc::channel(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", + body.len() + ); let server = thread::spawn(move || { let deadline = Instant::now() + Duration::from_secs(3); loop { @@ -59,29 +66,24 @@ fn refresh_server() -> ( Ok((mut stream, _)) => { stream .set_nonblocking(false) - .expect("set refresh stream blocking"); + .expect("set HTTP stream blocking"); stream .set_read_timeout(Some(Duration::from_secs(1))) - .expect("set refresh read timeout"); + .expect("set HTTP read timeout"); let mut request = [0_u8; 4096]; - let read = stream.read(&mut request).expect("read refresh request"); + let read = stream.read(&mut request).expect("read HTTP request"); let request = String::from_utf8_lossy(&request[..read]).into_owned(); - request_tx.send(request).expect("record refresh request"); + request_tx.send(request).expect("record HTTP request"); stream - .write_all( - b"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: 2\r\nconnection: close\r\n\r\n{}", - ) - .expect("write refresh response"); + .write_all(response.as_bytes()) + .expect("write HTTP response"); return; } Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => { - assert!( - Instant::now() < deadline, - "refresh server received no request" - ); + assert!(Instant::now() < deadline, "HTTP server received no request"); thread::sleep(Duration::from_millis(5)); } - Err(error) => panic!("accept refresh request: {error}"), + Err(error) => panic!("accept HTTP request: {error}"), } } }); @@ -98,6 +100,158 @@ async fn wait_for(timeout: Duration, mut predicate: impl FnMut() -> bool, descri .unwrap_or_else(|_| panic!("timed out waiting for {description}")); } +/// A local SFU WebSocket. It sends each received request to the channel. The +/// channel closes when the client socket closes. With `answer_join`, it answers +/// the `JoinRequest`. +async fn fake_sfu( + answer_join: bool, +) -> ( + Credentials, + tokio::sync::mpsc::UnboundedReceiver, +) { + use futures_util::{SinkExt, StreamExt}; + use prost::Message as _; + use tokio_tungstenite::tungstenite::Message; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind fake SFU"); + let address = listener.local_addr().expect("fake SFU address"); + let (requests, received) = tokio::sync::mpsc::unbounded_channel(); + tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept SFU client"); + let mut socket = tokio_tungstenite::accept_async(stream) + .await + .expect("SFU WebSocket handshake"); + while let Some(Ok(message)) = socket.next().await { + let Message::Binary(bytes) = message else { + continue; + }; + let request = event::SfuRequest::decode(bytes).expect("SFU request"); + if answer_join + && matches!( + request.request_payload, + Some(event::sfu_request::RequestPayload::JoinRequest(_)) + ) + { + let response = SfuEvent { + event_payload: Some(sfu_event::EventPayload::JoinResponse( + JoinResponse::default(), + )), + }; + socket + .send(Message::Binary(response.encode_to_vec().into())) + .await + .expect("send join response"); + } + let _ = requests.send(request); + } + }); + let credentials = Credentials { + server: coordinator::SfuServer { + edge_name: "fake-edge".to_owned(), + url: "http://127.0.0.1:9/twirp".to_owned(), + ws_endpoint: format!("ws://{address}/ws"), + }, + token: "sfu-token".to_owned(), + ice_servers: Vec::new(), + }; + (credentials, received) +} + +/// A local coordinator WebSocket that sends `connection.ok`. It returns the +/// REST base URL and a task that ends when the client socket closes. +async fn fake_coordinator() -> (String, tokio::task::JoinHandle<()>) { + use futures_util::{SinkExt, StreamExt}; + use tokio_tungstenite::tungstenite::Message; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind fake coordinator"); + let address = listener.local_addr().expect("fake coordinator address"); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept coordinator client"); + let mut socket = tokio_tungstenite::accept_async(stream) + .await + .expect("coordinator WebSocket handshake"); + socket + .next() + .await + .expect("auth frame") + .expect("valid auth frame"); + socket + .send(Message::Text( + json!({ "type": "connection.ok", "connection_id": "connection-1" }) + .to_string() + .into(), + )) + .await + .expect("send connection.ok"); + while let Some(Ok(_)) = socket.next().await {} + }); + (format!("http://{address}"), server) +} + +/// Every request the fake SFU received until the client socket closed. +async fn requests_until_close( + mut received: tokio::sync::mpsc::UnboundedReceiver, +) -> Vec { + tokio::time::timeout(Duration::from_secs(2), async move { + let mut requests = Vec::new(); + while let Some(request) = received.recv().await { + requests.push(request); + } + requests + }) + .await + .expect("SFU socket closed") +} + +fn alive_tasks() -> usize { + tokio::runtime::Handle::current() + .metrics() + .num_alive_tasks() +} + +async fn establish_fake( + core: &Arc, + generation: u64, +) -> ( + Connection, + tokio::sync::mpsc::UnboundedReceiver, +) { + let (credentials, sfu) = fake_sfu(true).await; + let connection = core + .clone() + .establish( + &credentials, + 0, + ReconnectStrategy::Fast, + None, + generation, + None, + ) + .await + .expect("establish against fake SFU"); + (connection, sfu) +} + +/// The event loop context of `connection` after a migration detached it. +fn detached_context(core: &Arc, connection: &Connection) -> EventLoopContext { + connection.reconnect_enabled.store(false, Ordering::SeqCst); + EventLoopContext { + core: core.clone(), + subscriber: connection.subscriber.clone(), + publisher: connection.publisher.clone(), + signal: connection.signal.clone(), + session_id: connection.session_id.clone(), + pending_ice: connection.pending_ice.clone(), + generation: connection.generation, + ws_healthy: connection.ws_healthy.clone(), + reconnect_enabled: connection.reconnect_enabled.clone(), + } +} + fn preferred_codec(core: &RtcCore, generation: u64) -> Option { core.preferred_publish_options(generation) .expect("current generation") @@ -189,6 +343,293 @@ async fn leave_cancels_join_generation_and_allows_later_join() { assert_eq!(core.lifecycle_snapshot(), (CallingState::Joining, second)); } +#[tokio::test] +async fn leave_tears_down_the_stored_connection() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let (connection, sfu) = establish_fake(&core, generation).await; + let (subscriber, publisher) = (connection.subscriber.clone(), connection.publisher.clone()); + *core.connection.lock().await = Some(connection); + + core.leave("test leave").await.expect("leave"); + + let requests = requests_until_close(sfu).await; + assert!(requests.iter().any(|request| matches!( + request.request_payload, + Some(event::sfu_request::RequestPayload::LeaveCallRequest(_)) + ))); + assert_eq!( + subscriber.connection_state(), + RTCPeerConnectionState::Closed + ); + assert_eq!(publisher.connection_state(), RTCPeerConnectionState::Closed); + let (active, spawned, completed) = core.runtime_task_snapshot(); + assert_eq!(active, 0); + assert_eq!(spawned, completed); +} + +#[tokio::test] +async fn leave_closes_a_connection_owned_by_a_cancelled_join() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let baseline = alive_tasks(); + let (connection, sfu) = establish_fake(&core, generation).await; + let owner_core = core.clone(); + let owner = tokio::spawn(async move { + owner_core + .while_generation(generation, async move { + let _connection = connection; + std::future::pending::<()>().await; + }) + .await + }); + tokio::task::yield_now().await; + + core.leave("cancel join").await.expect("leave"); + + assert!(owner.await.expect("owner task").is_err()); + requests_until_close(sfu).await; + wait_for( + Duration::from_secs(2), + || alive_tasks() == baseline, + "cancelled connection cleanup", + ) + .await; +} + +#[tokio::test] +async fn generation_change_closes_the_coordinator_socket() { + let (base_url, coordinator) = fake_coordinator().await; + let core = test_core_with_config(ClientConfig { + base_url, + ..ClientConfig::default() + }); + let generation = prepare_joined_core(&core, "alice"); + let token = core.current_user_token().expect("user token"); + core.connect_coordinator_events(generation, &token, "alice") + .await + .expect("coordinator events"); + + core.cancel_generation(); + + tokio::time::timeout(Duration::from_secs(2), coordinator) + .await + .expect("coordinator socket closed") + .expect("fake coordinator task"); +} + +#[tokio::test] +async fn generation_change_stops_the_connection_tasks() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let (connection, _sfu) = establish_fake(&core, generation).await; + assert!(core.runtime_task_snapshot().0 > 0); + + core.cancel_generation(); + + wait_for( + Duration::from_secs(2), + || core.runtime_task_snapshot().0 == 0, + "connection tasks stop", + ) + .await; + drop(connection); +} + +#[tokio::test] +async fn failed_establish_leaves_no_background_tasks() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + core.join_data + .lock() + .unwrap_or_else(|error| error.into_inner()) + .join_response_timeout = Duration::from_millis(50); + let baseline = alive_tasks(); + let (credentials, sfu) = fake_sfu(false).await; + + let result = core + .clone() + .establish( + &credentials, + 0, + ReconnectStrategy::Fast, + None, + generation, + None, + ) + .await; + + assert!(matches!(result, Err(RtcError::Timeout(_)))); + requests_until_close(sfu).await; + wait_for( + Duration::from_secs(2), + || alive_tasks() == baseline, + "failed establish cleanup", + ) + .await; +} + +#[tokio::test] +async fn leave_during_establish_leaves_no_background_tasks() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let baseline = alive_tasks(); + let (credentials, mut sfu) = fake_sfu(false).await; + let owner_core = core.clone(); + let owner = tokio::spawn(async move { + owner_core + .while_generation( + generation, + owner_core.clone().establish( + &credentials, + 0, + ReconnectStrategy::Fast, + None, + generation, + None, + ), + ) + .await + .map(|_| ()) + }); + let join_request = tokio::time::timeout(Duration::from_secs(2), sfu.recv()) + .await + .expect("join request") + .expect("SFU socket open"); + assert!(matches!( + join_request.request_payload, + Some(event::sfu_request::RequestPayload::JoinRequest(_)) + )); + + core.leave("cancel establish").await.expect("leave"); + + assert!(owner.await.expect("owner task").is_err()); + requests_until_close(sfu).await; + wait_for( + Duration::from_secs(2), + || alive_tasks() == baseline, + "cancelled establish cleanup", + ) + .await; +} + +#[tokio::test] +async fn leave_that_overlaps_a_new_join_keeps_the_new_join_state() { + let core = test_core(); + prepare_joined_core(&core, "alice"); + core.leave("first leave").await.expect("first leave"); + let connection_slot = core.connection.lock().await; + let cancelled = core.generation(); + let leave_core = core.clone(); + let leave = tokio::spawn(async move { leave_core.leave("second leave").await }); + wait_for( + Duration::from_secs(1), + || core.generation() != cancelled, + "second leave cancels its generation", + ) + .await; + + let second = prepare_joined_core(&core, "alice"); + *core + .coordinator_connection_id + .lock() + .unwrap_or_else(|error| error.into_inner()) = Some((second, "connection-2".to_owned())); + assert!(core.apply_join_call_state_if_current(second, "session-2", "alice", None)); + assert!(core.claim_reconnect(second)); + drop(connection_slot); + leave.await.expect("leave task").expect("second leave"); + + assert_eq!(core.state(), CallingState::Joined); + assert!(core.user_auth().is_some()); + assert!( + core.participants + .lock() + .unwrap_or_else(|error| error.into_inner()) + .contains_key("session-2") + ); + assert_eq!(core.active_reconnect_generation(), Some(second)); +} + +#[tokio::test] +async fn stale_coordinator_stop_keeps_the_current_coordinator() { + let (base_url, coordinator) = fake_coordinator().await; + let core = test_core_with_config(ClientConfig { + base_url, + ..ClientConfig::default() + }); + let first = prepare_joined_core(&core, "alice"); + core.leave("cancel first join").await.expect("leave"); + let second = prepare_joined_core(&core, "alice"); + let token = core.current_user_token().expect("user token"); + core.connect_coordinator_events(second, &token, "alice") + .await + .expect("coordinator events"); + + core.stop_coordinator_events(first).await; + + assert!(core.user_auth().is_some()); + assert!(!coordinator.is_finished()); + core.leave("cleanup").await.expect("cleanup leave"); + tokio::time::timeout(Duration::from_secs(2), coordinator) + .await + .expect("coordinator socket closed") + .expect("fake coordinator task"); +} + +#[tokio::test] +async fn detached_connection_ignores_publish_options_from_its_sfu() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let (old, _old_sfu) = establish_fake(&core, generation).await; + let (current, _sfu) = establish_fake(&core, generation).await; + *core.connection.lock().await = Some(current); + let context = detached_context(&core, &old); + + connection::handle_event( + &context, + sfu_event::EventPayload::ChangePublishOptions(event::ChangePublishOptions { + publish_options: vec![models::PublishOption { + id: 99, + ..Default::default() + }], + reason: "old SFU".to_owned(), + }), + ) + .await + .expect("handle event"); + + let connection = core.connection.lock().await; + assert!( + connection + .as_ref() + .expect("current connection") + .publish_options + .is_empty() + ); +} + +#[tokio::test] +async fn detached_connection_still_completes_the_migration() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let (old, _old_sfu) = establish_fake(&core, generation).await; + let context = detached_context(&core, &old); + let (sender, receiver) = tokio::sync::oneshot::channel(); + core.install_migration_waiter(generation, sender) + .expect("migration waiter"); + + connection::handle_event( + &context, + sfu_event::EventPayload::ParticipantMigrationComplete( + event::ParticipantMigrationComplete {}, + ), + ) + .await + .expect("handle event"); + + receiver.await.expect("migration complete"); +} + #[tokio::test] async fn forced_strategy_failures_reach_timeout_and_refresh_over_http() { for strategy in [ @@ -196,7 +637,7 @@ async fn forced_strategy_failures_reach_timeout_and_refresh_over_http() { ReconnectStrategy::Rejoin, ReconnectStrategy::Migrate, ] { - let (base_url, request_rx, server) = refresh_server(); + let (base_url, request_rx, server) = one_shot_http_server("{}"); let core = test_core_with_config(ClientConfig { base_url, request_timeout: Duration::from_secs(1), @@ -286,6 +727,429 @@ async fn leave_cancels_reconnect_task_before_next_generation() { assert_eq!(core.active_reconnect_generation(), None); } +#[test] +fn state_events_arrive_in_the_order_of_the_state_changes() { + let core = test_core(); + let generation = core.begin_join().expect("join generation"); + let mut events = core.subscribe(); + let rounds = 20_000; + let barrier = Arc::new(std::sync::Barrier::new(3)); + let workers = [CallingState::Joined, CallingState::Reconnecting].map(|state| { + let core = core.clone(); + let barrier = barrier.clone(); + thread::spawn(move || { + for _ in 0..rounds { + barrier.wait(); + core.set_state_if_current(generation, state); + barrier.wait(); + } + }) + }); + + for round in 0..rounds { + barrier.wait(); + barrier.wait(); + let mut last = None; + while let Ok(event) = events.try_recv() { + if let CallEvent::CallingStateChanged(state) = event { + last = Some(state); + } + } + assert_eq!(last, Some(core.state()), "round {round}"); + } + for worker in workers { + worker.join().expect("state worker"); + } +} + +#[test] +fn setting_the_same_state_again_sends_no_event() { + let core = test_core(); + let generation = core.begin_join().expect("join generation"); + let mut events = core.subscribe(); + + assert!(core.set_state_if_current(generation, CallingState::Reconnecting)); + assert!(core.set_state_if_current(generation, CallingState::Reconnecting)); + + assert!(matches!( + events.try_recv(), + Ok(CallEvent::CallingStateChanged(CallingState::Reconnecting)) + )); + assert!(events.try_recv().is_err()); +} + +#[test] +fn join_start_sends_joining() { + let core = test_core(); + let mut events = core.subscribe(); + + core.begin_join().expect("join generation"); + + assert!(matches!( + events.try_recv(), + Ok(CallEvent::CallingStateChanged(CallingState::Joining)) + )); +} + +#[tokio::test] +async fn state_during_leave_matches_the_last_state_event() { + let core = test_core(); + let mut events = core.subscribe(); + core.begin_join().expect("join generation"); + let connection_slot = core.connection.lock().await; + let generation = core.generation(); + let leave_core = core.clone(); + let leave = tokio::spawn(async move { leave_core.leave("leave during join").await }); + wait_for( + Duration::from_secs(1), + || core.generation() != generation, + "leave cancels the join", + ) + .await; + + let mut last = None; + while let Ok(event) = events.try_recv() { + if let CallEvent::CallingStateChanged(state) = event { + last = Some(state); + } + } + assert_eq!(last, Some(core.state())); + drop(connection_slot); + leave.await.expect("leave task").expect("leave"); +} + +#[tokio::test] +async fn late_reconnect_task_does_not_count_toward_the_next_join() { + let core = test_core(); + let first = prepare_joined_core(&core, "alice"); + core.trigger_reconnect( + first, + ReconnectStrategy::Fast, + reconnect::REASON_ICE_UNSUPPORTED.to_owned(), + ); + core.leave("leave before the reconnect task runs") + .await + .expect("leave"); + let second = prepare_joined_core(&core, "alice"); + wait_for( + Duration::from_secs(1), + || core.runtime_task_snapshot().0 == 0, + "late reconnect task ends", + ) + .await; + + core.trigger_reconnect( + second, + ReconnectStrategy::Fast, + reconnect::REASON_ICE_UNSUPPORTED.to_owned(), + ); + wait_for( + Duration::from_secs(1), + || core.state() != CallingState::Joined, + "second reconnect starts", + ) + .await; + + assert_eq!(core.state(), CallingState::Reconnecting); + core.leave("cleanup").await.expect("cleanup leave"); +} + +#[tokio::test] +async fn token_load_that_ends_after_a_new_join_keeps_the_new_token() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let old_token = core + .user_token + .lock() + .unwrap_or_else(|error| error.into_inner()) + .clone(); + let weak_core = Arc::downgrade(&core); + let provider = move || { + let weak_core = weak_core.clone(); + let old_token = old_token.clone(); + async move { + // A new join starts in the same poll in which this load ends. + if let Some(core) = weak_core.upgrade() { + core.cancel_generation(); + *core + .user_token + .lock() + .unwrap_or_else(|error| error.into_inner()) = "new-join-token".to_owned(); + } + Ok(old_token) + } + }; + *core + .token_source + .lock() + .unwrap_or_else(|error| error.into_inner()) = + Some(UserTokenSource::Provider(Arc::new(provider))); + + let result = core.reload_user_token(generation).await; + + assert!(result.is_err()); + assert_eq!( + *core + .user_token + .lock() + .unwrap_or_else(|error| error.into_inner()), + "new-join-token" + ); +} + +#[tokio::test] +async fn work_that_ends_after_its_generation_changed_is_cancelled() { + let core = test_core(); + let generation = core.begin_join().expect("join generation"); + let work_core = core.clone(); + + let result = core + .while_generation(generation, async move { + work_core.cancel_generation(); + }) + .await; + + assert!(result.is_err()); +} + +#[tokio::test] +async fn work_for_a_stale_generation_never_runs() { + let core = test_core(); + let generation = core.begin_join().expect("join generation"); + core.cancel_generation(); + let ran = Arc::new(AtomicBool::new(false)); + let work_ran = ran.clone(); + + let result = core + .while_generation(generation, async move { + work_ran.store(true, Ordering::SeqCst); + }) + .await; + + assert!(result.is_err()); + assert!(!ran.load(Ordering::SeqCst)); +} + +#[test] +fn concurrent_reconnect_claims_have_one_winner() { + let core = test_core(); + let generation = core.begin_join().expect("join generation"); + let rounds = 5_000; + let claimers = 4; + let barrier = Arc::new(std::sync::Barrier::new(claimers + 1)); + let wins = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let workers: Vec<_> = (0..claimers) + .map(|_| { + let core = core.clone(); + let barrier = barrier.clone(); + let wins = wins.clone(); + thread::spawn(move || { + for _ in 0..rounds { + barrier.wait(); + if core.claim_reconnect(generation) { + wins.fetch_add(1, Ordering::SeqCst); + } + barrier.wait(); + } + }) + }) + .collect(); + + for round in 0..rounds { + barrier.wait(); + barrier.wait(); + assert_eq!(wins.swap(0, Ordering::SeqCst), 1, "round {round}"); + core.release_reconnect(generation); + } + for worker in workers { + worker.join().expect("claim worker"); + } +} + +#[test] +fn second_migration_waiter_for_a_generation_is_rejected() { + let core = test_core(); + let generation = core.begin_join().expect("join generation"); + let (first, _first_receiver) = tokio::sync::oneshot::channel(); + let (second, _second_receiver) = tokio::sync::oneshot::channel(); + core.install_migration_waiter(generation, first) + .expect("first migration waiter"); + + assert!(matches!( + core.install_migration_waiter(generation, second), + Err(RtcError::IllegalState(_)) + )); +} + +#[tokio::test] +async fn migration_waiter_of_a_new_generation_replaces_the_old_one() { + let core = test_core(); + let first = core.begin_join().expect("first generation"); + let (old_sender, old_receiver) = tokio::sync::oneshot::channel(); + core.install_migration_waiter(first, old_sender) + .expect("old migration waiter"); + core.leave("next generation").await.expect("leave"); + let second = core.begin_join().expect("second generation"); + let (sender, mut receiver) = tokio::sync::oneshot::channel(); + + core.install_migration_waiter(second, sender) + .expect("new migration waiter"); + + assert!(old_receiver.await.is_err()); + core.complete_migration(second); + assert_eq!(receiver.try_recv(), Ok(())); +} + +#[test] +fn user_query_needs_the_coordinator_connection_of_the_current_generation() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + *core + .coordinator_connection_id + .lock() + .unwrap_or_else(|error| error.into_inner()) = + Some((generation.wrapping_sub(1), "old-connection".to_owned())); + + assert!(core.user_request_query().is_none()); + assert!(core.user_auth().is_none()); +} + +#[test] +fn user_query_needs_a_connection_id_and_a_user_id() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let set_connection_id = |id: &str| { + *core + .coordinator_connection_id + .lock() + .unwrap_or_else(|error| error.into_inner()) = Some((generation, id.to_owned())); + }; + + set_connection_id(""); + assert!(core.user_request_query().is_none()); + + set_connection_id("connection-1"); + core.join_data + .lock() + .unwrap_or_else(|error| error.into_inner()) + .user_id + .clear(); + assert!(core.user_request_query().is_none()); +} + +#[tokio::test] +async fn join_attempt_stores_the_stats_options_and_the_sfu_session() { + let (credentials, mut sfu) = fake_sfu(true).await; + let join_response = json!({ + "credentials": { + "server": { + "edge_name": credentials.server.edge_name, + "url": credentials.server.url, + "ws_endpoint": credentials.server.ws_endpoint, + }, + "token": credentials.token, + "ice_servers": [], + }, + "stats_options": { "reporting_interval_ms": 1234, "enable_rtc_stats": true }, + }) + .to_string(); + let (base_url, _requests, server) = one_shot_http_server(&join_response); + let core = test_core_with_config(ClientConfig { + base_url, + ..ClientConfig::default() + }); + let generation = prepare_joined_core(&core, "alice"); + *core + .coordinator_connection_id + .lock() + .unwrap_or_else(|error| error.into_inner()) = Some((generation, "connection-1".to_owned())); + let token = core.current_user_token().expect("user token"); + + core.clone() + .join_once(JoinOnceOptions { + user_token: &token, + request: &JoinCallRequest::default(), + attempt: 0, + strategy: ReconnectStrategy::Fast, + reconnect_details: None, + generation, + session_id: None, + retain_old: false, + }) + .await + .map_err(|(error, _)| error) + .expect("join attempt"); + + let stats_options = core.stats_options(); + assert_eq!(stats_options.reporting_interval_ms, 1234); + assert!(stats_options.enable_rtc_stats); + let Some(event::sfu_request::RequestPayload::JoinRequest(join_request)) = + sfu.recv().await.expect("SFU join request").request_payload + else { + panic!("first SFU request is not a join request"); + }; + assert_eq!(core.session_id().await, Some(join_request.session_id)); + core.leave("test leave").await.expect("leave"); + assert_eq!(core.session_id().await, None); + server.join().expect("coordinator server"); +} + +#[tokio::test] +async fn only_the_stored_connection_of_the_current_generation_is_current() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let (old, _old_sfu) = establish_fake(&core, generation).await; + let (current, _sfu) = establish_fake(&core, generation).await; + let (old_epoch, epoch) = (old.epoch, current.epoch); + *core.connection.lock().await = Some(current); + + assert!(core.is_connection_current(generation, epoch).await); + assert!(!core.is_connection_current(generation, old_epoch).await); + core.cancel_generation(); + assert!(!core.is_connection_current(generation, epoch).await); + drop(old); +} + +#[tokio::test] +async fn video_track_with_a_non_video_type_is_not_published() { + let core = test_core(); + let generation = prepare_joined_core(&core, "alice"); + let (connection, _sfu) = establish_fake(&core, generation).await; + *core.connection.lock().await = Some(connection); + core.own_capabilities + .lock() + .unwrap_or_else(|error| error.into_inner()) + .insert("send-video".to_owned()); + let track = + LocalVideoTrack::h264_with_config(LocalVideoTrackConfig::default().server_managed()) + .expect("video track"); + + let result = core + .publish(LocalTrack::Video { + track, + track_type: TrackType::Unspecified, + }) + .await; + + assert!( + matches!(result, Err(RtcError::IllegalState(_))), + "{result:?}" + ); + assert!(core.media.lock().await.publications.is_empty()); +} + +#[tokio::test] +async fn unspecified_track_type_cannot_be_muted() { + let core = test_core(); + + let result = core.set_track_muted(TrackType::Unspecified, true).await; + + assert!( + matches!(result, Err(RtcError::IllegalState(_))), + "{result:?}" + ); +} + #[test] fn stale_reconnect_completion_does_not_release_the_current_generation() { let core = test_core(); @@ -355,12 +1219,12 @@ fn participant_refresh_replaces_published_track_state() { published_tracks: vec![TrackType::Audio as i32, TrackType::Video as i32], ..Default::default() }; - core.roster_upsert(&participant); + core.upsert_participant(&participant); participant.published_tracks = vec![TrackType::Audio as i32]; - core.roster_upsert(&participant); + core.upsert_participant(&participant); - let roster = core.roster.lock().unwrap_or_else(|e| e.into_inner()); - let entry = roster.get("session-a").expect("participant"); + let participants = core.participants.lock().unwrap_or_else(|e| e.into_inner()); + let entry = participants.get("session-a").expect("participant"); assert_eq!(entry.published.len(), 1); assert!(entry.published.contains(&(TrackType::Audio as i32))); } @@ -820,7 +1684,7 @@ fn stop_state_is_retryable_until_mute_sync_commits() { media.set_status(&track_id, PublicationStatus::Published); // A stop marks the track `PendingStopMute`: it leaves the active set (so it - // is not re-announced on reconnect) but stays in the roster until the mute + // is not re-announced on reconnect) but stays in the publication list until the mute // RPC commits, so a failed mute can be retried without losing the track. media.set_status(&track_id, PublicationStatus::PendingStopMute); assert!(media.active_tracks().is_empty()); @@ -904,22 +1768,65 @@ async fn live_twirp_ice_trickle_framing_accepted() { } } -/// Live proof that reconnect surfaces media-restoration failures. -/// -/// A reconnect that fails while restoring the publisher must surface the error -/// instead of silently reporting `Joined`. This joins a live call, publishes -/// audio, injects a one-shot failure at the REJOIN published-restore hook, forces -/// a REJOIN, and proves (1) the restore hook was reached and the injected failure -/// propagated out of the attempt, and (2) the driver retried and recovered to -/// `Joined` rather than leaving a failed attempt marked joined. #[tokio::test] -async fn live_forced_media_restore_failure_is_surfaced_and_recovers() { +async fn live_rejoin_restore_failure_is_surfaced_and_recovers() { + assert_live_restore_failure_is_surfaced_and_recovers(ReconnectStrategy::Rejoin).await; +} + +#[tokio::test] +async fn live_fast_restore_failure_is_surfaced_and_recovers() { + assert_live_restore_failure_is_surfaced_and_recovers(ReconnectStrategy::Fast).await; +} + +#[tokio::test] +async fn live_migrate_restore_failure_is_surfaced_and_recovers() { + assert_live_restore_failure_is_surfaced_and_recovers(ReconnectStrategy::Migrate).await; +} + +#[tokio::test] +async fn live_fast_reconnect_during_ice_gathering_succeeds_at_once() { + let Some((attempts, final_state)) = live_forced_reconnect(ReconnectStrategy::Fast, None).await + else { + return; + }; + assert_eq!(attempts, [ReconnectStrategy::Fast]); + assert_eq!(final_state, CallingState::Joined); +} + +/// A reconnect that fails once after the published tracks are restored must +/// end that attempt, and a retry must bring the call back to `Joined`. +async fn assert_live_restore_failure_is_surfaced_and_recovers(strategy: ReconnectStrategy) { + let Some((attempts, final_state)) = + live_forced_reconnect(strategy, Some(ReconnectFaultPoint::AfterPublishedRestore)).await + else { + return; + }; + assert_eq!(attempts.first(), Some(&strategy)); + assert!( + attempts.len() >= 2, + "the restore failure did not end the {strategy:?} attempt" + ); + assert_eq!( + final_state, + CallingState::Joined, + "reconnect did not recover after a surfaced media-restoration failure" + ); +} + +/// Joins a live call, publishes audio, and forces a `strategy` reconnect at +/// once. The reconnect fails once at `fault`, if given. Returns the reconnect +/// attempts and the state after the reconnect settles, or `None` without +/// credentials. +async fn live_forced_reconnect( + strategy: ReconnectStrategy, + fault: Option, +) -> Option<(Vec, CallingState)> { let _ = dotenvy::dotenv(); let key = std::env::var("STREAM_API_KEY").unwrap_or_default(); let secret = std::env::var("STREAM_API_SECRET").unwrap_or_default(); if key.is_empty() || secret.is_empty() { - eprintln!("SKIP: STREAM creds absent; skipping live media-restore failure probe"); - return; + eprintln!("SKIP: STREAM creds absent; skipping live forced reconnect"); + return None; } let stream = crate::Stream::new(key, secret).expect("client"); @@ -971,40 +1878,44 @@ async fn live_forced_media_restore_failure_is_surfaced_and_recovers() { let generation = core.lifecycle_snapshot().1; let probe = Arc::new(ReconnectProbe::default()); - probe.fail_once( - ReconnectStrategy::Rejoin, - ReconnectFaultPoint::AfterPublishedRestore, - ); + if let Some(point) = fault { + probe.fail_once(strategy, point); + } core.install_reconnect_probe(probe.clone()); - core.trigger_reconnect( - generation, - ReconnectStrategy::Rejoin, - "forced media restore failure".to_owned(), - ); + core.trigger_reconnect(generation, strategy, "forced reconnect".to_owned()); - let reached = |probe: &Arc| { - probe - .restores - .lock() - .unwrap_or_else(|error| error.into_inner()) - .iter() - .any(|(strategy, point)| { - *strategy == ReconnectStrategy::Rejoin - && *point == ReconnectFaultPoint::AfterPublishedRestore - }) - }; wait_for( Duration::from_secs(45), - || reached(&probe), - "forced restore fault reached", + || { + !probe + .attempts + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_empty() + }, + "reconnect started", ) .await; + if let Some(fault) = fault { + wait_for( + Duration::from_secs(45), + || { + probe + .restores + .lock() + .unwrap_or_else(|error| error.into_inner()) + .contains(&(strategy, fault)) + }, + "forced fault reached", + ) + .await; + } wait_for( Duration::from_secs(45), || { matches!( core.state(), - CallingState::Joined | CallingState::ReconnectingFailed + CallingState::Joined | CallingState::ReconnectingFailed | CallingState::Left ) }, "reconnect settled", @@ -1012,12 +1923,12 @@ async fn live_forced_media_restore_failure_is_surfaced_and_recovers() { .await; feeder.abort(); - let restores = probe - .restores + let attempts = probe + .attempts .lock() .unwrap_or_else(|error| error.into_inner()) .clone(); - (restores, core.state()) + (attempts, core.state()) }) .await; @@ -1026,18 +1937,10 @@ async fn live_forced_media_restore_failure_is_surfaced_and_recovers() { .delete(crate::models::DeleteCallRequest { hard: Some(true) }) .await; - let (restores, final_state) = outcome.expect("media-restore test timed out (120s guard)"); - eprintln!("RESTORE FAILURE: restores={restores:?} final_state={final_state:?}"); - assert!( - restores.iter().any(|(strategy, point)| { - *strategy == ReconnectStrategy::Rejoin - && *point == ReconnectFaultPoint::AfterPublishedRestore - }), - "forced REJOIN published-restore failure was never reached/surfaced" - ); - assert_eq!( - final_state, - CallingState::Joined, - "reconnect did not recover after a surfaced media-restoration failure" + let outcome = outcome.expect("forced reconnect timed out (120s guard)"); + eprintln!( + "FORCED RECONNECT {strategy:?} fault={fault:?}: attempts={:?} final_state={:?}", + outcome.0, outcome.1 ); + Some(outcome) } diff --git a/src/rtc/mod.rs b/src/rtc/mod.rs index 2ca636e..929df56 100644 --- a/src/rtc/mod.rs +++ b/src/rtc/mod.rs @@ -71,115 +71,3 @@ pub use tracks::{ audio_level_dbov, }; pub use video_frame::VideoFrame; - -#[cfg(test)] -mod tests { - use super::error::RtcError; - use super::proto::event::{ - HealthCheckRequest, JoinRequest, SfuEvent, SfuRequest, sfu_event, sfu_request, - }; - use super::proto::{models, signal}; - use prost::Message; - - #[test] - fn sfu_request_join_round_trips() { - let request = SfuRequest { - request_payload: Some(sfu_request::RequestPayload::JoinRequest(JoinRequest { - token: "tok".to_owned(), - session_id: "sess-123".to_owned(), - subscriber_sdp: "v=0".to_owned(), - client_details: Some(super::identity::client_details()), - ..Default::default() - })), - }; - - let bytes = request.encode_to_vec(); - let decoded = SfuRequest::decode(bytes.as_slice()).expect("decode SfuRequest"); - assert_eq!(request, decoded); - - match decoded.request_payload { - Some(sfu_request::RequestPayload::JoinRequest(join)) => { - assert_eq!(join.session_id, "sess-123"); - let sdk = join.client_details.and_then(|d| d.sdk).expect("sdk"); - // AGENTS.md hard rule: never report Go to the SFU. - assert_ne!(sdk.r#type, models::SdkType::Go as i32); - } - other => panic!("unexpected payload: {other:?}"), - } - } - - #[test] - fn sfu_request_health_check_round_trips() { - let request = SfuRequest { - request_payload: Some(sfu_request::RequestPayload::HealthCheckRequest( - HealthCheckRequest {}, - )), - }; - let bytes = request.encode_to_vec(); - let decoded = SfuRequest::decode(bytes.as_slice()).expect("decode"); - assert_eq!(request, decoded); - } - - #[test] - fn sfu_event_error_round_trips() { - let event = SfuEvent { - event_payload: Some(sfu_event::EventPayload::Error(super::proto::event::Error { - error: Some(models::Error { - code: models::ErrorCode::ParticipantSignalLost as i32, - message: "signal lost".to_owned(), - should_retry: true, - }), - reconnect_strategy: models::WebsocketReconnectStrategy::Rejoin as i32, - })), - }; - let bytes = event.encode_to_vec(); - let decoded = SfuEvent::decode(bytes.as_slice()).expect("decode SfuEvent"); - assert_eq!(event, decoded); - } - - #[test] - fn set_publisher_request_round_trips() { - let request = signal::SetPublisherRequest { - sdp: "v=0".to_owned(), - session_id: "sess".to_owned(), - tracks: vec![], - }; - let bytes = request.encode_to_vec(); - let decoded = signal::SetPublisherRequest::decode(bytes.as_slice()).expect("decode"); - assert_eq!(request, decoded); - } - - #[test] - fn from_signal_error_maps_only_real_codes() { - // UNSPECIFIED (and absent) is success. - assert!(RtcError::from_signal_error(None).is_ok()); - assert!( - RtcError::from_signal_error(Some(models::Error { - code: models::ErrorCode::Unspecified as i32, - message: String::new(), - should_retry: false, - })) - .is_ok() - ); - // A real code becomes an error. - let err = RtcError::from_signal_error(Some(models::Error { - code: models::ErrorCode::ParticipantSignalLost as i32, - message: "boom".to_owned(), - should_retry: true, - })) - .expect_err("should be an error"); - assert!(matches!(err, RtcError::Signal { .. })); - } - - #[test] - fn ws_auth_message_serializes_video_product() { - let auth = - super::WsAuthMessage::video("jwt-token", super::ConnectUserDetails::new("agent")); - let json = serde_json::to_value(&auth).expect("serialize"); - assert_eq!(json["token"], "jwt-token"); - assert_eq!(json["user_details"]["id"], "agent"); - assert_eq!(json["products"][0], "video"); - // Optional user fields are omitted when unset. - assert!(json["user_details"].get("name").is_none()); - } -} diff --git a/src/rtc/peer/publisher.rs b/src/rtc/peer/publisher.rs index d7df3d2..587257b 100644 --- a/src/rtc/peer/publisher.rs +++ b/src/rtc/peer/publisher.rs @@ -8,6 +8,7 @@ use std::collections::{HashMap, HashSet}; use std::sync::Arc; +use std::time::Duration; use tokio::task::JoinHandle; use webrtc::peer_connection::RTCPeerConnection; @@ -15,7 +16,7 @@ use webrtc::peer_connection::sdp::sdp_type::RTCSdpType; use webrtc::peer_connection::sdp::session_description::RTCSessionDescription; use webrtc::peer_connection::signaling_state::RTCSignalingState; -use crate::rtc::error::{NegotiationError, Result, RtcError}; +use crate::rtc::error::{NegotiationError, Result, RtcError, SfuTimeoutError}; use crate::rtc::proto::models::{PublishOption, TrackInfo, TrackType}; use crate::rtc::proto::signal::SetPublisherRequest; use crate::rtc::sfu::signal::SignalClient; @@ -81,10 +82,31 @@ pub(crate) async fn restart_ice( if tracks.is_empty() { return Ok(()); } - publisher.restart_ice().await.map_err(neg)?; + // webrtc-rs rejects an ICE restart while it gathers candidates. + let deadline = tokio::time::Instant::now() + ICE_GATHERING_TIMEOUT; + loop { + match publisher.restart_ice().await { + Err(webrtc::Error::Ice(webrtc::ice::Error::ErrRestartWhenGathering)) => { + if tokio::time::Instant::now() >= deadline { + return Err(RtcError::Timeout(SfuTimeoutError::new( + "publisher ICE gathering", + ICE_GATHERING_TIMEOUT, + ))); + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + result => { + result.map_err(neg)?; + break; + } + } + } negotiate_publish(publisher, signal, session_id, tracks, publish_options).await } +/// Longer than the 5 s STUN timeout of webrtc-rs candidate gathering. +const ICE_GATHERING_TIMEOUT: Duration = Duration::from_secs(10); + fn neg(e: impl std::fmt::Display) -> RtcError { RtcError::Negotiation(NegotiationError(e.to_string())) } diff --git a/src/rtc/reconnect.rs b/src/rtc/reconnect.rs index 7b85ab9..cc11aa0 100644 --- a/src/rtc/reconnect.rs +++ b/src/rtc/reconnect.rs @@ -3,11 +3,11 @@ //! Ported from JS `Call.ts` + `coordinator/connection/utils.ts`. Everything in //! this module is deterministic (or jitter-only) and side-effect free so it can //! be unit-tested without a live SFU: backoff intervals, the rejoin rate -//! limiter, the ICE / negotiation failure caps, the join-retry decision, and +//! limiter, the ICE / negotiation failure limits, the join-retry decision, and //! the FAST→REJOIN escalation rule. The orchestration that *acts* on these //! decisions lives in [`super::join`]. -use std::collections::VecDeque; +use std::collections::{HashMap, VecDeque}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use super::proto::models::WebsocketReconnectStrategy; @@ -196,17 +196,17 @@ impl SlidingWindowRateLimiter { } } -/// Tracks the failure caps that force the reconnect loop to give up (JS +/// Tracks the failure limits that force the reconnect loop to give up (JS /// `iceFailuresWithoutConnect` / `consecutiveNegotiationFailures`). #[derive(Debug, Clone)] -pub struct FailureCaps { +pub struct FailureLimits { ice_failures_without_connect: u32, consecutive_negotiation_failures: u32, max_ice_failures: u32, max_consecutive_negotiation: u32, } -impl Default for FailureCaps { +impl Default for FailureLimits { fn default() -> Self { Self { ice_failures_without_connect: 0, @@ -217,8 +217,8 @@ impl Default for FailureCaps { } } -impl FailureCaps { - /// Record an ICE-never-connected failure. Returns `true` when the cap (2) is +impl FailureLimits { + /// Record an ICE-never-connected failure. Returns `true` when the limit (2) is /// reached and the caller must `leave` with `webrtc_unsupported_network`. pub fn record_ice_never_connected(&mut self) -> bool { self.ice_failures_without_connect += 1; @@ -230,7 +230,7 @@ impl FailureCaps { self.ice_failures_without_connect = 0; } - /// Record a negotiation failure. Returns `true` when the cap (3) is reached + /// Record a negotiation failure. Returns `true` when the limit (3) is reached /// and the caller must `leave` with `repeated_negotiation_failures`. pub fn record_negotiation_failure(&mut self) -> bool { self.consecutive_negotiation_failures += 1; @@ -243,6 +243,32 @@ impl FailureCaps { } } +/// SFUs that failed a rejoin during one reconnect. Two failures confirm an SFU +/// as bad; a join error code confirms it at once. +#[derive(Debug, Default)] +pub(crate) struct SfuRejoinFailures { + counts: HashMap, + confirmed: Vec, +} + +impl SfuRejoinFailures { + pub(crate) fn record(&mut self, edge: &str, force_switch: bool) { + let count = self.counts.entry(edge.to_owned()).or_insert(0); + *count = count.saturating_add(1); + if force_switch { + *count = (*count).max(2); + } + if *count >= 2 && !self.confirmed.iter().any(|known| known == edge) { + self.confirmed.push(edge.to_owned()); + } + } + + /// Confirmed bad SFUs, in the order they were confirmed. + pub(crate) fn confirmed(&self) -> &[String] { + &self.confirmed + } +} + /// Decide the strategy for the *next* reconnect attempt after the current one /// failed (JS `shouldRejoin` escalation). Once we fall back to `REJOIN` we stay /// there. @@ -361,23 +387,41 @@ mod tests { } #[test] - fn ice_cap_trips_on_second_failure() { - let mut caps = FailureCaps::default(); - assert!(!caps.record_ice_never_connected()); - assert!(caps.record_ice_never_connected()); + fn ice_limit_is_reached_on_second_failure() { + let mut limits = FailureLimits::default(); + assert!(!limits.record_ice_never_connected()); + assert!(limits.record_ice_never_connected()); // reset clears it - caps.reset_ice(); - assert!(!caps.record_ice_never_connected()); + limits.reset_ice(); + assert!(!limits.record_ice_never_connected()); + } + + #[test] + fn sfu_is_confirmed_bad_after_two_rejoin_failures() { + let mut failures = SfuRejoinFailures::default(); + failures.record("sfu-a", false); + assert!(failures.confirmed().is_empty()); + failures.record("sfu-a", false); + assert_eq!(failures.confirmed(), ["sfu-a"]); + } + + #[test] + fn join_error_code_confirms_the_sfu_at_once() { + let mut failures = SfuRejoinFailures::default(); + failures.record("sfu-a", true); + failures.record("sfu-b", true); + failures.record("sfu-a", true); + assert_eq!(failures.confirmed(), ["sfu-a", "sfu-b"]); } #[test] - fn negotiation_cap_trips_on_third_failure() { - let mut caps = FailureCaps::default(); - assert!(!caps.record_negotiation_failure()); - assert!(!caps.record_negotiation_failure()); - assert!(caps.record_negotiation_failure()); - caps.reset_negotiation(); - assert!(!caps.record_negotiation_failure()); + fn negotiation_limit_is_reached_on_third_failure() { + let mut limits = FailureLimits::default(); + assert!(!limits.record_negotiation_failure()); + assert!(!limits.record_negotiation_failure()); + assert!(limits.record_negotiation_failure()); + limits.reset_negotiation(); + assert!(!limits.record_negotiation_failure()); } #[test] diff --git a/src/rtc/subscriptions.rs b/src/rtc/subscriptions.rs index 182f803..44ad069 100644 --- a/src/rtc/subscriptions.rs +++ b/src/rtc/subscriptions.rs @@ -4,9 +4,9 @@ //! The SFU never auto-forwards media — without an explicit subscription no //! `on_track` fires (JS `DynascaleManager`, stream-py `SubscriptionManager`, //! videosdk `UpdateSubscriptions`). This module holds the declarative policy; -//! [`RtcCore`](super::join::RtcCore) turns it plus the live participant roster -//! into the concrete `TrackSubscriptionDetails` list and (re)sends it whenever -//! the roster changes. +//! [`RtcCore`](super::join::RtcCore) turns it plus the live participants into +//! the concrete `TrackSubscriptionDetails` list and (re)sends it whenever the +//! participants change. //! //! The default policy subscribes to remote **audio** only (the backend-bot //! default); video and screen-share are opt-in. diff --git a/tests/rtc_join.rs b/tests/rtc_join.rs index 578030b..ca839df 100644 --- a/tests/rtc_join.rs +++ b/tests/rtc_join.rs @@ -121,7 +121,7 @@ async fn two_sessions_join_and_observe_each_other() { .await .expect("session B join failed"); - // A should see B arrive as an event; B should see A via the initial roster. + // A should see B arrive as an event; B should see A via the initial participants. let saw_b = observe_participant(rx_a, user_b.clone(), Duration::from_secs(30)); let saw_a = observe_participant(rx_b, user_a.clone(), Duration::from_secs(30)); let (saw_b, saw_a) = tokio::join!(saw_b, saw_a);