From 0cbd5aa5b5fddba126e72b89f9d91bca18656067 Mon Sep 17 00:00:00 2001 From: euzu Date: Sun, 27 Apr 2025 13:26:21 +0200 Subject: [PATCH] provider connection limit is ignored --- src/api/api_utils.rs | 15 +- src/api/endpoints/hls_api.rs | 4 +- src/api/endpoints/m3u_api.rs | 2 +- src/api/endpoints/xtream_api.rs | 2 +- src/api/model/active_provider_manager.rs | 257 +++++------------------ src/api/model/active_user_manager.rs | 72 +++++-- src/api/model/mod.rs | 3 +- src/api/model/provider_config.rs | 195 +++++++++++++++++ 8 files changed, 320 insertions(+), 230 deletions(-) create mode 100644 src/api/model/provider_config.rs diff --git a/src/api/api_utils.rs b/src/api/api_utils.rs index cdad615ab..fe06631ef 100644 --- a/src/api/api_utils.rs +++ b/src/api/api_utils.rs @@ -1,5 +1,5 @@ use crate::api::endpoints::xtream_api::{get_xtream_player_api_stream_url, XtreamApiStreamContext}; -use crate::api::model::active_provider_manager::{ProviderAllocation, ProviderConfig, ProviderConnectionGuard}; +use crate::api::model::active_provider_manager::{ProviderAllocation, ProviderConnectionGuard}; use crate::api::model::app_state::AppState; use crate::api::model::model_utils::get_stream_response_with_headers; use crate::api::model::request::UserApiRequest; @@ -83,6 +83,7 @@ macro_rules! try_result_bad_request { pub use try_option_bad_request; pub use try_result_bad_request; +use crate::api::model::provider_config::ProviderConfig; pub fn get_server_time() -> String { chrono::offset::Local::now().with_timezone(&chrono::Local).format("%Y-%m-%d %H:%M:%S %Z").to_string() @@ -652,8 +653,10 @@ pub async fn stream_response(app_state: &AppState, } if let Some(provider) = provider_name { - if let Some(cookie_value) = create_session_cookie_for_provider(&app_state.config.t_encrypt_secret, virtual_id, &provider, stream_url) { - response = response.header(axum::http::header::SET_COOKIE, &cookie_value); + if matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash) { + if let Some(cookie_value) = create_session_cookie_for_provider(&app_state.config.t_encrypt_secret, virtual_id, &provider, stream_url) { + response = response.header(axum::http::header::SET_COOKIE, &cookie_value); + } } } @@ -846,7 +849,11 @@ pub fn is_seek_response( read_session_cookie(req_headers) } -pub async fn check_force_provider(app_state: &AppState, virtual_id: u32, req_headers: &HeaderMap, user: &ProxyUserCredentials) -> (Option, UserConnectionPermission) { +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) { + return (None, user.connection_permission(app_state).await); + } // if you have multi provider setup you need to delegate the same hls requests // to the same provider. Hls has alternating m3u8 and stream requests. diff --git a/src/api/endpoints/hls_api.rs b/src/api/endpoints/hls_api.rs index 5ec3eba67..df9b6fd76 100644 --- a/src/api/endpoints/hls_api.rs +++ b/src/api/endpoints/hls_api.rs @@ -1,6 +1,5 @@ use crate::api::api_utils::{bad_response_with_delete_cookie, check_force_provider, create_session_cookie_for_provider, force_provider_stream_response, get_stream_alternative_url}; use crate::api::api_utils::{get_stream_info_from_crypted_cookie, try_option_bad_request}; -use crate::api::model::active_provider_manager::ProviderConfig; use crate::api::model::app_state::AppState; use crate::api::model::streams::provider_stream::{create_custom_video_stream_response, CustomVideoStreamType}; use crate::model::api_proxy::{ProxyUserCredentials, UserConnectionPermission}; @@ -14,6 +13,7 @@ use axum::response::IntoResponse; use log::{debug, error}; use serde::Deserialize; use std::sync::Arc; +use crate::api::model::provider_config::ProviderConfig; #[derive(Debug, Deserialize)] struct HlsApiPathParams { @@ -104,7 +104,7 @@ async fn hls_api_stream( let virtual_id = params.stream_id; let input = try_option_bad_request!(app_state.config.get_input_by_id(params.input_id), true, format!("Cant find input for target {target_name}, context {}, stream_id {virtual_id}", XtreamCluster::Live)); - let (_provider_name, connection_permission) = check_force_provider(&app_state, virtual_id, &req_headers, &user).await; + let (_provider_name, connection_permission) = check_force_provider(&app_state, virtual_id, PlaylistItemType::LiveHls, &req_headers, &user).await; if connection_permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response(&app_state.config, &CustomVideoStreamType::UserConnectionsExhausted).into_response(); } diff --git a/src/api/endpoints/m3u_api.rs b/src/api/endpoints/m3u_api.rs index c5614603e..4ccef7a08 100644 --- a/src/api/endpoints/m3u_api.rs +++ b/src/api/endpoints/m3u_api.rs @@ -90,7 +90,7 @@ async fn m3u_api_stream( return force_provider_stream_response(&app_state, &cookie, pli.virtual_id, pli.item_type, &req_headers, input, &user).await.into_response() } - let (provider_name, connection_permission) = check_force_provider(&app_state, virtual_id, &req_headers, &user).await; + let (provider_name, connection_permission) = check_force_provider(&app_state, virtual_id, pli.item_type, &req_headers, &user).await; if connection_permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response(&app_state.config, &CustomVideoStreamType::UserConnectionsExhausted).into_response(); } diff --git a/src/api/endpoints/xtream_api.rs b/src/api/endpoints/xtream_api.rs index db3735359..c36c363fc 100644 --- a/src/api/endpoints/xtream_api.rs +++ b/src/api/endpoints/xtream_api.rs @@ -202,7 +202,7 @@ async fn xtream_player_api_stream( return force_provider_stream_response(app_state, &cookie, pli.virtual_id, pli.item_type, req_headers, input, &user).await.into_response() } - let (provider_name, connection_permission) = check_force_provider(app_state, virtual_id, req_headers, &user).await; + let (provider_name, connection_permission) = check_force_provider(app_state, virtual_id, pli.item_type, req_headers, &user).await; if connection_permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response(&app_state.config, &CustomVideoStreamType::UserConnectionsExhausted).into_response(); } diff --git a/src/api/model/active_provider_manager.rs b/src/api/model/active_provider_manager.rs index f77ecfee7..1ccf6e85b 100644 --- a/src/api/model/active_provider_manager.rs +++ b/src/api/model/active_provider_manager.rs @@ -1,10 +1,11 @@ -use crate::model::config::{Config, ConfigInput, ConfigInputAlias, InputType, InputUserInfo}; +use crate::model::config::{Config, ConfigInput}; use log::{debug, log_enabled}; use std::collections::HashMap; use std::ops::Deref; -use std::sync::atomic::{AtomicU16, AtomicUsize, Ordering}; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; -use tokio::sync::RwLock; +use tokio::sync::{Mutex, RwLock}; +use crate::api::model::provider_config::{ProviderConfig, ProviderConfigWrapper}; pub struct ProviderConnectionGuard { manager: Arc, @@ -62,156 +63,6 @@ pub enum ProviderAllocation { GracePeriod(Arc), } -/// This struct represents an individual provider configuration with fields like: -/// -/// `id`, `name`, `url`, `username`, `password` -/// `input_type`: Determines the type of input the provider supports. -/// `max_connections`: Maximum allowed concurrent connections. -/// `priority`: Priority level for selecting providers. -/// `current_connections`: A `RwLock` to safely track the number of active connections. -#[derive(Debug)] -pub struct ProviderConfig { - pub id: u16, - pub name: String, - pub url: String, - pub username: Option, - pub password: Option, - pub input_type: InputType, - max_connections: u16, - priority: i16, - current_connections: AtomicU16, -} - -impl ProviderConfig { - pub fn new(cfg: &ConfigInput) -> Self { - Self { - id: cfg.id, - name: cfg.name.clone(), - url: cfg.url.clone(), - username: cfg.username.clone(), - password: cfg.password.clone(), - input_type: cfg.input_type, - max_connections: cfg.max_connections, - priority: cfg.priority, - current_connections: AtomicU16::new(0), - } - } - - pub fn new_alias(cfg: &ConfigInput, alias: &ConfigInputAlias) -> Self { - Self { - id: alias.id, - name: alias.name.clone(), - url: alias.url.clone(), - username: alias.username.clone(), - password: alias.password.clone(), - input_type: cfg.input_type, - max_connections: alias.max_connections, - priority: alias.priority, - current_connections: AtomicU16::new(0), - } - } - - pub fn get_user_info(&self) -> Option { - InputUserInfo::new(self.input_type, self.username.as_deref(), self.password.as_deref(), &self.url) - } - - #[inline] - pub fn is_exhausted(&self) -> bool { - self.max_connections > 0 && self.current_connections.load(Ordering::SeqCst) >= self.max_connections - } - - #[inline] - pub fn is_over_limit(&self) -> bool { - self.max_connections > 0 && self.current_connections.load(Ordering::SeqCst) > self.max_connections - } - - // - // #[inline] - // pub fn has_capacity(&self) -> bool { - // !self.is_exhausted() - // } - - fn try_allocate(&self, grace: bool) -> u8 { - let connections = self.current_connections.load(Ordering::SeqCst); - if self.max_connections == 0 { - self.current_connections.fetch_add(1, Ordering::SeqCst); - return 1; - } - if (!grace && connections < self.max_connections) || (grace && connections <= self.max_connections) { - self.current_connections.fetch_add(1, Ordering::SeqCst); - return if connections < self.max_connections { 1 } else { 2 }; - } - 3 - } - - fn force_allocate(&self) { - self.current_connections.fetch_add(1, Ordering::SeqCst); - } - - // is intended to use with redirects, to cycle through provider - fn get_next(&self, grace: bool) -> bool { - let connections = self.current_connections.load(Ordering::SeqCst); - if self.max_connections == 0 { - return true; - } - if (!grace && connections < self.max_connections) || (grace && connections <= self.max_connections) { - return true; - } - false - } - - pub fn release(&self) { - let connections = self.current_connections.load(Ordering::SeqCst); - if connections > 0 { - self.current_connections.fetch_sub(1, Ordering::SeqCst); - } - } - - pub fn get_connection(&self) -> u16 { - self.current_connections.load(Ordering::SeqCst) - } -} - -#[derive(Clone, Debug)] -struct ProviderConfigWrapper { - inner: Arc, -} - - -impl ProviderConfigWrapper { - pub fn new(cfg: ProviderConfig) -> Self { - Self { - inner: Arc::new(cfg) - } - } - - pub fn force_allocate(&self) -> ProviderAllocation { - self.inner.force_allocate(); - ProviderAllocation::Available(Arc::clone(&self.inner)) - } - - pub fn try_allocate(&self, grace: bool) -> ProviderAllocation { - match self.inner.try_allocate(grace) { - 1 => ProviderAllocation::Available(Arc::clone(&self.inner)), - 2 => ProviderAllocation::GracePeriod(Arc::clone(&self.inner)), - _ => ProviderAllocation::Exhausted, - } - } - - pub fn get_next(&self, grace: bool) -> Option> { - if self.inner.get_next(grace) { - return Some(Arc::clone(&self.inner)); - } - None - } -} -impl Deref for ProviderConfigWrapper { - type Target = ProviderConfig; - - fn deref(&self) -> &Self::Target { - &self.inner - } -} /// This manages different types of provider lineups: /// @@ -224,24 +75,24 @@ enum ProviderLineup { } impl ProviderLineup { - fn get_next(&self) -> Option> { + async fn get_next(&self) -> Option> { match self { - ProviderLineup::Single(lineup) => lineup.get_next(), - ProviderLineup::Multi(lineup) => lineup.get_next(), + ProviderLineup::Single(lineup) => lineup.get_next().await, + ProviderLineup::Multi(lineup) => lineup.get_next().await, } } - fn acquire(&self) -> ProviderAllocation { + async fn acquire(&self) -> ProviderAllocation { match self { - ProviderLineup::Single(lineup) => lineup.acquire(), - ProviderLineup::Multi(lineup) => lineup.acquire(), + ProviderLineup::Single(lineup) => lineup.acquire().await, + ProviderLineup::Multi(lineup) => lineup.acquire().await, } } - fn release(&self, provider_name: &str) { + async fn release(&self, provider_name: &str) { match self { - ProviderLineup::Single(lineup) => lineup.release(provider_name), - ProviderLineup::Multi(lineup) => lineup.release(provider_name), + ProviderLineup::Single(lineup) => lineup.release(provider_name).await, + ProviderLineup::Multi(lineup) => lineup.release(provider_name).await, } } } @@ -259,17 +110,17 @@ impl SingleProviderLineup { } } - fn get_next(&self) -> Option> { - self.provider.get_next(false) + async fn get_next(&self) -> Option> { + self.provider.get_next(false).await } - fn acquire(&self) -> ProviderAllocation { - self.provider.try_allocate(true) + async fn acquire(&self) -> ProviderAllocation { + self.provider.try_allocate(true).await } - fn release(&self, provider_name: &str) { + async fn release(&self, provider_name: &str) { if self.provider.name == provider_name { - self.provider.release(); + self.provider.release().await; } } } @@ -282,7 +133,7 @@ impl SingleProviderLineup { #[derive(Debug)] enum ProviderPriorityGroup { SingleProviderGroup(ProviderConfigWrapper), - MultiProviderGroup(AtomicUsize, Vec), + MultiProviderGroup(Mutex, Vec), } impl ProviderPriorityGroup { @@ -320,7 +171,7 @@ impl MultiProviderLineup { } let mut providers = HashMap::new(); for provider in inputs { - let priority = provider.priority; + let priority = provider.get_priority(); providers.entry(priority) .or_insert_with(Vec::new) .push(provider); @@ -329,7 +180,7 @@ impl MultiProviderLineup { values.sort_by(|(p1, _), (p2, _)| p1.cmp(p2)); let providers: Vec = values.into_iter().map(|(_, mut group)| { if group.len() > 1 { - ProviderPriorityGroup::MultiProviderGroup(AtomicUsize::new(0), group) + ProviderPriorityGroup::MultiProviderGroup(Mutex::new(0), group) } else { ProviderPriorityGroup::SingleProviderGroup(group.remove(0)) } @@ -368,57 +219,59 @@ impl MultiProviderLineup { /// } /// } /// ``` - fn acquire_next_provider_from_group(priority_group: &ProviderPriorityGroup, grace: bool) -> ProviderAllocation { + async fn acquire_next_provider_from_group(priority_group: &ProviderPriorityGroup, grace: bool) -> ProviderAllocation { match priority_group { ProviderPriorityGroup::SingleProviderGroup(p) => { - let result = p.try_allocate(grace); + let result = p.try_allocate(grace).await; match result { ProviderAllocation::Exhausted => {} ProviderAllocation::Available(_) | ProviderAllocation::GracePeriod(_) => return result } } ProviderPriorityGroup::MultiProviderGroup(index, pg) => { - let mut idx = index.load(Ordering::SeqCst); + let mut idx_guard = index.lock().await; + let mut idx = *idx_guard; let provider_count = pg.len(); let start = idx; for _ in start..provider_count { let p = pg.get(idx).unwrap(); idx = (idx + 1) % provider_count; - let result = p.try_allocate(grace); + let result = p.try_allocate(grace).await; match result { ProviderAllocation::Exhausted => {} ProviderAllocation::Available(_) | ProviderAllocation::GracePeriod(_) => { - index.store(idx, Ordering::SeqCst); + *idx_guard = idx; return result; } } } - index.store(idx, Ordering::SeqCst); + *idx_guard = idx; } } ProviderAllocation::Exhausted } // Used for redirect to cylce through provider - fn get_next_provider_from_group(priority_group: &ProviderPriorityGroup, grace: bool) -> Option> { + async fn get_next_provider_from_group(priority_group: &ProviderPriorityGroup, grace: bool) -> Option> { match priority_group { ProviderPriorityGroup::SingleProviderGroup(p) => { - return p.get_next(grace); + return p.get_next(grace).await; } ProviderPriorityGroup::MultiProviderGroup(index, pg) => { - let mut idx = index.load(Ordering::SeqCst); + let mut idx_guard = index.lock().await; + let mut idx = *idx_guard; let provider_count = pg.len(); let start = idx; for _ in start..provider_count { let p = pg.get(idx).unwrap(); idx = (idx + 1) % provider_count; - let result = p.get_next(grace); + let result = p.get_next(grace).await; if result.is_some() { - index.store(idx, Ordering::SeqCst); + *idx_guard = idx; return result; } } - index.store(idx, Ordering::SeqCst); + *idx_guard = idx; } } None @@ -449,16 +302,16 @@ impl MultiProviderLineup { /// ProviderAllocation::GracePeriodprovider) => println!("Provider with grace period {}", provider.name), /// } /// ``` - fn acquire(&self) -> ProviderAllocation { + async fn acquire(&self) -> ProviderAllocation { let main_idx = self.index.load(Ordering::SeqCst); let provider_count = self.providers.len(); for index in main_idx..provider_count { let priority_group = &self.providers[index]; let allocation = { - let without_grace_allocation = Self::acquire_next_provider_from_group(priority_group, false); + let without_grace_allocation = Self::acquire_next_provider_from_group(priority_group, false).await; if matches!(without_grace_allocation, ProviderAllocation::Exhausted) { - Self::acquire_next_provider_from_group(priority_group, true) + Self::acquire_next_provider_from_group(priority_group, true).await } else { without_grace_allocation } @@ -479,16 +332,16 @@ impl MultiProviderLineup { } // it intended to use with redirects to cycle through provider - fn get_next(&self) -> Option> { + async fn get_next(&self) -> Option> { let main_idx = self.index.load(Ordering::SeqCst); let provider_count = self.providers.len(); for index in main_idx..provider_count { let priority_group = &self.providers[index]; let allocation = { - let config = Self::get_next_provider_from_group(priority_group, false); + let config = Self::get_next_provider_from_group(priority_group, false).await; if config.is_none() { - Self::get_next_provider_from_group(priority_group, true) + Self::get_next_provider_from_group(priority_group, true).await } else { config } @@ -508,19 +361,19 @@ impl MultiProviderLineup { } - fn release(&self, provider_name: &str) { + async fn release(&self, provider_name: &str) { for g in &self.providers { match g { ProviderPriorityGroup::SingleProviderGroup(pc) => { if pc.name == provider_name { - pc.release(); + pc.release().await; break; } } ProviderPriorityGroup::MultiProviderGroup(_, group) => { for pc in group { if pc.name == provider_name { - pc.release(); + pc.release().await; return; } } @@ -611,7 +464,7 @@ impl ActiveProviderManager { let providers = self.providers.read().await; let allocation = match Self::get_provider_config(input_name, &providers) { None => ProviderAllocation::Exhausted, // No Name matched, we don't have this provider - Some((lineup, _config)) => lineup.acquire() + Some((lineup, _config)) => lineup.acquire().await }; if log_enabled!(log::Level::Debug) { @@ -637,7 +490,7 @@ impl ActiveProviderManager { match Self::get_provider_config(input_name, &providers) { None => None, Some((lineup, _config)) => { - let cfg = lineup.get_next(); + let cfg = lineup.get_next().await; if log_enabled!(log::Level::Debug) { if let Some(ref c) = cfg { debug!("Using provider {}", c.name); @@ -652,14 +505,14 @@ impl ActiveProviderManager { pub async fn release_connection(&self, provider_name: &str) { let providers = self.providers.read().await; if let Some((lineup, _config)) = Self::get_provider_config(provider_name, &providers) { - lineup.release(provider_name); + lineup.release(provider_name).await; } } pub async fn active_connections(&self) -> Option> { let mut result = HashMap::::new(); let mut add_provider = |provider: &ProviderConfig| { - let count = provider.current_connections.load(Ordering::SeqCst); + let count = provider.get_current_connections(); if count > 0 { result.insert(provider.name.to_string(), count); } @@ -712,7 +565,8 @@ mod tests { macro_rules! should_available { ($lineup:expr, $provider_id:expr) => { - match $lineup.acquire() { + thread::sleep(std::time::Duration::from_millis(200)); + match $lineup.acquire() { ProviderAllocation::Exhausted => assert!(false, "Should available and not exhausted"), ProviderAllocation::Available(provider) => assert_eq!(provider.id, $provider_id), ProviderAllocation::GracePeriod(provider) => assert!(false, "Should available and not grace period: {}", provider.id), @@ -721,7 +575,8 @@ mod tests { } macro_rules! should_grace_period { ($lineup:expr, $provider_id:expr) => { - match $lineup.acquire() { + thread::sleep(std::time::Duration::from_millis(200)); + match $lineup.acquire() { ProviderAllocation::Exhausted => assert!(false, "Should grace period and not exhausted"), ProviderAllocation::Available(provider) => assert!(false, "Should grace period and not available: {}", provider.id), ProviderAllocation::GracePeriod(provider) => assert_eq!(provider.id, $provider_id), @@ -731,7 +586,8 @@ mod tests { macro_rules! should_exhausted { ($lineup:expr) => { - match $lineup.acquire() { + thread::sleep(std::time::Duration::from_millis(200)); + match $lineup.acquire() { ProviderAllocation::Exhausted => {}, ProviderAllocation::Available(provider) => assert!(false, "Should exhausted and not available: {}", provider.id), ProviderAllocation::GracePeriod(provider) => assert!(false, "Should exhausted and not grace period: {}", provider.id), @@ -739,8 +595,6 @@ mod tests { }; } - - // Helper function to create a ConfigInput instance fn create_config_input(id: u16, name: &str, priority: i16, max_connections: u16) -> ConfigInput { ConfigInput { @@ -901,6 +755,7 @@ mod tests { should_available!(lineup, 1); should_available!(lineup, 1); should_grace_period!(lineup, 1); + should_exhausted!(lineup); lineup.release("provider7_1"); should_grace_period!(lineup, 1); lineup.release("provider7_1"); diff --git a/src/api/model/active_user_manager.rs b/src/api/model/active_user_manager.rs index 509bbd980..607432657 100644 --- a/src/api/model/active_user_manager.rs +++ b/src/api/model/active_user_manager.rs @@ -1,9 +1,9 @@ -use std::collections::HashMap; -use std::sync::Arc; +use crate::model::api_proxy::UserConnectionPermission; use jsonwebtoken::get_current_timestamp; use log::{debug, info}; +use std::collections::HashMap; +use std::sync::Arc; use tokio::sync::RwLock; -use crate::model::api_proxy::UserConnectionPermission; pub struct UserConnectionGuard { manager: Arc, @@ -21,7 +21,7 @@ impl Drop for UserConnectionGuard { struct UserConnectionData { connections: u32, - granted_grace: u32, + granted_grace: bool, grace_ts: u64, } @@ -29,7 +29,7 @@ impl UserConnectionData { fn new() -> Self { Self { connections: 1, - granted_grace: 0, + granted_grace: false, grace_ts: 0, } } @@ -68,32 +68,59 @@ impl ActiveUserManager { 0 } - pub async fn connection_permission(&self, username: &str, max_connections: u32, grace_period: bool, grace_period_timeout_secs: u64) -> UserConnectionPermission { + pub async fn connection_permission( + &self, + username: &str, + max_connections: u32, + grace_period: bool, + grace_period_timeout_secs: u64, + ) -> UserConnectionPermission { if let Some(connection_data) = self.user.write().await.get_mut(username) { let current_connections = connection_data.connections; - if connection_data.granted_grace > 0 { - if (get_current_timestamp() - connection_data.grace_ts) <= grace_period_timeout_secs { + let now = get_current_timestamp(); + + if current_connections < max_connections { + // Reset grace period because user is back under max_connections + connection_data.granted_grace = false; + connection_data.grace_ts = 0; + return UserConnectionPermission::Allowed; + } + + // Check if user already used grace period + if connection_data.granted_grace { + if now - connection_data.grace_ts <= grace_period_timeout_secs { + // Grace timeout still active, deny connection debug!("User access denied, grace exhausted, too many connections: {username}"); return UserConnectionPermission::Exhausted; + } else { + // Grace timeout expired, reset grace counters + connection_data.granted_grace = false; + connection_data.grace_ts = 0; } - connection_data.granted_grace = 0; } - let extra_con = u32::from(grace_period); - if current_connections < max_connections + extra_con { - connection_data.granted_grace += 1; - connection_data.grace_ts = get_current_timestamp(); + + if current_connections < max_connections { + // everything ok + return UserConnectionPermission::Allowed; + } + + if grace_period && current_connections == max_connections { + // Allow grace period once + connection_data.granted_grace = true; + connection_data.grace_ts = now; debug!("Granted grace period for user access: {username}"); return UserConnectionPermission::GracePeriod; } - if current_connections >= max_connections { - debug!("User access denied, too many connections: {username}"); - return UserConnectionPermission::Exhausted; - } - connection_data.granted_grace = 0; + + // Too many connections, no grace allowed + debug!("User access denied, too many connections: {username}"); + return UserConnectionPermission::Exhausted; } + UserConnectionPermission::Allowed } + pub async fn active_users(&self) -> usize { self.user.read().await.len() } @@ -122,8 +149,13 @@ impl ActiveUserManager { async fn remove_connection(&self, username: &str) { let mut lock = self.user.write().await; if let Some(connection_data) = lock.get_mut(username) { - connection_data.connections -= 1; - connection_data.granted_grace = 0; + if connection_data.connections > 0 { + connection_data.connections -= 1; + } + + // DO NOT reset granted_grace or grace_ts here! + // We must preserve the grace period state until connection_permission() checks it. + if connection_data.connections == 0 { lock.remove(username); } diff --git a/src/api/model/mod.rs b/src/api/model/mod.rs index cc20b3d02..14f4fc5d8 100644 --- a/src/api/model/mod.rs +++ b/src/api/model/mod.rs @@ -8,4 +8,5 @@ pub(in crate::api) mod stream_error; pub(crate) mod streams; pub(in crate::api) mod active_user_manager; pub(in crate::api) mod active_provider_manager; -pub(in crate::api) mod stream; \ No newline at end of file +pub(in crate::api) mod stream; +pub(in crate::api) mod provider_config; \ No newline at end of file diff --git a/src/api/model/provider_config.rs b/src/api/model/provider_config.rs new file mode 100644 index 000000000..68516ddaf --- /dev/null +++ b/src/api/model/provider_config.rs @@ -0,0 +1,195 @@ +use std::ops::Deref; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, AtomicU16, AtomicU64, Ordering}; +use tokio::sync::RwLock; +use crate::api::model::active_provider_manager::ProviderAllocation; +use crate::model::config::{ConfigInput, ConfigInputAlias, InputType, InputUserInfo}; + +#[derive(Debug)] +pub enum ProviderConfigAllocation { + Exhausted, + Available, + GracePeriod, +} + +/// This struct represents an individual provider configuration with fields like: +/// +/// `id`, `name`, `url`, `username`, `password` +/// `input_type`: Determines the type of input the provider supports. +/// `max_connections`: Maximum allowed concurrent connections. +/// `priority`: Priority level for selecting providers. +/// `current_connections`: A `RwLock` to safely track the number of active connections. +#[derive(Debug)] +pub struct ProviderConfig { + pub id: u16, + pub name: String, + pub url: String, + pub username: Option, + pub password: Option, + pub input_type: InputType, + max_connections: u16, + priority: i16, + current_connections: AtomicU16, + granted_grace: AtomicBool, + grace_ts: AtomicU64, + lock: RwLock<()>, +} + +impl ProviderConfig { + pub fn new(cfg: &ConfigInput) -> Self { + Self { + id: cfg.id, + name: cfg.name.clone(), + url: cfg.url.clone(), + username: cfg.username.clone(), + password: cfg.password.clone(), + input_type: cfg.input_type, + max_connections: cfg.max_connections, + priority: cfg.priority, + current_connections: AtomicU16::new(0), + granted_grace: AtomicBool::new(false), + grace_ts: AtomicU64::new(0), + lock: RwLock::new(()), + } + } + + pub fn new_alias(cfg: &ConfigInput, alias: &ConfigInputAlias) -> Self { + Self { + id: alias.id, + name: alias.name.clone(), + url: alias.url.clone(), + username: alias.username.clone(), + password: alias.password.clone(), + input_type: cfg.input_type, + max_connections: alias.max_connections, + priority: alias.priority, + current_connections: AtomicU16::new(0), + granted_grace: AtomicBool::new(false), + grace_ts: AtomicU64::new(0), + lock: RwLock::new(()), + } + } + + pub fn get_user_info(&self) -> Option { + InputUserInfo::new(self.input_type, self.username.as_deref(), self.password.as_deref(), &self.url) + } + + #[inline] + pub fn is_exhausted(&self) -> bool { + let max = self.max_connections; + if max == 0 { + return false; + } + self.current_connections.load(Ordering::Acquire) >= max + } + + #[inline] + pub fn is_over_limit(&self) -> bool { + let max = self.max_connections; + if max == 0 { + return false; + } + self.current_connections.load(Ordering::Acquire) > max + } + + // + // #[inline] + // pub fn has_capacity(&self) -> bool { + // !self.is_exhausted() + // } + + + fn force_allocate(&self) { + self.current_connections.fetch_add(1, Ordering::SeqCst); + } + + async fn try_allocate(&self, grace: bool) -> ProviderConfigAllocation { + let _lock = self.lock.write().await; + let connections = self.current_connections.load(Ordering::SeqCst); + if self.max_connections == 0 { + self.current_connections.fetch_add(1, Ordering::SeqCst); + return ProviderConfigAllocation::Available; + } + if (!grace && connections < self.max_connections) || (grace && connections <= self.max_connections) { + self.current_connections.fetch_add(1, Ordering::SeqCst); + return if connections < self.max_connections { ProviderConfigAllocation::Available } else { ProviderConfigAllocation::GracePeriod }; + } + ProviderConfigAllocation::Exhausted + } + + // is intended to use with redirects, to cycle through provider + async fn get_next(&self, grace: bool) -> bool { + let _lock = self.lock.write().await; + let connections = self.current_connections.load(Ordering::SeqCst); + if self.max_connections == 0 { + return true; + } + if (!grace && connections < self.max_connections) || (grace && connections <= self.max_connections) { + return true; + } + false + } + + pub async fn release(&self) { + let _ = self.current_connections.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| { + if current > 0 { + Some(current - 1) + } else { + None + } + }); + } + + #[inline] + pub(crate) fn get_current_connections(&self) -> u16 { + self.current_connections.load(Ordering::SeqCst) + } + + #[inline] + pub(crate) fn get_priority(&self) -> i16 { + self.priority + } + +} + +#[derive(Clone, Debug)] +pub(in crate::api::model) struct ProviderConfigWrapper { + inner: Arc, +} + + +impl ProviderConfigWrapper { + pub fn new(cfg: ProviderConfig) -> Self { + Self { + inner: Arc::new(cfg) + } + } + + pub fn force_allocate(&self) -> ProviderAllocation { + self.inner.force_allocate(); + ProviderAllocation::Available(Arc::clone(&self.inner)) + } + + pub async fn try_allocate(&self, grace: bool) -> ProviderAllocation { + match self.inner.try_allocate(grace).await { + ProviderConfigAllocation::Available => ProviderAllocation::Available(Arc::clone(&self.inner)), + ProviderConfigAllocation::GracePeriod => ProviderAllocation::GracePeriod(Arc::clone(&self.inner)), + ProviderConfigAllocation::Exhausted => ProviderAllocation::Exhausted, + } + } + + pub async fn get_next(&self, grace: bool) -> Option> { + if self.inner.get_next(grace).await { + return Some(Arc::clone(&self.inner)); + } + None + } + +} +impl Deref for ProviderConfigWrapper { + type Target = ProviderConfig; + + fn deref(&self) -> &Self::Target { + &self.inner + } +} \ No newline at end of file