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/xmltv_api.rs b/src/api/endpoints/xmltv_api.rs index 2f762bf5a..66b762714 100644 --- a/src/api/endpoints/xmltv_api.rs +++ b/src/api/endpoints/xmltv_api.rs @@ -1,6 +1,3 @@ -use std::fs::File; -use std::path::{Path, PathBuf}; -use std::sync::Arc; use axum::response::IntoResponse; use chrono::{Duration, NaiveDateTime, TimeDelta}; use flate2::write::GzEncoder; @@ -8,6 +5,9 @@ use flate2::Compression; use log::{error, trace}; use quick_xml::events::{BytesStart, Event}; use quick_xml::{Reader, Writer}; +use std::fs::File; +use std::path::{Path, PathBuf}; +use std::sync::Arc; use crate::api::api_utils::{get_user_target, serve_file}; use crate::api::model::app_state::AppState; @@ -170,25 +170,48 @@ fn serve_epg_with_timeshift(epg_file: File, offset_minutes: i32) -> impl axum::r .into_response() } +/// Handles XMLTV EPG API requests, serving the appropriate EPG file with optional time-shifting based on user configuration. +/// +/// Returns a 403 Forbidden response if the user or target is invalid or if the user lacks permission. If no EPG file is configured for the target, returns an empty EPG response. Otherwise, serves the EPG file, applying a time shift if specified by the user. +/// +/// # Examples +/// +/// ``` +/// // Example usage within an Axum router: +/// let router = xmltv_api_register(); +/// // A GET request to /xmltv.php with valid query parameters will invoke this handler. +/// ``` async fn xmltv_api( axum::extract::Query(api_req): axum::extract::Query, axum::extract::State(app_state): axum::extract::State>, -) -> impl axum::response::IntoResponse + Send { - if let Some((user, target)) = get_user_target(&api_req, &app_state).await { - if user.permission_denied(&app_state) { - return axum::http::StatusCode::FORBIDDEN.into_response(); - } - match get_epg_path_for_target(&app_state.config, target) { - None => { - // No epg configured, No processing or timeshift, epg can't be mapped to the channels. - // we do not deliver epg - } - Some(epg_path) => return serve_epg(&epg_path, &user).await.into_response() - } +) -> impl IntoResponse + Send { + let Some((user, target)) = get_user_target(&api_req, &app_state).await else { + return axum::http::StatusCode::FORBIDDEN.into_response(); + }; + + if user.permission_denied(&app_state) { + return axum::http::StatusCode::FORBIDDEN.into_response(); } - get_empty_epg_response().into_response() + + let Some(epg_path) = get_epg_path_for_target(&app_state.config, target) else { + // No epg configured, No processing or timeshift, epg can't be mapped to the channels. + // we do not deliver epg + return get_empty_epg_response().into_response(); + }; + + serve_epg(&epg_path, &user).await.into_response() } +/// Registers the XMLTV EPG API routes for handling HTTP GET requests. +/// +/// The returned router maps the `/xmltv.php`, `/update/epg.php`, and `/epg` endpoints to the `xmltv_api` handler, enabling XMLTV EPG data retrieval with optional time-shifting and compression. +/// +/// # Examples +/// +/// ``` +/// let router = xmltv_api_register(); +/// // The router can now be used with an Axum server. +/// ``` pub fn xmltv_api_register() -> axum::Router> { axum::Router::new() .route("/xmltv.php", axum::routing::get(xmltv_api)) 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/main_api.rs b/src/api/main_api.rs index 4d7b98299..7719f6687 100644 --- a/src/api/main_api.rs +++ b/src/api/main_api.rs @@ -68,7 +68,7 @@ async fn create_shared_data(cfg: &Arc) -> AppState { } }); - let active_users = Arc::new(ActiveUserManager::new(cfg.log.as_ref().is_some_and(|l| l.log_active_user))); + let active_users = Arc::new(ActiveUserManager::new(cfg)); let active_provider = Arc::new(ActiveProviderManager::new(cfg).await); let mut builder = Client::builder().http1_only(); diff --git a/src/api/model/active_provider_manager.rs b/src/api/model/active_provider_manager.rs index f77ecfee7..8dc20c7f4 100644 --- a/src/api/model/active_provider_manager.rs +++ b/src/api/model/active_provider_manager.rs @@ -1,10 +1,12 @@ -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}; +use crate::utils::default_utils::{default_grace_period_millis, default_grace_period_timeout_secs}; pub struct ProviderConnectionGuard { manager: Arc, @@ -62,156 +64,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 +76,24 @@ enum ProviderLineup { } impl ProviderLineup { - fn get_next(&self) -> Option> { + async fn get_next(&self, grace_period_timeout_secs: u64) -> Option> { match self { - ProviderLineup::Single(lineup) => lineup.get_next(), - ProviderLineup::Multi(lineup) => lineup.get_next(), + ProviderLineup::Single(lineup) => lineup.get_next(grace_period_timeout_secs).await, + ProviderLineup::Multi(lineup) => lineup.get_next(grace_period_timeout_secs).await, } } - fn acquire(&self) -> ProviderAllocation { + async fn acquire(&self, grace_period_timeout_secs: u64) -> ProviderAllocation { match self { - ProviderLineup::Single(lineup) => lineup.acquire(), - ProviderLineup::Multi(lineup) => lineup.acquire(), + ProviderLineup::Single(lineup) => lineup.acquire(grace_period_timeout_secs).await, + ProviderLineup::Multi(lineup) => lineup.acquire(grace_period_timeout_secs).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 +111,17 @@ impl SingleProviderLineup { } } - fn get_next(&self) -> Option> { - self.provider.get_next(false) + async fn get_next(&self, grace_period_timeout_secs: u64) -> Option> { + self.provider.get_next(false, grace_period_timeout_secs).await } - fn acquire(&self) -> ProviderAllocation { - self.provider.try_allocate(true) + async fn acquire(&self, grace_period_timeout_secs: u64) -> ProviderAllocation { + self.provider.try_allocate(true, grace_period_timeout_secs).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,16 +134,16 @@ impl SingleProviderLineup { #[derive(Debug)] enum ProviderPriorityGroup { SingleProviderGroup(ProviderConfigWrapper), - MultiProviderGroup(AtomicUsize, Vec), + MultiProviderGroup(Mutex, Vec), } impl ProviderPriorityGroup { - fn is_exhausted(&self) -> bool { + async fn is_exhausted(&self) -> bool { match self { - ProviderPriorityGroup::SingleProviderGroup(g) => g.is_exhausted(), + ProviderPriorityGroup::SingleProviderGroup(g) => g.is_exhausted().await, ProviderPriorityGroup::MultiProviderGroup(_, groups) => { for g in groups { - if !g.is_exhausted() { + if !g.is_exhausted().await { return false; } } @@ -320,7 +172,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 +181,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 +220,61 @@ 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, grace_period_timeout_secs: u64) -> ProviderAllocation { match priority_group { ProviderPriorityGroup::SingleProviderGroup(p) => { - let result = p.try_allocate(grace); + let result = p.try_allocate(grace, grace_period_timeout_secs).await; match result { ProviderAllocation::Exhausted => {} ProviderAllocation::Available(_) | ProviderAllocation::GracePeriod(_) => return result } } ProviderPriorityGroup::MultiProviderGroup(index, pg) => { - let mut idx = index.load(Ordering::SeqCst); let provider_count = pg.len(); + let mut idx = { + *index.lock().await + }; 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, grace_period_timeout_secs).await; match result { ProviderAllocation::Exhausted => {} ProviderAllocation::Available(_) | ProviderAllocation::GracePeriod(_) => { - index.store(idx, Ordering::SeqCst); + *index.lock().await = idx; return result; } } } - index.store(idx, Ordering::SeqCst); + *index.lock().await = 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, grace_period_timeout_secs: u64) -> Option> { match priority_group { ProviderPriorityGroup::SingleProviderGroup(p) => { - return p.get_next(grace); + return p.get_next(grace, grace_period_timeout_secs).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, grace_period_timeout_secs).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 +305,16 @@ impl MultiProviderLineup { /// ProviderAllocation::GracePeriodprovider) => println!("Provider with grace period {}", provider.name), /// } /// ``` - fn acquire(&self) -> ProviderAllocation { + async fn acquire(&self, grace_period_timeout_secs: u64) -> 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, grace_period_timeout_secs).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, grace_period_timeout_secs).await } else { without_grace_allocation } @@ -467,7 +323,7 @@ impl MultiProviderLineup { ProviderAllocation::Exhausted => {} ProviderAllocation::Available(_) | ProviderAllocation::GracePeriod(_) => { - if priority_group.is_exhausted() { + if priority_group.is_exhausted().await { self.index.store((index + 1) % provider_count, Ordering::SeqCst); } return allocation; @@ -479,16 +335,16 @@ impl MultiProviderLineup { } // it intended to use with redirects to cycle through provider - fn get_next(&self) -> Option> { + async fn get_next(&self, grace_period_timeout_secs: u64) -> 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, grace_period_timeout_secs).await; if config.is_none() { - Self::get_next_provider_from_group(priority_group, true) + Self::get_next_provider_from_group(priority_group, true, grace_period_timeout_secs).await } else { config } @@ -496,7 +352,7 @@ impl MultiProviderLineup { match allocation { None => {} Some(config) => { - if priority_group.is_exhausted() { + if priority_group.is_exhausted().await { self.index.store((index + 1) % provider_count, Ordering::SeqCst); } return Some(config); @@ -508,19 +364,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; } } @@ -531,12 +387,20 @@ impl MultiProviderLineup { } pub struct ActiveProviderManager { + grace_period_millis: u64, + grace_period_timeout_secs: u64, providers: Arc>>, } impl ActiveProviderManager { pub async fn new(cfg: &Config) -> Self { + let (grace_period_millis, grace_period_timeout_secs) = cfg.reverse_proxy.as_ref() + .and_then(|r| r.stream.as_ref()) + .map_or_else(|| (default_grace_period_millis(), default_grace_period_timeout_secs()), |s| (s.grace_period_millis, s.grace_period_timeout_secs)); + let mut this = Self { + grace_period_millis, + grace_period_timeout_secs, providers: Arc::new(RwLock::new(Vec::new())), }; for source in &cfg.sources { @@ -549,6 +413,8 @@ impl ActiveProviderManager { fn clone_inner(&self) -> Self { Self { + grace_period_millis: self.grace_period_millis, + grace_period_timeout_secs: self.grace_period_timeout_secs, providers: Arc::clone(&self.providers), } } @@ -597,7 +463,7 @@ impl ActiveProviderManager { let providers = self.providers.read().await; let allocation = match Self::get_provider_config(provider_name, &providers) { None => ProviderAllocation::Exhausted, // No Name matched, we don't have this provider - Some((_lineup, config)) => config.force_allocate(), + Some((_lineup, config)) => config.force_allocate().await, }; ProviderConnectionGuard { @@ -611,7 +477,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(self.grace_period_timeout_secs).await }; if log_enabled!(log::Level::Debug) { @@ -637,7 +503,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(self.grace_period_timeout_secs).await; if log_enabled!(log::Level::Debug) { if let Some(ref c) = cfg { debug!("Using provider {}", c.name); @@ -652,14 +518,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); + pub async fn active_connections(&self) -> Option> { + let mut result = HashMap::::new(); + let mut add_provider = async |provider: &ProviderConfig| { + let count = provider.get_current_connections().await; if count > 0 { result.insert(provider.name.to_string(), count); } @@ -668,17 +534,17 @@ impl ActiveProviderManager { for lineup in &*providers { match lineup { ProviderLineup::Single(provider_lineup) => { - add_provider(&provider_lineup.provider); + add_provider(&provider_lineup.provider).await; } ProviderLineup::Multi(provider_lineup) => { for provider_group in &provider_lineup.providers { match provider_group { ProviderPriorityGroup::SingleProviderGroup(provider) => { - add_provider(provider); + add_provider(provider).await; } ProviderPriorityGroup::MultiProviderGroup(_, providers) => { for provider in providers { - add_provider(provider); + add_provider(provider).await; } } } @@ -696,7 +562,7 @@ impl ActiveProviderManager { pub async fn is_over_limit(&self, provider_name: &str) -> bool { let providers = self.providers.read().await; if let Some((_, config)) = Self::get_provider_config(provider_name, &providers) { - config.is_over_limit() + config.is_over_limit().await } else { false } @@ -705,14 +571,16 @@ impl ActiveProviderManager { #[cfg(test)] mod tests { + use std::sync::atomic::AtomicU16; use super::*; - use crate::model::config::InputFetchMethod; + use crate::model::config::{ConfigInputAlias, InputFetchMethod, InputType}; use crate::Arc; use std::thread; macro_rules! should_available { - ($lineup:expr, $provider_id:expr) => { - match $lineup.acquire() { + ($lineup:expr, $provider_id:expr, $grace_period_timeout_secs: expr) => { + thread::sleep(std::time::Duration::from_millis(200)); + match $lineup.acquire($grace_period_timeout_secs).await { 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), @@ -720,8 +588,9 @@ mod tests { }; } macro_rules! should_grace_period { - ($lineup:expr, $provider_id:expr) => { - match $lineup.acquire() { + ($lineup:expr, $provider_id:expr, $grace_period_timeout_secs: expr) => { + thread::sleep(std::time::Duration::from_millis(200)); + match $lineup.acquire($grace_period_timeout_secs).await { 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), @@ -730,8 +599,9 @@ mod tests { } macro_rules! should_exhausted { - ($lineup:expr) => { - match $lineup.acquire() { + ($lineup:expr, $grace_period_timeout_secs: expr) => { + thread::sleep(std::time::Duration::from_millis(200)); + match $lineup.acquire($grace_period_timeout_secs).await { 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 +609,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 { @@ -790,15 +658,19 @@ mod tests { // Create MultiProviderLineup with the provider and alias let lineup = MultiProviderLineup::new(&input); + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { + // Test that the alias provider is available + should_available!(lineup, 1, 5); + // Try acquiring again + should_available!(lineup, 2, 5); + should_available!(lineup, 2, 5); + should_grace_period!(lineup, 1, 5); + should_grace_period!(lineup, 2, 5); + should_exhausted!(lineup, 5); + should_exhausted!(lineup, 5); + }); - // Test that the alias provider is available - should_available!(lineup, 1); - // Try acquiring again - should_available!(lineup, 2); - should_available!(lineup, 2); - should_grace_period!(lineup, 1); - should_grace_period!(lineup, 2); - should_exhausted!(lineup); } // // Test acquiring from a MultiProviderLineup where the alias has a different priority @@ -810,10 +682,13 @@ mod tests { input.aliases = Some(vec![alias]); let lineup = MultiProviderLineup::new(&input); // The alias has a higher priority, so the alias should be acquired first - for _ in 0..2 { - should_available!(lineup, 2); - } - should_available!(lineup, 1); + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { + for _ in 0..2 { + should_available!(lineup, 2, 5); + } + should_available!(lineup, 1, 5); + }); } // Test provider when there are multiple aliases, all with distinct priorities @@ -827,20 +702,22 @@ mod tests { input.aliases = Some(vec![alias1, alias2]); let lineup = MultiProviderLineup::new(&input); + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { + // The alias with priority 0 should be acquired first (higher priority) + should_available!(lineup, 3, 5); + // Acquire again, and provider should still be available (with remaining capacity) + should_available!(lineup, 1, 5); + // // Check that the second alias with priority 2 is considered next + should_available!(lineup, 2, 5); + should_available!(lineup, 2, 5); - // The alias with priority 0 should be acquired first (higher priority) - should_available!(lineup, 3); - // Acquire again, and provider should still be available (with remaining capacity) - should_available!(lineup, 1); - // // Check that the second alias with priority 2 is considered next - should_available!(lineup, 2); - should_available!(lineup, 2); + should_grace_period!(lineup, 3, 5); + should_grace_period!(lineup, 1, 5); + should_grace_period!(lineup, 2, 5); - should_grace_period!(lineup, 3); - should_grace_period!(lineup, 1); - should_grace_period!(lineup, 2); - - should_exhausted!(lineup); + should_exhausted!(lineup, 5); + }); } @@ -855,23 +732,25 @@ mod tests { input.aliases = Some(vec![alias1, alias2]); let lineup = MultiProviderLineup::new(&input); + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { + // Acquire connection from alias2 + should_available!(lineup, 3, 5); + // Acquire connection from provider1 + should_available!(lineup, 1, 5); + // Acquire connection from alias1 + should_available!(lineup, 2, 5); - // Acquire connection from alias2 - should_available!(lineup, 3); - // Acquire connection from provider1 - should_available!(lineup, 1); - // Acquire connection from alias1 - should_available!(lineup, 2); + // Acquire connection from alias2 + should_grace_period!(lineup, 3, 5); + // Acquire connection from provider1 + should_grace_period!(lineup, 1, 5); + // Acquire connection from alias1 + should_grace_period!(lineup, 2, 5); - // Acquire connection from alias2 - should_grace_period!(lineup, 3); - // Acquire connection from provider1 - should_grace_period!(lineup, 1); - // Acquire connection from alias1 - should_grace_period!(lineup, 2); - - // Now, all are exhausted - should_exhausted!(lineup); + // Now, all are exhausted + should_exhausted!(lineup, 5); + }); } // Test acquiring a connection when there is available capacity @@ -879,15 +758,17 @@ mod tests { fn test_acquire_when_capacity_available() { let cfg = create_config_input(1, "provider5_1", 1, 2); let lineup = SingleProviderLineup::new(&cfg); - - // First acquire attempt should succeed - should_available!(lineup, 1); - // Second acquire attempt should succeed as well - should_available!(lineup, 1); - // Third with grace time - should_grace_period!(lineup, 1); - // Fourth acquire attempt should fail as the provider is exhausted - should_exhausted!(lineup); + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { + // First acquire attempt should succeed + should_available!(lineup, 1, 5); + // Second acquire attempt should succeed as well + should_available!(lineup, 1, 5); + // Third with grace time + should_grace_period!(lineup, 1, 5); + // Fourth acquire attempt should fail as the provider is exhausted + should_exhausted!(lineup, 5); + }); } @@ -896,18 +777,21 @@ mod tests { fn test_release_connection() { let cfg = create_config_input(1, "provider7_1", 1, 2); let lineup = SingleProviderLineup::new(&cfg); - - // Acquire two connections - should_available!(lineup, 1); - should_available!(lineup, 1); - should_grace_period!(lineup, 1); - lineup.release("provider7_1"); - should_grace_period!(lineup, 1); - lineup.release("provider7_1"); - lineup.release("provider7_1"); - should_available!(lineup, 1); - should_grace_period!(lineup, 1); - should_exhausted!(lineup); + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { + // Acquire two connections + should_available!(lineup, 1, 5); + should_available!(lineup, 1, 5); + should_grace_period!(lineup, 1, 5); + should_exhausted!(lineup, 5); + lineup.release("provider7_1").await; + should_grace_period!(lineup, 1, 5); + lineup.release("provider7_1").await; + lineup.release("provider7_1").await; + should_available!(lineup, 1, 5); + should_grace_period!(lineup, 1, 5); + should_exhausted!(lineup, 5); + }); } // Test acquiring with MultiProviderLineup and round-robin allocation @@ -921,28 +805,31 @@ mod tests { // Create MultiProviderLineup with the provider and alias let lineup = MultiProviderLineup::new(&cfg1); + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { - // Test acquiring the first provider - should_available!(lineup, 1); + // Test acquiring the first provider + should_available!(lineup, 1, 5); - // Test acquiring the second provider - should_available!(lineup, 2); + // Test acquiring the second provider + should_available!(lineup, 2, 5); - // Test acquiring the first provider - should_available!(lineup, 1); + // Test acquiring the first provider + should_available!(lineup, 1, 5); - should_grace_period!(lineup, 1); - should_grace_period!(lineup, 2); + should_grace_period!(lineup, 1, 5); + should_grace_period!(lineup, 2, 5); - lineup.release("provider8_1"); - lineup.release("alias_2"); - lineup.release("provider8_1"); + lineup.release("provider8_1").await; + lineup.release("alias_2").await; + lineup.release("provider8_1").await; - should_available!(lineup, 1); - should_grace_period!(lineup, 1); - should_grace_period!(lineup, 2); + should_available!(lineup, 1, 5); + should_grace_period!(lineup, 1, 5); + should_grace_period!(lineup, 2, 5); - should_exhausted!(lineup); + should_exhausted!(lineup, 5); + }); } // Test concurrent access to `acquire` using multiple threads @@ -951,8 +838,6 @@ mod tests { let cfg = create_config_input(1, "provider9_1", 1, 2); let lineup = Arc::new(SingleProviderLineup::new(&cfg)); - let mut handles = vec![]; - let available_count = Arc::new(AtomicU16::new(2)); let grace_period_count = Arc::new(AtomicU16::new(1)); let exhausted_count = Arc::new(AtomicU16::new(2)); @@ -962,22 +847,15 @@ mod tests { let available = Arc::clone(&available_count); let grace_period = Arc::clone(&grace_period_count); let exhausted = Arc::clone(&exhausted_count); - let handle = thread::spawn(move || { - // Each thread tries to acquire a connection - match lineup_clone.acquire() { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async move { + match lineup_clone.acquire(5).await { ProviderAllocation::Exhausted => exhausted.fetch_sub(1, Ordering::SeqCst), ProviderAllocation::Available(_) => available.fetch_sub(1, Ordering::SeqCst), ProviderAllocation::GracePeriod(_) => grace_period.fetch_sub(1, Ordering::SeqCst), } }); - handles.push(handle); } - - // Join all threads to ensure completion - for handle in handles { - handle.join().unwrap(); - } - assert_eq!(exhausted_count.load(Ordering::SeqCst), 0); assert_eq!(available_count.load(Ordering::SeqCst), 0); assert_eq!(grace_period_count.load(Ordering::SeqCst), 0); diff --git a/src/api/model/active_user_manager.rs b/src/api/model/active_user_manager.rs index 509bbd980..bd7a9f0af 100644 --- a/src/api/model/active_user_manager.rs +++ b/src/api/model/active_user_manager.rs @@ -1,9 +1,11 @@ -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; +use crate::model::config::Config; +use crate::utils::default_utils::{default_grace_period_millis, default_grace_period_timeout_secs}; pub struct UserConnectionGuard { manager: Arc, @@ -21,7 +23,7 @@ impl Drop for UserConnectionGuard { struct UserConnectionData { connections: u32, - granted_grace: u32, + granted_grace: bool, grace_ts: u64, } @@ -29,26 +31,29 @@ impl UserConnectionData { fn new() -> Self { Self { connections: 1, - granted_grace: 0, + granted_grace: false, grace_ts: 0, } } } pub struct ActiveUserManager { + grace_period_millis: u64, + grace_period_timeout_secs: u64, log_active_user: bool, user: Arc>>, } -impl Default for ActiveUserManager { - fn default() -> Self { - Self::new(false) - } -} - impl ActiveUserManager { - pub fn new(log_active_user: bool) -> Self { + pub fn new(config: &Config) -> Self { + let log_active_user = config.log.as_ref().is_some_and(|l| l.log_active_user); + let (grace_period_millis, grace_period_timeout_secs) = config.reverse_proxy.as_ref() + .and_then(|r| r.stream.as_ref()) + .map_or_else(|| (default_grace_period_millis(), default_grace_period_timeout_secs()), |s| (s.grace_period_millis, s.grace_period_timeout_secs)); + Self { + grace_period_millis, + grace_period_timeout_secs, log_active_user, user: Arc::new(RwLock::new(HashMap::new())), } @@ -56,6 +61,8 @@ impl ActiveUserManager { fn clone_inner(&self) -> Self { Self { + grace_period_millis: self.grace_period_millis, + grace_period_timeout_secs: self.grace_period_timeout_secs, log_active_user: self.log_active_user, user: Arc::clone(&self.user), } @@ -68,32 +75,53 @@ impl ActiveUserManager { 0 } - 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 { - debug!("User access denied, grace exhausted, too many connections: {username}"); - return UserConnectionPermission::Exhausted; + pub async fn connection_permission( + &self, + username: &str, + max_connections: u32, + ) -> UserConnectionPermission { + if max_connections > 0 { + if let Some(connection_data) = self.user.write().await.get_mut(username) { + let current_connections = connection_data.connections; + + 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; } - 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(); - debug!("Granted grace period for user access: {username}"); - return UserConnectionPermission::GracePeriod; - } - if current_connections >= max_connections { + + let now = get_current_timestamp(); + // Check if user already used grace period + if connection_data.granted_grace { + if now - connection_data.grace_ts <= self.grace_period_timeout_secs { + // Grace timeout still active, deny connection + debug!("User access denied, grace exhausted, too many connections: {username}"); + return UserConnectionPermission::Exhausted; + } + // Grace timeout expired, reset grace counters + connection_data.granted_grace = false; + connection_data.grace_ts = 0; + } + + if self.grace_period_millis > 0 && 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; + } + + // Too many connections, no grace allowed debug!("User access denied, too many connections: {username}"); return UserConnectionPermission::Exhausted; } - connection_data.granted_grace = 0; } + UserConnectionPermission::Allowed } + pub async fn active_users(&self) -> usize { self.user.read().await.len() } @@ -122,8 +150,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/app_state.rs b/src/api/model/app_state.rs index 829679109..9eefff0a0 100644 --- a/src/api/model/app_state.rs +++ b/src/api/model/app_state.rs @@ -8,7 +8,6 @@ use crate::model::api_proxy::UserConnectionPermission; use crate::model::config::{Config}; use crate::model::hdhomerun_config::HdHomeRunDeviceConfig; use crate::tools::lru_cache::LRUResourceCache; -use crate::utils::default_utils::{default_grace_period_millis, default_grace_period_timeout_secs}; #[derive(Clone)] pub struct AppState { @@ -27,10 +26,7 @@ impl AppState { } pub async fn get_connection_permission(&self, username: &str, max_connections: u32) -> UserConnectionPermission { - let (grace_period_millis, grace_period_timeout_secs) = self.config.reverse_proxy.as_ref() - .and_then(|r| r.stream.as_ref()) - .map_or_else(|| (default_grace_period_millis(), default_grace_period_timeout_secs()), |s| (s.grace_period_millis, s.grace_period_timeout_secs)); - self.active_users.connection_permission(username, max_connections, grace_period_millis > 0, grace_period_timeout_secs).await + self.active_users.connection_permission(username, max_connections).await } } 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..ba519923e --- /dev/null +++ b/src/api/model/provider_config.rs @@ -0,0 +1,223 @@ +use crate::api::model::active_provider_manager::ProviderAllocation; +use crate::model::config::{ConfigInput, ConfigInputAlias, InputType, InputUserInfo}; +use std::ops::Deref; +use std::sync::Arc; +use jsonwebtoken::get_current_timestamp; +use log::debug; +use tokio::sync::RwLock; + +#[derive(Debug)] +pub enum ProviderConfigAllocation { + Exhausted, + Available, + GracePeriod, +} + +#[derive(Debug, Default)] +struct ProviderConfigConnection { + current_connections: usize, + granted_grace: bool, + grace_ts: u64, +} + +/// 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: usize, + priority: i16, + connection: 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 as usize, + priority: cfg.priority, + connection: RwLock::new(ProviderConfigConnection::default()), + } + } + + 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 as usize, + priority: alias.priority, + connection: RwLock::new(ProviderConfigConnection::default()), + } + } + + pub fn get_user_info(&self) -> Option { + InputUserInfo::new(self.input_type, self.username.as_deref(), self.password.as_deref(), &self.url) + } + + #[inline] + pub async fn is_exhausted(&self) -> bool { + let max = self.max_connections; + if max == 0 { + return false; + } + self.connection.read().await.current_connections >= max + } + + #[inline] + pub async fn is_over_limit(&self) -> bool { + let max = self.max_connections; + if max == 0 { + return false; + } + self.connection.read().await.current_connections > max + } + + // + // #[inline] + // pub fn has_capacity(&self) -> bool { + // !self.is_exhausted() + // } + + + async fn force_allocate(&self) { + let mut guard = self.connection.write().await; + guard.current_connections += 1; + } + + async fn try_allocate(&self, grace: bool, grace_period_timeout_secs: u64) -> ProviderConfigAllocation { + let mut guard = self.connection.write().await; + if self.max_connections == 0 { + guard.current_connections += 1; + return ProviderConfigAllocation::Available; + } + 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; + return ProviderConfigAllocation::Available; + } + + let now = get_current_timestamp(); + if guard.granted_grace { + if now - guard.grace_ts <= grace_period_timeout_secs { + // Grace timeout still active, deny connection + debug!("Provider access denied, grace exhausted, too many connections: {}", self.name); + return ProviderConfigAllocation::Exhausted; + } + // Grace timeout expired, reset grace counters + guard.granted_grace = false; + guard.grace_ts = 0; + } + guard.granted_grace = true; + guard.grace_ts = now; + guard.current_connections += 1; + return ProviderConfigAllocation::GracePeriod + } + ProviderConfigAllocation::Exhausted + } + + // is intended to use with redirects, to cycle through provider + async fn get_next(&self, grace: bool, grace_period_timeout_secs: u64) -> bool { + if self.max_connections == 0 { + return true; + } + let mut guard = self.connection.write().await; + let connections = guard.current_connections; + if connections < self.max_connections || (grace && connections <= self.max_connections) { + let now = get_current_timestamp(); + if guard.granted_grace { + if now - guard.grace_ts <= grace_period_timeout_secs { + // Grace timeout still active, deny connection + debug!("Provider access denied, grace exhausted, too many connections: {}", self.name); + return false; + } + // Grace timeout expired, reset grace counters + guard.granted_grace = false; + guard.grace_ts = 0; + } + return true; + } + false + } + + pub async fn release(&self) { + // DO NOT reset granted_grace or grace_ts here! + // We must preserve the grace period state until allocate() checks it. + let mut guard = self.connection.write().await; + if guard.current_connections > 0 { + guard.current_connections -= 1; + } + } + + #[inline] + pub(crate) async fn get_current_connections(&self) -> usize { + self.connection.read().await.current_connections + } + + #[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 async fn force_allocate(&self) -> ProviderAllocation { + self.inner.force_allocate().await; + ProviderAllocation::Available(Arc::clone(&self.inner)) + } + + pub async fn try_allocate(&self, grace: bool, grace_period_timeout_secs: u64) -> ProviderAllocation { + match self.inner.try_allocate(grace, grace_period_timeout_secs).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, grace_period_timeout_secs: u64) -> Option> { + if self.inner.get_next(grace, grace_period_timeout_secs).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 diff --git a/src/model/healthcheck.rs b/src/model/healthcheck.rs index 47d68896a..d6892bbe5 100644 --- a/src/model/healthcheck.rs +++ b/src/model/healthcheck.rs @@ -23,5 +23,5 @@ pub struct StatusCheck { pub active_users: usize, pub active_user_connections: usize, #[serde(skip_serializing_if = "Option::is_none")] - pub active_provider_connections: Option>, + pub active_provider_connections: Option>, } \ No newline at end of file diff --git a/src/processing/parser/m3u.rs b/src/processing/parser/m3u.rs index eadd26afc..1e5fd401c 100644 --- a/src/processing/parser/m3u.rs +++ b/src/processing/parser/m3u.rs @@ -5,37 +5,37 @@ use crate::utils::hash_utils::extract_id_from_url; use crate::utils::string_utils; #[inline] -fn token_value(it: &mut std::str::Chars) -> String { +fn token_value(stack: &mut String, it: &mut std::str::Chars) -> String { // Use .find() to skip until the first double quote (") character. if it.any(|ch| ch == '"') { // If a quote is found, call get_value to extract the value. - return get_value(it); + return get_value(stack, it); } // If no double quote is found, return an empty string. String::new() } -fn get_value(it: &mut std::str::Chars) -> String { - let mut result = String::with_capacity(128); - for oc in it.by_ref() { - if oc == '"' { +fn get_value(stack: &mut String, it: &mut std::str::Chars) -> String { + for c in it.skip_while(|c| c.is_whitespace()) { + if c == '"' { break; } - result.push(oc); + stack.push(c); } - result.shrink_to_fit(); + + let result = (*stack).to_string(); + stack.clear(); result } -fn token_till(it: &mut std::str::Chars, stop_char: char, start_with_alpha: bool) -> Option { - let mut result = String::with_capacity(128); +fn token_till(stack: &mut String, it: &mut std::str::Chars, stop_char: char, start_with_alpha: bool) -> Option { let mut skip_non_alpha = start_with_alpha; for ch in it.by_ref() { if ch == stop_char { break; } - if result.is_empty() && ch.is_whitespace() { + if stack.is_empty() && ch.is_whitespace() { continue; } @@ -46,13 +46,14 @@ fn token_till(it: &mut std::str::Chars, stop_char: char, start_with_alpha: bool) continue; } } - result.push(ch); + stack.push(ch); } - if result.is_empty() { + if stack.is_empty() { None } else { - result.shrink_to_fit(); + let result = (*stack).to_string(); + stack.clear(); Some(result) } } @@ -91,22 +92,27 @@ macro_rules! process_header_fields { }; } -fn process_header(input: &ConfigInput, video_suffixes: &[&str], content: &str, url: &str) -> PlaylistItemHeader { - let mut plih = create_empty_playlistitem_header(input.name.as_str(), url); +fn process_header(input_name: &str, video_suffixes: &[&str], content: &str, url: &str) -> PlaylistItemHeader { + let mut plih = create_empty_playlistitem_header(input_name, url); let mut it = content.chars(); - let line_token = token_till(&mut it, ':', false); + let mut stack = String::with_capacity(64); + let line_token = token_till(&mut stack, &mut it, ':', false); if line_token.as_deref() == Some("#EXTINF") { let mut c = skip_digit(&mut it); loop { if c.is_none() { break; } - if c.unwrap() == ',' { - plih.title = get_value(&mut it); + let chr = c.unwrap(); + if chr.is_whitespace() { + // skip + } else if chr == ',' { + plih.title = get_value(&mut stack, &mut it); } else { - let token = token_till(&mut it, '=', true); + stack.push(chr); + let token = token_till(&mut stack, &mut it, '=', true); if let Some(t) = token { - let value = token_value(&mut it); + let value = token_value(&mut stack, &mut it); process_header_fields!(plih, t.to_lowercase().as_str(), (id, "tvg-id"), (group, "group-title"), @@ -161,6 +167,7 @@ where { let mut header: Option = None; let mut group: Option = None; + let input_name = input.name.as_str(); let video_suffixes = cfg.video.as_ref().unwrap().extensions.iter().map(String::as_str).collect::>(); for line in lines { @@ -176,7 +183,7 @@ where continue; } if let Some(header_value) = header { - let mut item = PlaylistItem { header: process_header(input, &video_suffixes, &header_value, line) }; + let mut item = PlaylistItem { header: process_header(input_name, &video_suffixes, &header_value, line) }; let header = &mut item.header; if header.group.is_empty() { if let Some(group_value) = group { @@ -222,9 +229,44 @@ where // create a group based on the first playlist item let channel = channels.first(); let (cluster, group_title) = channel.map(|pli| - (pli.header.xtream_cluster, &pli.header.group)).unwrap(); + (pli.header.xtream_cluster, &pli.header.group)).unwrap(); grp_id += 1; PlaylistGroup { id: grp_id, xtream_cluster: cluster, title: group_title.to_string(), channels } }).collect(); result } + +#[cfg(test)] +mod test { + use crate::processing::parser::m3u::process_header; + + #[test] + fn test_process_header_1() { + let input: &str = "hello"; + let video_suffixes = Vec::new(); + let url = "http://hello.de/hello.ts"; + let line = r#"#EXTINF:-1 channel-id="abc-seven" tvg-id="abc-seven" tvg-logo="https://abc.nz/.images/seven.png" tvg-chno="7" group-title="Sydney" , Seven"#; + + let pli = process_header(input, &video_suffixes, line, url); + assert_eq!(pli.title, "Seven"); + assert_eq!(pli.id, "abc-seven"); + assert_eq!(pli.logo, "https://abc.nz/.images/seven.png"); + assert_eq!(pli.chno, "7"); + assert_eq!(pli.group, "Sydney"); + } + + #[test] + fn test_process_header_2() { + let input: &str = "hello"; + let video_suffixes = Vec::new(); + let url = "http://hello.de/hello.ts"; + let line = r#"#EXTINF:-1 channel-id="abc-seven" tvg-id="abc-seven" tvg-logo="https://abc.nz/.images/seven.png" tvg-chno="7" group-title="Sydney", Seven"#; + + let pli = process_header(input, &video_suffixes, line, url); + assert_eq!(pli.title, "Seven"); + assert_eq!(pli.id, "abc-seven"); + assert_eq!(pli.logo, "https://abc.nz/.images/seven.png"); + assert_eq!(pli.chno, "7"); + assert_eq!(pli.group, "Sydney"); + } +} \ No newline at end of file diff --git a/src/utils/default_utils.rs b/src/utils/default_utils.rs index 08b196411..3127c94a0 100644 --- a/src/utils/default_utils.rs +++ b/src/utils/default_utils.rs @@ -8,6 +8,6 @@ pub const fn default_resolve_delay_secs() -> u16 { 2 } // Default grace values to accommodate rapid channel changes and seek requests, // helping avoid triggering hard max_connection enforcement. -pub const fn default_grace_period_millis() -> u64 { 2000 } -pub const fn default_grace_period_timeout_secs() -> u64 { 5 } -pub const fn default_connect_timeout_secs() -> u32 { 10 } \ No newline at end of file +pub const fn default_grace_period_millis() -> u64 { 500 } +pub const fn default_grace_period_timeout_secs() -> u64 { 10 } +pub const fn default_connect_timeout_secs() -> u32 { 6 } \ No newline at end of file