diff --git a/src/api/api_utils.rs b/src/api/api_utils.rs index b53f176eb..1bc8a3214 100644 --- a/src/api/api_utils.rs +++ b/src/api/api_utils.rs @@ -658,9 +658,10 @@ pub async fn stream_response(app_state: &AppState, if let Some(provider) = provider_name { if matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash | PlaylistItemType::Video | PlaylistItemType::Series) { - let grace_token = app_state.active_users.create_token(&user.username).await; - if let Some(cookie_value) = create_session_cookie_for_provider(&app_state.config.t_encrypt_secret, &grace_token, virtual_id, &provider, stream_url) { - response = response.header(axum::http::header::SET_COOKIE, &cookie_value); + if let Some(grace_token) = app_state.active_users.get_or_create_token(&user.username).await { + if let Some(cookie_value) = create_session_cookie_for_provider(&app_state.config.t_encrypt_secret, &grace_token, virtual_id, &provider, stream_url) { + response = response.header(axum::http::header::SET_COOKIE, &cookie_value); + } } } } @@ -859,8 +860,7 @@ pub async fn is_seek_response( } pub async fn check_force_provider(app_state: &AppState, virtual_id: u32, item_type: PlaylistItemType, req_headers: &HeaderMap, user: &ProxyUserCredentials) -> (Option, UserConnectionPermission) { - - if ! matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash | PlaylistItemType::Series | PlaylistItemType::Video) { + if !matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash | PlaylistItemType::Series | PlaylistItemType::Video) { return (None, user.connection_permission(app_state).await); } @@ -873,20 +873,17 @@ pub async fn check_force_provider(app_state: &AppState, virtual_id: u32, item_ty } } - let connection_permission = match provider_name { - Some(_) => UserConnectionPermission::Allowed, - None => { - let permission = user.connection_permission(app_state).await; - match permission { - UserConnectionPermission::GracePeriod => { - if app_state.active_users.get_token(&user.username).await.is_some() { - UserConnectionPermission::Exhausted - } else { - UserConnectionPermission::GracePeriod - } + let connection_permission = if provider_name.is_some() { UserConnectionPermission::Allowed } else { + let permission = user.connection_permission(app_state).await; + match permission { + UserConnectionPermission::GracePeriod => { + if app_state.active_users.get_token(&user.username).await.is_some() { + UserConnectionPermission::Exhausted + } else { + UserConnectionPermission::GracePeriod } - _ => permission, } + _ => permission, } }; diff --git a/src/api/endpoints/hls_api.rs b/src/api/endpoints/hls_api.rs index bd3ed33ee..1f07acded 100644 --- a/src/api/endpoints/hls_api.rs +++ b/src/api/endpoints/hls_api.rs @@ -45,12 +45,12 @@ pub(in crate::api) async fn handle_hls_stream_request(app_state: &Arc, let url = replace_url_extension(hls_url, HLS_EXT); let server_info = app_state.config.get_user_server_info(user).await; - let grace_token = app_state.active_users.create_token(&user.username).await; + let grace_token = app_state.active_users.get_or_create_token(&user.username).await; let create_stream_and_cookie = |provider_cfg: &Arc| { let stream_url = get_stream_alternative_url(&url, input, provider_cfg); let cookie = create_session_cookie_for_provider( &app_state.config.t_encrypt_secret, - &grace_token, + &grace_token.clone().unwrap_or_default(), virtual_id, &provider_cfg.name, &stream_url, @@ -79,7 +79,7 @@ pub(in crate::api) async fn handle_hls_stream_request(app_state: &Arc, virtual_id, input_id: input.id, provider_name: provider.unwrap_or_default(), // this should not happen - user_token: grace_token + user_token: grace_token.unwrap_or_default().to_string(), }; let hls_content = rewrite_hls(user, &rewrite_hls_props); hls_response(hls_content, cookie).into_response() diff --git a/src/api/model/active_user_manager.rs b/src/api/model/active_user_manager.rs index 7373d89f4..4b6cb1ee1 100644 --- a/src/api/model/active_user_manager.rs +++ b/src/api/model/active_user_manager.rs @@ -138,6 +138,7 @@ impl ActiveUserManager { let mut lock = self.user.write().await; if let Some(connection_data) = lock.get_mut(username) { connection_data.connections += 1; + connection_data.max_connections = max_connections; } else { lock.insert(username.to_string(), UserConnectionData::new(max_connections)); } @@ -163,10 +164,8 @@ impl ActiveUserManager { if connection_data.connections == 0 { lock.remove(username); - } else { - if connection_data.connections < connection_data.max_connections { - connection_data.token = None; - } + } else if connection_data.connections < connection_data.max_connections { + connection_data.token = None; } } drop(lock); @@ -174,11 +173,17 @@ impl ActiveUserManager { self.log_active_user().await; } - pub async fn create_token(&self, username: &str) -> String { - let result = crate::utils::string_utils::generate_random_string(6); + pub async fn get_or_create_token(&self, username: &str) -> Option { + let token = crate::utils::string_utils::generate_random_string(6); + let mut result = None; let mut lock = self.user.write().await; if let Some(connection_data) = lock.get_mut(username) { - connection_data.token = Some(result.to_string()); + result = if connection_data.token.is_some() { + connection_data.token.clone() + } else { + connection_data.token = Some(token.to_string()); + Some(token) + }; } drop(lock); result @@ -226,6 +231,8 @@ impl ActiveUserManager { // }) // .collect(); // + + // for handle in handles { // handle.join().unwrap(); // } diff --git a/src/api/model/provider_config.rs b/src/api/model/provider_config.rs index ba519923e..33700caf8 100644 --- a/src/api/model/provider_config.rs +++ b/src/api/model/provider_config.rs @@ -145,6 +145,12 @@ impl ProviderConfig { let mut guard = self.connection.write().await; let connections = guard.current_connections; if connections < self.max_connections || (grace && connections <= self.max_connections) { + if connections < self.max_connections { + guard.granted_grace = false; + guard.grace_ts = 0; + guard.current_connections += 1; + } + let now = get_current_timestamp(); if guard.granted_grace { if now - guard.grace_ts <= grace_period_timeout_secs {