diff --git a/backend/src/api/endpoints/v1_api.rs b/backend/src/api/endpoints/v1_api.rs index aab542b3a..cd5934d87 100644 --- a/backend/src/api/endpoints/v1_api.rs +++ b/backend/src/api/endpoints/v1_api.rs @@ -186,3 +186,85 @@ pub fn v1_api_register( .nest(&api_prefix, public_router) .nest(&api_prefix, router) } + +#[cfg(test)] +mod tests { + use super::create_status_check; + use crate::{ + api::model::{create_test_app_state, ConnectionKind, ConnectionParams}, + auth::Fingerprint, + model::Config, + }; + use shared::{ + model::{PlaylistItemType, StreamChannel, XtreamCluster}, + utils::Internable, + }; + use std::{borrow::Cow, net::SocketAddr}; + + #[tokio::test] + async fn status_snapshot_removes_released_direct_series_stream() { + let app_state = create_test_app_state(Config::default()); + let addr: SocketAddr = "127.0.0.1:55070".parse().expect("test address"); + let fingerprint = Fingerprint::new("status-series".to_string(), "127.0.0.1".to_string(), addr); + let channel = StreamChannel { + target_id: 1, + virtual_id: 70, + provider_id: 1, + input_name: "input".intern(), + item_type: PlaylistItemType::Series, + cluster: XtreamCluster::Series, + group: "Series".intern(), + title: "Episode".intern(), + url: "http://provider.example/series/70.mkv".intern(), + shared: false, + shared_joined_existing: None, + shared_stream_id: None, + technical: None, + epg_channel_id: None, + epg_reference_ts: None, + }; + app_state.connection_manager.add_connection(&addr).await; + let registered = app_state + .connection_manager + .update_connection(ConnectionParams { + meter_uid: 0, + username: "status-user", + max_connections: 1, + soft_connections: 0, + connection_kind: ConnectionKind::Normal, + priority: 0, + soft_priority: 0, + fingerprint: &fingerprint, + provider: "provider".intern(), + stream_channel: &channel, + user_agent: Cow::Borrowed("ua"), + session_token: None, + }) + .await + .expect("direct Series stream should register"); + + let active = create_status_check(&app_state).await; + assert_eq!(active.active_users, 1); + assert_eq!(active.active_user_connections, 1); + assert_eq!(active.active_user_streams.len(), 1); + assert_eq!(active.active_user_streams[0].uid, registered.uid); + + app_state + .active_users + .release_stream_by_uid(&addr, registered.uid) + .await + .expect("registered stream should release"); + let clean = create_status_check(&app_state).await; + assert_eq!(clean.active_users, 0); + assert_eq!(clean.active_user_connections, 0); + assert!(clean.active_user_streams.is_empty()); + assert_eq!( + clean + .active_provider_connections + .unwrap_or_default() + .values() + .sum::(), + 0 + ); + } +} diff --git a/backend/src/api/endpoints/websocket_api.rs b/backend/src/api/endpoints/websocket_api.rs index 22e6744ea..7135547c9 100644 --- a/backend/src/api/endpoints/websocket_api.rs +++ b/backend/src/api/endpoints/websocket_api.rs @@ -119,6 +119,26 @@ fn websocket_can_receive_runtime_events(mem: &ProtocolHandlerMemory, event: &Eve } } +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +enum MainEventReceiveErrorAction { + Continue, + ResyncStatus, + Terminate, +} + +fn main_event_receive_error_action( + handler: &ProtocolHandler, + error: &tokio::sync::broadcast::error::RecvError, +) -> MainEventReceiveErrorAction { + match error { + tokio::sync::broadcast::error::RecvError::Lagged(_) if matches!(handler, ProtocolHandler::Default(mem) if mem.permissions.contains(Permission::SystemRead)) => { + MainEventReceiveErrorAction::ResyncStatus + } + tokio::sync::broadcast::error::RecvError::Lagged(_) => MainEventReceiveErrorAction::Continue, + tokio::sync::broadcast::error::RecvError::Closed => MainEventReceiveErrorAction::Terminate, + } +} + fn get_secret_key(app_state: &AppState, auth: bool) -> Option> { if !auth { return None; @@ -457,10 +477,30 @@ async fn handle_socket(mut socket: WebSocket, app_state: Arc, auth_req break; } } - Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => { - trace!("Main websocket event receiver lagged by {skipped} messages"); + Err(error) => { + if let tokio::sync::broadcast::error::RecvError::Lagged(skipped) = &error { + trace!("Main websocket event receiver lagged by {skipped} messages"); + } + match main_event_receive_error_action(&handler, &error) { + MainEventReceiveErrorAction::Continue => {} + MainEventReceiveErrorAction::ResyncStatus => { + // Drop retained pre-snapshot deltas so they cannot be replayed after the authoritative state. + event_rx = app_state.event_manager.get_event_channel(); + let status = create_status_check(&app_state).await; + if let Err(e) = send_event_response( + &mut socket, + ProtocolMessage::StatusResponse(status), + "Status resync after lagged main websocket event receiver", + ) + .await + { + trace!("Failed to send ws status resync: {e}"); + break; + } + } + MainEventReceiveErrorAction::Terminate => break, + } } - Err(tokio::sync::broadcast::error::RecvError::Closed) => break, } } @@ -514,13 +554,57 @@ async fn handle_user_action(app_state: &Arc, cmd: UserCommand) -> bool #[cfg(test)] mod tests { - use super::websocket_can_receive_runtime_events; + use super::{main_event_receive_error_action, websocket_can_receive_runtime_events, MainEventReceiveErrorAction}; use crate::api::model::EventMessage; use shared::model::{ DownloadsDelta, DownloadsResponse, FileDownloadDto, LibraryScanProgressEvent, LibraryScanSummary, - LibraryScanSummaryStatus, Permission, PlaylistUpdateProgressEvent, ProtocolHandlerMemory, TaskKindDto, - TaskPriorityDto, TransferStatusDto, UserRole, + LibraryScanSummaryStatus, Permission, PlaylistUpdateProgressEvent, ProtocolHandler, ProtocolHandlerMemory, + TaskKindDto, TaskPriorityDto, TransferStatusDto, UserRole, PROTOCOL_VERSION, }; + use tokio::sync::broadcast::error::RecvError; + + #[test] + fn lagged_main_event_receiver_resyncs_authorized_system_reader() { + let handler = ProtocolHandler::Default(ProtocolHandlerMemory { + token: Some("token".to_string()), + permissions: Permission::SystemRead.into(), + role: UserRole::User, + ..ProtocolHandlerMemory::default() + }); + + assert_eq!( + main_event_receive_error_action(&handler, &RecvError::Lagged(3)), + MainEventReceiveErrorAction::ResyncStatus + ); + } + + #[test] + fn closed_main_event_receiver_terminates() { + let handler = ProtocolHandler::Default(ProtocolHandlerMemory { + permissions: Permission::SystemRead.into(), + ..ProtocolHandlerMemory::default() + }); + + assert_eq!( + main_event_receive_error_action(&handler, &RecvError::Closed), + MainEventReceiveErrorAction::Terminate + ); + } + + #[test] + fn lagged_main_event_receiver_does_not_resync_before_handshake_or_authorization() { + let version_handler = ProtocolHandler::Version(PROTOCOL_VERSION); + let unauthorized_handler = ProtocolHandler::Default(ProtocolHandlerMemory::default()); + + assert_eq!( + main_event_receive_error_action(&version_handler, &RecvError::Lagged(1)), + MainEventReceiveErrorAction::Continue + ); + assert_eq!( + main_event_receive_error_action(&unauthorized_handler, &RecvError::Lagged(1)), + MainEventReceiveErrorAction::Continue + ); + } #[test] fn test_websocket_runtime_events_allowed_for_admin() { diff --git a/backend/src/api/model/active_user_manager.rs b/backend/src/api/model/active_user_manager.rs index baed302b6..36ddf297e 100644 --- a/backend/src/api/model/active_user_manager.rs +++ b/backend/src/api/model/active_user_manager.rs @@ -441,7 +441,7 @@ struct AdaptiveExpiryKey { pub struct ReleasedConnection { pub addr_removed: bool, pub removed_streams: Vec, - pub disconnected_user: Option, + pub disconnected_users: Vec, } pub struct ActiveUserConnectionParams<'a> { @@ -729,17 +729,14 @@ impl ActiveUserManager { async fn log_active_user(&self) { let is_log_user_enabled = self.is_log_user_enabled(); - // Skip the full connection-map snapshot + event send entirely when logging - // is disabled — this runs on every connection add/release and dominates - // lock contention on the active-user path at high segment rates. - if !is_log_user_enabled { - return; - } let (user_count, user_connection_count) = { self.active_users_and_connections().await }; self.event_manager.send_event(EventMessage::ActiveUser(ActiveUserConnectionChange::Connections( user_count, user_connection_count, ))); + if !is_log_user_enabled { + return; + } let last_user_count = self.last_logged_user_count.load(Ordering::Relaxed); let last_connection_count = self.last_logged_user_connection_count.load(Ordering::Relaxed); if last_user_count != user_count || last_connection_count != user_connection_count { @@ -786,22 +783,29 @@ impl ActiveUserManager { let (removed_stream, username, expiry_entry, preserved_update, connection_changed, promotion, divergence_snapshot) = { let mut user_connections = self.connections.write().await; - let username = user_connections - .key_by_addr - .get(addr) - .filter(|reg| !reg.username.is_empty()) - .map(|reg| reg.username.clone()) - .or_else(|| { - stream_uid.and_then(|uid| { + let username = match stream_uid { + Some(uid) => user_connections.by_key.iter().find_map(|(username, connection_data)| { + connection_data + .streams + .iter() + .any(|stream| !stream.preserved && stream.uid == uid && stream.addr == *addr) + .then(|| username.clone()) + }), + None => user_connections + .key_by_addr + .get(addr) + .filter(|reg| !reg.username.is_empty()) + .map(|reg| reg.username.clone()) + .or_else(|| { user_connections.by_key.iter().find_map(|(username, connection_data)| { connection_data .streams .iter() - .any(|stream| stream.uid == uid && stream.addr == *addr) + .any(|stream| !stream.preserved && stream.addr == *addr) .then(|| username.clone()) }) - }) - })?; + }), + }?; let mut removed_stream = None; let mut expiry_entry = None; @@ -920,31 +924,45 @@ impl ActiveUserManager { #[allow(clippy::too_many_lines)] async fn release_connection_inner(&self, addr: &SocketAddr, preserve_session_streams: bool) -> ReleasedConnection { - let (addr_removed, disconnected_user, removed_streams, expiry_entries, preserved_updates, promotions) = { + let ( + addr_removed, + connection_count_changed, + disconnected_users, + removed_streams, + expiry_entries, + preserved_updates, + promotions, + ) = { let mut user_connections = self.connections.write().await; let registration = user_connections.key_by_addr.remove(addr); let had_registration = registration.is_some(); - let fallback_username = if had_registration { - None - } else { - user_connections.by_key.iter().find_map(|(username, connection_data)| { - connection_data - .streams - .iter() - .any(|stream| stream.addr == *addr) - .then(|| username.clone()) - }) - }; + let mut disconnected_users = registration + .map(|registration| registration.username) + .filter(|username| !username.is_empty()) + .into_iter() + .collect::>(); + disconnected_users.extend( + user_connections + .by_key + .iter() + .filter(|(_, connection_data)| { + connection_data.has_session_addr(addr) + || connection_data.streams.iter().any(|stream| stream.addr == *addr) + }) + .map(|(username, _)| username.clone()), + ); + disconnected_users.sort_unstable(); + disconnected_users.dedup(); - let username = registration.map(|registration| registration.username).or(fallback_username); - - if let Some(username) = username { - let mut removed_streams = Vec::new(); - let mut expiry_entries = Vec::new(); - let mut preserved_updates = Vec::new(); - let mut promotions = Vec::new(); - if let Some(connection_data) = user_connections.by_key.get_mut(&username) { + let mut removed_streams = Vec::new(); + let mut expiry_entries = Vec::new(); + let mut preserved_updates = Vec::new(); + let mut promotions = Vec::new(); + let mut connection_count_changed = false; + for username in &disconnected_users { + if let Some(connection_data) = user_connections.by_key.get_mut(username) { + let previous_connection_count = connection_data.connections; let migrated_session_addrs = connection_data.release_addr_from_sessions(addr); let mut remaining_streams = Vec::with_capacity(connection_data.streams.len()); let mut released_kinds = Vec::new(); @@ -965,7 +983,7 @@ impl ActiveUserManager { remaining_streams.push(stream_info); } else if preserve_session_streams && Self::should_preserve_session_stream(&stream_info) { if let Some(entry) = - self.build_preserved_stream_expiry(&username, &stream_info, &connection_data.sessions) + self.build_preserved_stream_expiry(username, &stream_info, &connection_data.sessions) { if let Some(kind) = connection_data.stream_kinds.remove(&stream_info.uid) { released_kinds.push(kind); @@ -1019,7 +1037,7 @@ impl ActiveUserManager { if let Some(stream) = promoted_stream.as_ref() { Self::promote_session_for_stream(connection_data, stream); } - promotions.push(action); + promotions.push((username.clone(), action)); } for session_token in &removed_session_tokens { Self::clear_session_counted_without_stream(connection_data, session_token); @@ -1033,12 +1051,19 @@ impl ActiveUserManager { connection_data.granted_grace = false; connection_data.grace_ts = 0; } + connection_count_changed |= connection_data.connections != previous_connection_count; } - let state_changed = had_registration || !removed_streams.is_empty(); - (state_changed, Some(username), removed_streams, expiry_entries, preserved_updates, promotions) - } else { - (false, None, Vec::new(), Vec::new(), Vec::new(), Vec::new()) } + let state_changed = had_registration || !disconnected_users.is_empty(); + ( + state_changed, + connection_count_changed, + disconnected_users, + removed_streams, + expiry_entries, + preserved_updates, + promotions, + ) }; for entry in expiry_entries { @@ -1049,29 +1074,31 @@ impl ActiveUserManager { self.event_manager.send_event(EventMessage::ActiveUser(ActiveUserConnectionChange::Updated(stream_info))); } - if let Some(ref username) = disconnected_user { + for username in &disconnected_users { if !username.is_empty() { debug_if_enabled!( "Released connection for user {username} at {}", sanitize_sensitive_info(&addr.to_string()) ); } - if addr_removed { - self.log_active_user().await; - for action in promotions { - self.emit_promotion_update(username, action).await; - } + } + if connection_count_changed { + self.log_active_user().await; + } + if addr_removed { + for (username, action) in promotions { + self.emit_promotion_update(&username, action).await; } } - ReleasedConnection { addr_removed, removed_streams, disconnected_user } + ReleasedConnection { addr_removed, removed_streams, disconnected_users } } pub async fn release_connection(&self, addr: &SocketAddr) -> ReleasedConnection { let released = self.release_connection_inner(addr, true).await; // divergence check after connection release if released.addr_removed { - if let Some(ref username) = released.disconnected_user { + for username in &released.disconnected_users { self.check_and_log_divergence_for_user(username).await; } } @@ -1082,7 +1109,7 @@ impl ActiveUserManager { let released = self.release_connection_inner(addr, false).await; // divergence check after connection release if released.addr_removed { - if let Some(ref username) = released.disconnected_user { + for username in &released.disconnected_users { self.check_and_log_divergence_for_user(username).await; } } @@ -1469,7 +1496,7 @@ impl ActiveUserManager { user_agent, session_token, } = update; - let (stream_info, divergence_snapshot) = { + let (stream_info, divergence_snapshot, connection_count_changed) = { let mut user_connections = self.connections.write().await; let now = current_time_secs(); @@ -1493,6 +1520,7 @@ impl ActiveUserManager { .or_insert_with(|| UserConnectionData::new(0, max_connections, soft_connections)); connection_data.max_connections = max_connections; connection_data.soft_connections = soft_connections; + let previous_connection_count = connection_data.connections; if let Some(token) = session_token { if let Some(session) = connection_data.sessions.iter_mut().find(|session| session.token == token) { @@ -1567,7 +1595,7 @@ impl ActiveUserManager { stream_info.previous_session_id = None; (result, was_preserved) }); - if let Some((stream_info, was_preserved)) = existing_stream_info { + let (stream_info, divergence_snapshot) = if let Some((stream_info, was_preserved)) = existing_stream_info { let effective_connection_kind = reserved_session_kind.unwrap_or(connection_kind); if was_preserved { connection_data.increment_kind(effective_connection_kind); @@ -1618,12 +1646,16 @@ impl ActiveUserManager { Self::log_connection_added(username, &fingerprint.addr, connection_data, tracked_socket_count); let divergence_snapshot = Self::collect_divergence_snapshot(connection_data, username); (stream_info, divergence_snapshot) - } + }; + let connection_count_changed = connection_data.connections != previous_connection_count; + (stream_info, divergence_snapshot, connection_count_changed) }; self.log_divergence_snapshot(divergence_snapshot).await; - self.log_active_user().await; + if connection_count_changed { + self.log_active_user().await; + } Some(stream_info) } @@ -3489,6 +3521,15 @@ mod tests { } } + fn test_series_channel(virtual_id: u32) -> StreamChannel { + StreamChannel { + item_type: PlaylistItemType::Series, + cluster: XtreamCluster::Series, + url: "http://localhost/series/episode.mkv".intern(), + ..test_channel(virtual_id) + } + } + /// Session refresh normalizes Expired -> Prepared. /// When a new request arrives on an expired session, the lifecycle should be /// reset to Prepared so that full activation evaluation happens. @@ -5717,6 +5758,189 @@ mod tests { assert_eq!(streams[0].uid, 41); } + #[tokio::test] + async fn release_stream_by_uid_finds_original_user_after_shared_addr_owner_changes() { + let config = Config::default(); + let geoip = Arc::new(ArcSwapOption::::default()); + let event_manager = Arc::new(EventManager::new()); + let manager = ActiveUserManager::new(&config, &geoip, &event_manager); + + let addr: SocketAddr = "127.0.0.1:55034".parse().unwrap(); + let fingerprint = Fingerprint::new("fp-cross-user-stream".to_string(), "127.0.0.1".to_string(), addr); + manager.add_connection(&addr).await; + + for (uid, username) in [(43, "user-a"), (44, "user-b")] { + manager + .update_connection(ActiveUserConnectionParams { + uid, + meter_uid: 0, + username, + max_connections: 1, + soft_connections: 0, + connection_kind: ConnectionKind::Normal, + priority: 0, + soft_priority: 0, + fingerprint: &fingerprint, + provider: "provider-a".intern(), + stream_channel: &test_series_channel(3005), + user_agent: Cow::Borrowed("ua"), + session_token: None, + }) + .await + .expect("direct Series stream should register"); + } + + assert_eq!(manager.active_users_and_connections().await, (2, 2)); + assert_eq!(manager.active_streams().await.len(), 2); + + let removed = manager.release_stream_by_uid(&addr, 43).await; + assert!(removed.as_ref().is_some_and(|stream| stream.uid == 43)); + assert_eq!(manager.active_users_and_connections().await, (1, 1)); + let streams = manager.active_streams().await; + assert_eq!(streams.len(), 1); + assert_eq!(streams[0].uid, 44); + + assert!(manager.release_stream_by_uid(&addr, 44).await.is_some()); + assert_eq!(manager.active_users_and_connections().await, (0, 0)); + assert!(manager.active_streams().await.is_empty()); + } + + #[tokio::test] + async fn release_connection_cleans_every_user_stream_for_reused_addr_only() { + let config = Config::default(); + let geoip = Arc::new(ArcSwapOption::::default()); + let event_manager = Arc::new(EventManager::new()); + let manager = ActiveUserManager::new(&config, &geoip, &event_manager); + + let reused_addr: SocketAddr = "127.0.0.1:55035".parse().unwrap(); + let unrelated_addr: SocketAddr = "127.0.0.1:55036".parse().unwrap(); + let reused_fingerprint = + Fingerprint::new("fp-cross-user-socket".to_string(), "127.0.0.1".to_string(), reused_addr); + let unrelated_fingerprint = + Fingerprint::new("fp-unrelated-socket".to_string(), "127.0.0.1".to_string(), unrelated_addr); + manager.add_connection(&reused_addr).await; + manager.add_connection(&unrelated_addr).await; + + for (uid, username, fingerprint) in [ + (45, "user-a", &reused_fingerprint), + (46, "user-b", &reused_fingerprint), + (47, "user-a", &unrelated_fingerprint), + ] { + manager + .update_connection(ActiveUserConnectionParams { + uid, + meter_uid: 0, + username, + max_connections: 2, + soft_connections: 0, + connection_kind: ConnectionKind::Normal, + priority: 0, + soft_priority: 0, + fingerprint, + provider: "provider-a".intern(), + stream_channel: &test_series_channel(3006 + uid), + user_agent: Cow::Borrowed("ua"), + session_token: None, + }) + .await + .expect("direct Series stream should register"); + } + + let released = manager.release_connection(&reused_addr).await; + let mut removed_uids = released.removed_streams.iter().map(|stream| stream.uid).collect::>(); + removed_uids.sort_unstable(); + assert!(released.addr_removed); + assert_eq!(removed_uids, vec![45, 46]); + assert_eq!(manager.active_users_and_connections().await, (1, 1)); + let streams = manager.active_streams().await; + assert_eq!(streams.len(), 1); + assert_eq!(streams[0].uid, 47); + + manager.release_connection(&unrelated_addr).await; + assert_eq!(manager.active_users_and_connections().await, (0, 0)); + assert!(manager.active_streams().await.is_empty()); + } + + #[tokio::test] + async fn connection_counts_are_broadcast_when_active_user_logging_is_disabled() { + let config = Config::default(); + let geoip = Arc::new(ArcSwapOption::::default()); + let event_manager = Arc::new(EventManager::new()); + let mut events = event_manager.get_event_channel(); + let manager = ActiveUserManager::new(&config, &geoip, &event_manager); + + let addr: SocketAddr = "127.0.0.1:55037".parse().unwrap(); + let fingerprint = Fingerprint::new("fp-count-events".to_string(), "127.0.0.1".to_string(), addr); + manager.add_connection(&addr).await; + manager.release_connection(&addr).await; + assert!( + tokio::time::timeout(Duration::from_millis(50), events.recv()).await.is_err(), + "closing an unowned socket must not broadcast unchanged connection counts" + ); + manager.add_connection(&addr).await; + manager + .update_connection(ActiveUserConnectionParams { + uid: 48, + meter_uid: 0, + username: "event-user", + max_connections: 1, + soft_connections: 0, + connection_kind: ConnectionKind::Normal, + priority: 0, + soft_priority: 0, + fingerprint: &fingerprint, + provider: "provider-a".intern(), + stream_channel: &test_series_channel(3048), + user_agent: Cow::Borrowed("ua"), + session_token: None, + }) + .await + .expect("direct Series stream should register"); + + assert_eq!( + tokio::time::timeout(Duration::from_secs(1), events.recv()) + .await + .expect("connection count update should be broadcast") + .expect("event channel should remain open"), + EventMessage::ActiveUser(ActiveUserConnectionChange::Connections(1, 1)) + ); + + manager + .update_connection(ActiveUserConnectionParams { + uid: 49, + meter_uid: 0, + username: "event-user", + max_connections: 1, + soft_connections: 0, + connection_kind: ConnectionKind::Normal, + priority: 0, + soft_priority: 0, + fingerprint: &fingerprint, + provider: "provider-a".intern(), + stream_channel: &test_series_channel(3048), + user_agent: Cow::Borrowed("ua"), + session_token: None, + }) + .await + .expect("same direct Series stream should be reused"); + assert!( + tokio::time::timeout(Duration::from_millis(50), events.recv()).await.is_err(), + "unchanged connection counts must not broadcast another full snapshot" + ); + + manager + .release_stream_by_uid(&addr, 48) + .await + .expect("direct Series stream should release"); + assert_eq!( + tokio::time::timeout(Duration::from_secs(1), events.recv()) + .await + .expect("released count update should be broadcast") + .expect("event channel should remain open"), + EventMessage::ActiveUser(ActiveUserConnectionChange::Connections(0, 0)) + ); + } + #[tokio::test] async fn test_adaptive_session_release_connection_preserves_logical_stream_and_start_time() { let config = Config::default(); diff --git a/backend/src/api/model/connection_manager.rs b/backend/src/api/model/connection_manager.rs index b5008ebe5..d6a6f678b 100644 --- a/backend/src/api/model/connection_manager.rs +++ b/backend/src/api/model/connection_manager.rs @@ -312,7 +312,7 @@ async fn release_connection_parts( } // Explicitly terminate all sessions for the kicked addr. This expires them // immediately rather than leaving them for TTL-based GC cleanup. - if let Some(ref username) = removed.disconnected_user { + for username in &removed.disconnected_users { deps.user_manager.terminate_sessions_for_addr(username, addr).await; } } @@ -953,7 +953,7 @@ impl ConnectionManager { self.provider_manager.clear_provider_reservation(session_token).await; } } - if let Some(ref username) = removed.disconnected_user { + for username in &removed.disconnected_users { self.user_manager.terminate_sessions_for_addr(username, addr).await; } for stream_info in &removed.removed_streams { diff --git a/backend/src/api/model/streams/active_client_stream.rs b/backend/src/api/model/streams/active_client_stream.rs index a47f73cf9..6f3eecf7f 100644 --- a/backend/src/api/model/streams/active_client_stream.rs +++ b/backend/src/api/model/streams/active_client_stream.rs @@ -1333,6 +1333,7 @@ fn stream_grace_period(request: GracePeriodParams) -> (Option>, Op #[cfg(test)] mod tests { + use super::super::buffered_stream::BufferedStream; use super::{ create_active_client_stream, create_deferred_provider_open_future, should_use_direct_body_idle_timeout, stream_grace_period, ActiveClientStream, ActiveClientStreamParams, ActiveClientStreamState, CustomVideoBuffers, @@ -1344,9 +1345,10 @@ mod tests { api_utils::GraceResolutionContext, model::{ connection_manager::PROVIDER_END_NOT_SET, ActiveProviderManager, ActiveUserManager, AppState, - CancelTokens, ConnectionManager, CreateUserSessionParams, CustomVideoStreamType, DownloadQueue, - EventManager, MetadataUpdateManager, PlaylistStorageState, ProviderContentRepresentationMode, - SharedStreamManager, StreamDetails, StreamError, UpdateGuard, + BoxedProviderStream, CancelTokens, ConnectionManager, CreateUserSessionParams, CustomVideoStreamType, + DownloadQueue, EventManager, MetadataUpdateManager, PlaylistStorageState, + ProviderContentRepresentationMode, ProviderHandle, SharedStreamManager, StreamDetails, StreamError, + UpdateGuard, }, }, auth::Fingerprint, @@ -1357,9 +1359,10 @@ mod tests { utils::{FileLockManager, GeoIp}, }; use arc_swap::{ArcSwap, ArcSwapOption}; - use axum::http::HeaderMap; + use axum::{body::Body, http::HeaderMap}; use bytes::Bytes; use futures::{pin_mut, StreamExt}; + use http_body_util::BodyExt; use reqwest::Client; use shared::{ model::{ @@ -1370,13 +1373,17 @@ mod tests { }; use std::{ collections::HashMap, + net::SocketAddr, + pin::Pin, sync::{ - atomic::{AtomicU8, Ordering}, + atomic::{AtomicBool, AtomicU8, Ordering}, Arc, }, + task::{Context, Poll}, time::Duration, }; - use tokio::sync::mpsc; + use tokio::sync::{mpsc, oneshot, Notify}; + use tokio_util::sync::CancellationToken; fn create_test_app_config() -> AppConfig { let input = Arc::new(ConfigInput { @@ -1602,6 +1609,180 @@ mod tests { channel } + fn create_test_series_stream_channel(virtual_id: u32, url: &str) -> StreamChannel { + let mut channel = create_test_stream_channel(virtual_id, url); + channel.item_type = PlaylistItemType::Series; + channel.cluster = XtreamCluster::Series; + channel.group = "Series".intern(); + channel.url = url.into(); + channel + } + + #[derive(Clone, Default)] + struct DropTracker(Arc); + + impl DropTracker { + fn is_dropped(&self) -> bool { self.0.load(Ordering::Acquire) } + } + + struct DropTrackedProviderStream { + inner: BoxedProviderStream, + tracker: DropTracker, + } + + impl futures::Stream for DropTrackedProviderStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.inner.as_mut().poll_next(cx) + } + } + + impl Drop for DropTrackedProviderStream { + fn drop(&mut self) { self.tracker.0.store(true, Ordering::Release); } + } + + fn track_provider_stream(stream: BoxedProviderStream) -> (BoxedProviderStream, DropTracker) { + let tracker = DropTracker::default(); + let tracked_stream = DropTrackedProviderStream { inner: stream, tracker: tracker.clone() }.boxed(); + (tracked_stream, tracker) + } + + struct TestDirectStreamParams<'a> { + app_state: &'a Arc, + username: &'a str, + max_connections: u32, + addr: SocketAddr, + stream_channel: StreamChannel, + provider_stream: BoxedProviderStream, + provider_handle: Option, + } + + struct TestDirectStream { + stream: BoxedProviderStream, + uid: u32, + } + + async fn create_test_active_direct_stream(params: TestDirectStreamParams<'_>) -> TestDirectStream { + let TestDirectStreamParams { + app_state, + username, + max_connections, + addr, + stream_channel, + provider_stream, + provider_handle, + } = params; + let mut user = create_test_user(username); + user.max_connections = max_connections; + let fingerprint = create_test_fingerprint(addr); + let mut stream_details = StreamDetails::from_stream(provider_stream, GracePeriodOptions::default()); + if provider_handle.is_some() { + stream_details.provider_name = Some("provider_1".intern()); + } + stream_details.provider_handle = provider_handle; + let virtual_id = stream_channel.virtual_id; + + let stream = create_active_client_stream(ActiveClientStreamParams { + stream_details, + app_state, + user: &user, + connection_permission: UserConnectionPermission::Allowed, + connection_kind: crate::api::model::ConnectionKind::Normal, + fingerprint: &fingerprint, + stream_channel, + socket_bound: false, + session_token: None, + req_headers: &HeaderMap::default(), + meter_uid: 0, + meter_stream: false, + }) + .await; + let uid = app_state + .active_users + .active_streams() + .await + .into_iter() + .find(|active| { + active.username == username && active.addr == addr && active.channel.virtual_id == virtual_id + }) + .map(|active| active.uid) + .expect("direct test stream should be registered"); + + TestDirectStream { stream, uid } + } + + async fn acquire_test_provider_handle(app_state: &Arc, addr: SocketAddr) -> ProviderHandle { + app_state + .active_provider + .acquire_exact_connection_with_grace( + &"provider_1".intern(), + &addr, + false, + 0, + crate::api::model::ConnectionKind::Normal, + ) + .await + .expect("direct test stream should acquire the provider slot") + } + + #[derive(Debug, Eq, PartialEq)] + struct TestLifecycleSnapshot { + active_counts: (usize, usize), + stream_uids: Vec, + provider_connections: usize, + } + + async fn lifecycle_snapshot(app_state: &Arc) -> TestLifecycleSnapshot { + let active_counts = app_state.active_users.active_users_and_connections().await; + let mut stream_uids = + app_state.active_users.active_streams().await.into_iter().map(|stream| stream.uid).collect::>(); + stream_uids.sort_unstable(); + TestLifecycleSnapshot { + active_counts, + stream_uids, + provider_connections: app_state.active_provider.get_provider_connections_count().await, + } + } + + struct ExpectedLifecycle<'a> { + description: &'static str, + active_counts: (usize, usize), + stream_uids: &'a [u32], + provider_connections: usize, + dropped_streams: &'a [&'a DropTracker], + } + + async fn wait_for_lifecycle(app_state: &Arc, expected: ExpectedLifecycle<'_>) { + let mut expected_stream_uids = expected.stream_uids.to_vec(); + expected_stream_uids.sort_unstable(); + let expected_snapshot = TestLifecycleSnapshot { + active_counts: expected.active_counts, + stream_uids: expected_stream_uids, + provider_connections: expected.provider_connections, + }; + let completed = tokio::time::timeout(Duration::from_secs(1), async { + loop { + let snapshot = lifecycle_snapshot(app_state).await; + if snapshot == expected_snapshot && expected.dropped_streams.iter().all(|tracker| tracker.is_dropped()) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await; + + if completed.is_err() { + let snapshot = lifecycle_snapshot(app_state).await; + let dropped = expected.dropped_streams.iter().map(|tracker| tracker.is_dropped()).collect::>(); + panic!( + "{} did not converge: expected={expected_snapshot:?}, actual={snapshot:?}, dropped={dropped:?}", + expected.description + ); + } + } + fn create_deferred_provider_grace_details( provider_name: &Arc, provider_handle: crate::api::model::ProviderHandle, @@ -2643,36 +2824,302 @@ mod tests { assert_eq!(entries[0].total_kb, 3); } + #[tokio::test] + async fn direct_series_normal_eof_releases_full_lifecycle() { + let app_state = create_test_app_state(); + let addr = "127.0.0.1:55031".parse().unwrap_or_else(|_| unreachable!()); + let provider_handle = acquire_test_provider_handle(&app_state, addr).await; + let (provider_stream, tracker) = + track_provider_stream(futures::stream::once(async { Ok(Bytes::from_static(b"series-eof")) }).boxed()); + let direct = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "series-eof-user", + max_connections: 1, + addr, + stream_channel: create_test_series_stream_channel(1, "http://provider-1.example/series/1.mkv"), + provider_stream, + provider_handle: Some(provider_handle), + }) + .await; + let active_uids = [direct.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "registered Series EOF stream", + active_counts: (1, 1), + stream_uids: &active_uids, + provider_connections: 1, + dropped_streams: &[], + }, + ) + .await; + assert!(!tracker.is_dropped()); + + let body = Body::from_stream(direct.stream); + let collected = body.collect().await.expect("normal provider EOF should complete the response body"); + assert_eq!(collected.to_bytes(), Bytes::from_static(b"series-eof")); + + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "normal Series EOF cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&tracker], + }, + ) + .await; + } + + #[tokio::test] + async fn direct_series_upstream_error_after_registration_releases_full_lifecycle() { + let app_state = create_test_app_state(); + let addr = "127.0.0.1:55032".parse().unwrap_or_else(|_| unreachable!()); + let provider_handle = acquire_test_provider_handle(&app_state, addr).await; + let provider_items = vec![ + Ok(Bytes::from_static(b"before-error")), + Err(StreamError::Stream("controlled upstream failure".to_string())), + ]; + let (provider_stream, tracker) = track_provider_stream(futures::stream::iter(provider_items).boxed()); + let direct = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "series-error-user", + max_connections: 1, + addr, + stream_channel: create_test_series_stream_channel(2, "http://provider-1.example/series/2.mkv"), + provider_stream, + provider_handle: Some(provider_handle), + }) + .await; + let active_uids = [direct.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "registered Series error stream", + active_counts: (1, 1), + stream_uids: &active_uids, + provider_connections: 1, + dropped_streams: &[], + }, + ) + .await; + + let body = Body::from_stream(direct.stream); + let collected = body + .collect() + .await + .expect("ActiveClientStream should convert the controlled upstream error into stream termination"); + assert_eq!(collected.to_bytes(), Bytes::from_static(b"before-error")); + + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "Series upstream error cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&tracker], + }, + ) + .await; + } + + #[tokio::test] + async fn closed_buffered_consumer_releases_direct_series_full_lifecycle() { + let app_state = create_test_app_state(); + let addr = "127.0.0.1:55039".parse().unwrap_or_else(|_| unreachable!()); + let provider_handle = acquire_test_provider_handle(&app_state, addr).await; + let (gate_tx, gate_rx) = oneshot::channel(); + let gated_provider = futures::stream::once(async move { + gate_rx.await.expect("test should release the upstream chunk"); + Ok(Bytes::from_static(b"after-consumer-drop")) + }) + .chain(futures::stream::pending()) + .boxed(); + let (tracked_provider, tracker) = track_provider_stream(gated_provider); + let producer_cancel = CancellationToken::new(); + let buffered_provider = + BufferedStream::new(tracked_provider, 1, producer_cancel.clone(), "controlled-test-stream").boxed(); + let direct = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "series-closed-consumer-user", + max_connections: 1, + addr, + stream_channel: create_test_series_stream_channel(10, "http://provider-1.example/series/10.mkv"), + provider_stream: buffered_provider, + provider_handle: Some(provider_handle), + }) + .await; + let active_uids = [direct.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "registered buffered Series stream", + active_counts: (1, 1), + stream_uids: &active_uids, + provider_connections: 1, + dropped_streams: &[], + }, + ) + .await; + + drop(Body::from_stream(direct.stream)); + gate_tx.send(()).expect("buffered producer should still own the gated upstream"); + + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "closed buffered Series consumer cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&tracker], + }, + ) + .await; + assert!(producer_cancel.is_cancelled(), "closed consumer must cancel the buffered producer"); + } + + #[tokio::test] + async fn aborting_task_polling_direct_series_body_releases_full_lifecycle() { + let app_state = create_test_app_state(); + let addr = "127.0.0.1:55033".parse().unwrap_or_else(|_| unreachable!()); + let provider_handle = acquire_test_provider_handle(&app_state, addr).await; + let (provider_stream, tracker) = + track_provider_stream(futures::stream::pending::>().boxed()); + let direct = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "series-abort-user", + max_connections: 1, + addr, + stream_channel: create_test_series_stream_channel(3, "http://provider-1.example/series/3.mkv"), + provider_stream, + provider_handle: Some(provider_handle), + }) + .await; + let active_uids = [direct.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "registered Series task-abort stream", + active_counts: (1, 1), + stream_uids: &active_uids, + provider_connections: 1, + dropped_streams: &[], + }, + ) + .await; + + let started = Arc::new(Notify::new()); + let started_in_task = Arc::clone(&started); + let task = tokio::spawn(async move { + let mut body = Body::from_stream(direct.stream); + started_in_task.notify_one(); + let _ = body.frame().await; + }); + tokio::time::timeout(Duration::from_secs(1), started.notified()).await.expect("body polling task should start"); + task.abort(); + let join_error = tokio::time::timeout(Duration::from_secs(1), task) + .await + .expect("aborted body polling task should finish") + .expect_err("pending body polling task should be cancelled"); + assert!(join_error.is_cancelled()); + + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "aborted Series body task cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&tracker], + }, + ) + .await; + } + + #[tokio::test] + async fn response_builder_error_after_registration_drops_direct_series_body() { + let app_state = create_test_app_state(); + let addr = "127.0.0.1:55034".parse().unwrap_or_else(|_| unreachable!()); + let provider_handle = acquire_test_provider_handle(&app_state, addr).await; + let (provider_stream, tracker) = + track_provider_stream(futures::stream::pending::>().boxed()); + let direct = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "series-builder-error-user", + max_connections: 1, + addr, + stream_channel: create_test_series_stream_channel(4, "http://provider-1.example/series/4.mkv"), + provider_stream, + provider_handle: Some(provider_handle), + }) + .await; + let active_uids = [direct.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "registered Series response-builder stream", + active_counts: (1, 1), + stream_uids: &active_uids, + provider_connections: 1, + dropped_streams: &[], + }, + ) + .await; + + let response = axum::response::Response::builder().status(10_000_u16).body(Body::from_stream(direct.stream)); + assert!(response.is_err(), "invalid status should fail response construction"); + + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "response-builder error cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&tracker], + }, + ) + .await; + } + #[tokio::test(start_paused = true)] async fn test_direct_vod_body_idle_timeout_releases_active_stream() { let app_state = create_test_app_state(); - let addr = "127.0.0.1:55031".parse().unwrap_or_else(|_| unreachable!()); - let test_user = create_test_user("vod-user"); - let test_fingerprint = create_test_fingerprint(addr); - let provider_stream = + let addr = "127.0.0.1:55035".parse().unwrap_or_else(|_| unreachable!()); + let provider_handle = acquire_test_provider_handle(&app_state, addr).await; + let pending_provider = futures::stream::once(async { Ok(Bytes::from_static(b"vod")) }).chain(futures::stream::pending()).boxed(); - let stream_details = StreamDetails::from_stream(provider_stream, GracePeriodOptions::default()); - - let stream = create_active_client_stream(ActiveClientStreamParams { - stream_details, + let (provider_stream, tracker) = track_provider_stream(pending_provider); + let direct = create_test_active_direct_stream(TestDirectStreamParams { app_state: &app_state, - user: &test_user, - connection_permission: UserConnectionPermission::Allowed, - connection_kind: crate::api::model::ConnectionKind::Normal, - fingerprint: &test_fingerprint, - stream_channel: create_test_video_stream_channel(1, "http://provider-1.example/movie/1.mkv"), - socket_bound: false, - session_token: None, - req_headers: &HeaderMap::default(), - meter_uid: 0, - meter_stream: false, + username: "vod-user", + max_connections: 1, + addr, + stream_channel: create_test_video_stream_channel(5, "http://provider-1.example/movie/5.mkv"), + provider_stream, + provider_handle: Some(provider_handle), }) .await; + let active_uids = [direct.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "registered idle-timeout VOD stream", + active_counts: (1, 1), + stream_uids: &active_uids, + provider_connections: 1, + dropped_streams: &[], + }, + ) + .await; + let stream = direct.stream; pin_mut!(stream); let first_chunk = stream.next().await; assert!(matches!(first_chunk, Some(Ok(ref bytes)) if bytes.as_ref() == b"vod")); - assert_eq!(app_state.active_users.active_streams().await.len(), 1); assert!( matches!(futures::poll!(stream.next()), std::task::Poll::Pending), "pending VOD body should wait until the direct body idle timeout elapses" @@ -2680,18 +3127,165 @@ mod tests { tokio::time::advance(Duration::from_secs(DIRECT_BODY_IDLE_TIMEOUT_SECS)).await; tokio::task::yield_now().await; - assert!(stream.next().await.is_none(), "VOD body idle timeout should terminate the stream"); + tokio::time::resume(); - for _ in 0..20 { - if app_state.active_users.active_streams().await.is_empty() { - return; - } - tokio::task::yield_now().await; - } - assert!( - app_state.active_users.active_streams().await.is_empty(), - "cleanup worker should remove the timed-out VOD stream from active streams" - ); + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "VOD body idle-timeout cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&tracker], + }, + ) + .await; + } + + #[tokio::test] + async fn dropping_direct_series_body_releases_original_user_after_socket_owner_changes() { + let app_state = create_test_app_state(); + let addr = "127.0.0.1:55036".parse().unwrap_or_else(|_| unreachable!()); + let provider_handle = acquire_test_provider_handle(&app_state, addr).await; + let (first_provider_stream, first_tracker) = + track_provider_stream(futures::stream::pending::>().boxed()); + let first = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "series-user-a", + max_connections: 1, + addr, + stream_channel: create_test_series_stream_channel(6, "http://provider-1.example/series/6.mkv"), + provider_stream: first_provider_stream, + provider_handle: Some(provider_handle), + }) + .await; + let (second_provider_stream, second_tracker) = + track_provider_stream(futures::stream::pending::>().boxed()); + let second = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "series-user-b", + max_connections: 1, + addr, + stream_channel: create_test_series_stream_channel(7, "http://provider-1.example/series/7.mkv"), + provider_stream: second_provider_stream, + provider_handle: None, + }) + .await; + let both_uids = [first.uid, second.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "two users sharing a reused socket address", + active_counts: (2, 2), + stream_uids: &both_uids, + provider_connections: 1, + dropped_streams: &[], + }, + ) + .await; + + let first_body = Body::from_stream(first.stream); + let second_body = Body::from_stream(second.stream); + drop(first_body); + let second_uid = [second.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "original user Body drop after socket owner replacement", + active_counts: (1, 1), + stream_uids: &second_uid, + provider_connections: 0, + dropped_streams: &[&first_tracker], + }, + ) + .await; + assert!(!second_tracker.is_dropped()); + + drop(second_body); + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "replacement user Body drop cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&first_tracker, &second_tracker], + }, + ) + .await; + } + + #[tokio::test] + async fn multiple_same_user_direct_series_connections_release_independently() { + let app_state = create_test_app_state(); + let first_addr = "127.0.0.1:55037".parse().unwrap_or_else(|_| unreachable!()); + let second_addr = "127.0.0.1:55038".parse().unwrap_or_else(|_| unreachable!()); + let (first_provider_stream, first_tracker) = + track_provider_stream(futures::stream::pending::>().boxed()); + let first = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "multi-series-user", + max_connections: 2, + addr: first_addr, + stream_channel: create_test_series_stream_channel(8, "http://provider-1.example/series/8.mkv"), + provider_stream: first_provider_stream, + provider_handle: None, + }) + .await; + let (second_provider_stream, second_tracker) = + track_provider_stream(futures::stream::pending::>().boxed()); + let second = create_test_active_direct_stream(TestDirectStreamParams { + app_state: &app_state, + username: "multi-series-user", + max_connections: 2, + addr: second_addr, + stream_channel: create_test_series_stream_channel(9, "http://provider-1.example/series/9.mkv"), + provider_stream: second_provider_stream, + provider_handle: None, + }) + .await; + let both_uids = [first.uid, second.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "two connections for one Series user", + active_counts: (1, 2), + stream_uids: &both_uids, + provider_connections: 0, + dropped_streams: &[], + }, + ) + .await; + + let first_body = Body::from_stream(first.stream); + let second_body = Body::from_stream(second.stream); + drop(first_body); + let second_uid = [second.uid]; + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "first same-user Series connection cleanup", + active_counts: (1, 1), + stream_uids: &second_uid, + provider_connections: 0, + dropped_streams: &[&first_tracker], + }, + ) + .await; + assert!(!second_tracker.is_dropped()); + + drop(second_body); + wait_for_lifecycle( + &app_state, + ExpectedLifecycle { + description: "last same-user Series connection cleanup", + active_counts: (0, 0), + stream_uids: &[], + provider_connections: 0, + dropped_streams: &[&first_tracker, &second_tracker], + }, + ) + .await; } } diff --git a/backend/src/api/model/streams/buffered_stream.rs b/backend/src/api/model/streams/buffered_stream.rs index 5715083cf..28809197b 100644 --- a/backend/src/api/model/streams/buffered_stream.rs +++ b/backend/src/api/model/streams/buffered_stream.rs @@ -170,3 +170,68 @@ impl Stream for BufferedStream { } } } + +#[cfg(test)] +mod tests { + use super::BufferedStream; + use crate::api::model::StreamError; + use bytes::Bytes; + use futures::Stream; + use std::{future::Future, pin::Pin, task::{Context, Poll}, time::Duration}; + use tokio::sync::oneshot; + use tokio_util::sync::CancellationToken; + + struct GatedDropProbeStream { + gate: oneshot::Receiver<()>, + dropped: Option>, + yielded: bool, + } + + impl Stream for GatedDropProbeStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if self.yielded { + return Poll::Pending; + } + match Pin::new(&mut self.gate).poll(cx) { + Poll::Ready(Ok(())) => { + self.yielded = true; + Poll::Ready(Some(Ok(Bytes::from_static(b"chunk")))) + } + Poll::Ready(Err(_)) => Poll::Ready(None), + Poll::Pending => Poll::Pending, + } + } + } + + impl Drop for GatedDropProbeStream { + fn drop(&mut self) { + if let Some(dropped) = self.dropped.take() { + let _ = dropped.send(()); + } + } + } + + #[tokio::test] + async fn closed_consumer_cancels_and_terminates_buffered_producer() { + let (gate_tx, gate_rx) = oneshot::channel(); + let (dropped_tx, dropped_rx) = oneshot::channel(); + let cancel = CancellationToken::new(); + let upstream = GatedDropProbeStream { + gate: gate_rx, + dropped: Some(dropped_tx), + yielded: false, + }; + let buffered = BufferedStream::new(Box::pin(upstream), 1, cancel.clone(), "test"); + + drop(buffered); + gate_tx.send(()).expect("producer should still own the gated upstream"); + + tokio::time::timeout(Duration::from_secs(1), dropped_rx) + .await + .expect("producer should stop after sending to the closed consumer") + .expect("upstream Drop probe should be delivered"); + assert!(cancel.is_cancelled(), "closed receiver must cancel the buffered producer token"); + } +} diff --git a/backend/src/api/serve.rs b/backend/src/api/serve.rs index 173f6fe52..24904a4ce 100644 --- a/backend/src/api/serve.rs +++ b/backend/src/api/serve.rs @@ -182,3 +182,203 @@ async fn handle_connection( } }); } + +#[cfg(test)] +mod tests { + use super::serve; + use crate::{ + api::model::{ + create_active_client_stream, create_test_app_state, ActiveClientStreamParams, AppState, ConnectionKind, + StreamDetails, StreamError, + }, + auth::Fingerprint, + model::{Config, GracePeriodOptions, ProxyUserCredentials}, + }; + use axum::{ + body::Body, + extract::{ConnectInfo, State}, + http::HeaderMap, + response::Response, + routing::get, + Router, + }; + use bytes::Bytes; + use futures::Stream; + use shared::{ + model::{PlaylistItemType, StreamChannel, UserConnectionPermission, XtreamCluster}, + utils::Internable, + }; + use socket2::SockRef; + use std::{ + pin::Pin, + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }, + task::{Context, Poll}, + time::Duration, + }; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio_util::sync::CancellationToken; + + #[derive(Clone)] + struct DisconnectTestState { + app_state: Arc, + upstream_dropped: Arc, + } + + struct PendingDropProbeStream { + dropped: Arc, + } + + impl Stream for PendingDropProbeStream { + type Item = Result; + + fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { Poll::Pending } + } + + impl Drop for PendingDropProbeStream { + fn drop(&mut self) { self.dropped.store(true, Ordering::Release); } + } + + async fn pending_direct_series_response( + State(state): State, + ConnectInfo(addr): ConnectInfo, + ) -> Response { + let mut user = ProxyUserCredentials::default(); + user.username = "socket-series-user".to_string(); + user.max_connections = 1; + let fingerprint = Fingerprint::new(format!("socket-{addr}"), addr.ip().to_string(), addr); + let stream_channel = StreamChannel { + target_id: 1, + virtual_id: 1, + provider_id: 1, + input_name: "input".intern(), + item_type: PlaylistItemType::Series, + cluster: XtreamCluster::Series, + group: "Series".intern(), + title: "Episode".intern(), + url: "http://provider.example/series/1.mkv".intern(), + shared: false, + shared_joined_existing: None, + shared_stream_id: None, + technical: None, + epg_channel_id: None, + epg_reference_ts: None, + }; + let upstream = PendingDropProbeStream { + dropped: Arc::clone(&state.upstream_dropped), + }; + let stream_details = + StreamDetails::from_stream(Box::pin(upstream), GracePeriodOptions::default()); + let stream = create_active_client_stream(ActiveClientStreamParams { + stream_details, + app_state: &state.app_state, + user: &user, + connection_permission: UserConnectionPermission::Allowed, + connection_kind: ConnectionKind::Normal, + fingerprint: &fingerprint, + stream_channel, + socket_bound: false, + session_token: None, + req_headers: &HeaderMap::new(), + meter_uid: 0, + meter_stream: false, + }) + .await; + Response::new(Body::from_stream(stream)) + } + + #[derive(Clone, Copy)] + enum ClientDisconnect { + Fin, + Reset, + } + + async fn assert_socket_disconnect_cleans_direct_series(disconnect: ClientDisconnect) { + let app_state = create_test_app_state(Config::default()); + let upstream_dropped = Arc::new(AtomicBool::new(false)); + let router = Router::new() + .route("/series", get(pending_direct_series_response)) + .with_state(DisconnectTestState { + app_state: Arc::clone(&app_state), + upstream_dropped: Arc::clone(&upstream_dropped), + }); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("test listener"); + let server_addr = listener.local_addr().expect("listener address"); + let server_cancel = CancellationToken::new(); + let server_cancel_task = server_cancel.clone(); + let connection_manager = Arc::clone(&app_state.connection_manager); + let server = tokio::spawn(async move { + serve(listener, router, Some(server_cancel_task), &connection_manager).await; + }); + + let mut client = tokio::net::TcpStream::connect(server_addr).await.expect("test client connection"); + client + .write_all(b"GET /series HTTP/1.1\r\nHost: localhost\r\nConnection: keep-alive\r\n\r\n") + .await + .expect("test request"); + let mut response_head = [0_u8; 1024]; + let read = tokio::time::timeout(Duration::from_secs(1), client.read(&mut response_head)) + .await + .expect("response head timeout") + .expect("response head read"); + assert!(read > 0, "streaming response should begin before client disconnects"); + assert_eq!(app_state.active_users.active_users_and_connections().await, (1, 1)); + assert_eq!(app_state.active_users.active_streams().await.len(), 1); + + match disconnect { + ClientDisconnect::Fin => { + client.shutdown().await.expect("client FIN"); + drop(client); + } + ClientDisconnect::Reset => { + let client = client.into_std().expect("convert test client to std socket"); + SockRef::from(&client) + .set_linger(Some(Duration::ZERO)) + .expect("configure reset-on-close"); + drop(client); + } + } + + tokio::time::timeout(Duration::from_secs(1), async { + loop { + if app_state.active_users.active_users_and_connections().await == (0, 0) + && app_state.active_users.active_streams().await.is_empty() + && upstream_dropped.load(Ordering::Acquire) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("socket disconnect should release registry state and drop the upstream"); + assert_eq!( + app_state + .active_provider + .active_connections() + .await + .unwrap_or_default() + .values() + .sum::(), + 0 + ); + + server_cancel.cancel(); + tokio::time::timeout(Duration::from_secs(1), server) + .await + .expect("server shutdown timeout") + .expect("server task"); + } + + #[tokio::test] + async fn client_fin_cleans_direct_series_response() { + assert_socket_disconnect_cleans_direct_series(ClientDisconnect::Fin).await; + } + + #[tokio::test] + async fn client_reset_cleans_direct_series_response() { + assert_socket_disconnect_cleans_direct_series(ClientDisconnect::Reset).await; + } +} diff --git a/frontend/src/hooks/use_server_status.rs b/frontend/src/hooks/use_server_status.rs index 3f87dae4a..a7bad0594 100644 --- a/frontend/src/hooks/use_server_status.rs +++ b/frontend/src/hooks/use_server_status.rs @@ -132,6 +132,18 @@ fn merge_aux_streams(server_status: &mut StatusCheck, download_streams: &[Stream dedupe_streams_by_identity(&mut server_status.active_user_streams); } +fn replace_server_status_snapshot( + status_holder: &mut Option>, + mut server_status: StatusCheck, + download_streams: &[StreamInfo], +) -> Rc { + dedupe_streams_by_identity(&mut server_status.active_user_streams); + merge_aux_streams(&mut server_status, download_streams); + let server_status = Rc::new(server_status); + *status_holder = Some(Rc::clone(&server_status)); + server_status +} + fn rebuild_status_with_downloads( status_holder: &UseStateHandle>>>, status_signal: &UseStateHandle>>, @@ -265,11 +277,11 @@ pub fn use_server_status( subid = Some(services_ctx.event.subscribe(move |msg| match msg { EventMessage::ServerStatus(server_status) => { - let mut server_status = (*server_status).clone(); - dedupe_streams_by_identity(&mut server_status.active_user_streams); - merge_aux_streams(&mut server_status, download_streams_holder_signal.borrow().as_slice()); - let server_status = Rc::new(server_status); - *status_holder_signal.borrow_mut() = Some(Rc::clone(&server_status)); + let server_status = replace_server_status_snapshot( + &mut status_holder_signal.borrow_mut(), + (*server_status).clone(), + download_streams_holder_signal.borrow().as_slice(), + ); status_signal.set(Some(server_status)); } EventMessage::ActiveUser(event) => { @@ -347,7 +359,7 @@ pub fn use_server_status( mod tests { use super::{ apply_active_user_change, apply_downloads_delta, apply_downloads_snapshot, dedupe_streams_by_identity, - download_task_to_stream_with_ts, find_stream_update_index, + download_task_to_stream_with_ts, find_stream_update_index, replace_server_status_snapshot, }; use shared::{ model::{ @@ -356,7 +368,7 @@ mod tests { }, utils::Internable, }; - use std::net::SocketAddr; + use std::{net::SocketAddr, rc::Rc}; fn test_stream(uid: u32, addr: &str, session_token: Option<&str>, item_type: PlaylistItemType) -> StreamInfo { StreamInfo { @@ -569,6 +581,33 @@ mod tests { assert_eq!(status.active_user_streams, vec![kept]); } + #[test] + fn server_status_snapshot_replaces_stale_backend_streams_and_readds_current_downloads() { + let stale_stream = test_stream(1, "127.0.0.1:1234", Some("tok-series"), PlaylistItemType::Series); + let download_stream = download_task_to_stream_with_ts( + &test_download("running", TransferStatusDto::Running, TaskKindDto::Download), + 123, + ); + let mut status_holder = Some(Rc::new(shared::model::StatusCheck { + active_users: 1, + active_user_connections: 1, + active_user_streams: vec![stale_stream], + ..Default::default() + })); + let clean_backend_snapshot = shared::model::StatusCheck::default(); + + let status = replace_server_status_snapshot( + &mut status_holder, + clean_backend_snapshot, + std::slice::from_ref(&download_stream), + ); + + assert_eq!(status.active_users, 0); + assert_eq!(status.active_user_connections, 0); + assert_eq!(status.active_user_streams, vec![download_stream]); + assert_eq!(status_holder.as_deref(), Some(status.as_ref())); + } + fn test_download(id: &str, status: TransferStatusDto, kind: TaskKindDto) -> FileDownloadDto { FileDownloadDto { id: id.to_string(),