diff --git a/backend/src/api/api_utils.rs b/backend/src/api/api_utils.rs index 6a602bd18..f7efcf8e9 100644 --- a/backend/src/api/api_utils.rs +++ b/backend/src/api/api_utils.rs @@ -1,47 +1,59 @@ -use crate::api::endpoints::xtream_api::{get_xtream_player_api_stream_url, ApiStreamContext}; -use crate::api::model::{ - create_active_client_stream, create_channel_unavailable_stream, create_custom_video_stream_response, - create_provider_connections_exhausted_stream, create_provider_stream, get_stream_response_with_headers, AppState, - CustomVideoStreamType, ProviderStreamFactoryOptions, SharedStreamManager, StreamError, ThrottledStream, - UserApiRequest, +use crate::{ + api::{ + endpoints::xtream_api::{get_xtream_player_api_stream_url, ApiStreamContext}, + model::{ + create_active_client_stream, create_channel_unavailable_stream, create_custom_video_stream_response, + create_provider_connections_exhausted_stream, create_provider_stream, get_stream_response_with_headers, + tee_stream, AppState, CustomVideoStreamType, ProviderAllocation, ProviderConfig, + ProviderStreamFactoryOptions, ProviderStreamState, SharedStreamManager, StreamDetails, StreamError, + StreamingStrategy, ThrottledStream, UserApiRequest, UserSession, + }, + }, + auth::Fingerprint, + model::{ConfigInput, ConfigTarget, ProxyUserCredentials}, + tools::lru_cache::LRUResourceCache, + utils::{ + async_file_reader, async_file_writer, create_new_file_for_write, debug_if_enabled, get_file_extension, request, + request::{content_type_from_ext, parse_range, send_with_retry_and_provider}, + trace_if_enabled, + }, + BUILD_TIMESTAMP, }; -use crate::api::model::{tee_stream, UserSession}; -use crate::api::model::{ProviderAllocation, ProviderConfig, ProviderStreamState, StreamDetails, StreamingStrategy}; -use crate::auth::Fingerprint; -use crate::model::ConfigInput; -use crate::model::{ConfigTarget, ProxyUserCredentials}; -use crate::tools::lru_cache::LRUResourceCache; -use crate::utils::request; -use crate::utils::request::{content_type_from_ext, parse_range, send_with_retry_and_provider}; -use crate::utils::{async_file_reader, async_file_writer, create_new_file_for_write, get_file_extension}; -use crate::utils::{debug_if_enabled, trace_if_enabled}; -use crate::BUILD_TIMESTAMP; - use arc_swap::ArcSwapOption; -use axum::body::Body; -use axum::http::{header, HeaderMap, HeaderValue, Response, StatusCode}; -use axum::response::IntoResponse; +use axum::{ + body::Body, + http::{header, HeaderMap, HeaderValue, Response, StatusCode}, + response::IntoResponse, +}; use bytes::Bytes; use chrono::{DateTime, Utc}; use futures::{stream, Stream, StreamExt, TryStreamExt}; use jsonwebtoken::{decode, Algorithm, DecodingKey, Validation}; use log::{debug, error, info, log_enabled, trace, warn}; use serde::Serialize; -use shared::concat_string; -use shared::model::{ - Claims, InputFetchMethod, PlaylistEntry, PlaylistItemType, ProxyType, StreamChannel, TargetType, - UserConnectionPermission, VirtualId, XtreamCluster, +use shared::{ + concat_string, + model::{ + Claims, InputFetchMethod, PlaylistEntry, PlaylistItemType, ProxyType, StreamChannel, TargetType, + UserConnectionPermission, VirtualId, XtreamCluster, + }, + utils::{ + bin_serialize, extract_extension_from_url, human_readable_kbps, replace_url_extension, sanitize_sensitive_info, + trim_slash, Internable, CONTENT_TYPE_CBOR, CONTENT_TYPE_JSON, DASH_EXT, HLS_EXT, + }, +}; +use std::{ + borrow::Cow, + collections::HashMap, + convert::Infallible, + io::SeekFrom, + path::{Path, PathBuf}, + sync::Arc, +}; +use tokio::{ + io::{AsyncReadExt, AsyncSeekExt}, + sync::Mutex, }; -use shared::utils::{bin_serialize, human_readable_kbps, trim_slash, Internable, CONTENT_TYPE_CBOR, CONTENT_TYPE_JSON}; -use shared::utils::{extract_extension_from_url, replace_url_extension, sanitize_sensitive_info, DASH_EXT, HLS_EXT}; -use std::borrow::Cow; -use std::collections::HashMap; -use std::convert::Infallible; -use std::io::SeekFrom; -use std::path::{Path, PathBuf}; -use std::sync::Arc; -use tokio::io::{AsyncReadExt, AsyncSeekExt}; -use tokio::sync::Mutex; use tokio_util::io::ReaderStream; use url::Url; @@ -444,7 +456,8 @@ async fn create_stream_response_details( virtual_id: VirtualId, ) -> Result { let mut streaming_strategy = - resolve_streaming_strategy(app_state, stream_url, fingerprint, input, force_provider, allow_provider_grace).await; + resolve_streaming_strategy(app_state, stream_url, fingerprint, input, force_provider, allow_provider_grace) + .await; let mut grace_period_options = app_state.get_grace_options(); grace_period_options.period_millis = get_grace_period_millis( connection_permission, @@ -915,10 +928,7 @@ pub async fn stream_response( .as_ref() .map(|(h, sc, response_url, cvt)| (h.clone(), *sc, response_url.clone(), *cvt)); let provider_name = stream_details.provider_name.clone(); - let actual_request_url = stream_details - .request_url - .clone() - .unwrap_or_else(|| Arc::::from(stream_url)); + let actual_request_url = stream_details.request_url.clone().unwrap_or_else(|| Arc::::from(stream_url)); debug_if_enabled!( "Provider request mapping: allocated_provider={} actual_request_url={}", @@ -1009,10 +1019,7 @@ pub async fn stream_response( }; if log_enabled!(log::Level::Debug) { if session_url.eq(actual_request_url.as_ref()) { - debug!( - "Streaming stream request from {}", - sanitize_sensitive_info(actual_request_url.as_ref()) - ); + debug!("Streaming stream request from {}", sanitize_sensitive_info(actual_request_url.as_ref())); } else { debug!( "Streaming stream request for {} from {}", diff --git a/backend/src/api/config_file.rs b/backend/src/api/config_file.rs index d5f694a28..002617105 100644 --- a/backend/src/api/config_file.rs +++ b/backend/src/api/config_file.rs @@ -1,15 +1,21 @@ -use crate::api::model::{update_app_state_config, update_app_state_sources, AppState, EventMessage}; -use crate::model::{Config, Mappings, SourcesConfig}; -use crate::utils; -use crate::utils::{ - read_templates, prepare_sources_batch, read_config_file, read_mappings_file_unprepared, - read_mappings_file_with_templates, read_sources_file, read_sources_file_from_path_with_templates, +use crate::{ + api::model::{update_app_state_config, update_app_state_sources, AppState, EventMessage}, + model::{Config, Mappings, SourcesConfig}, + utils, + utils::{ + prepare_sources_batch, read_config_file, read_mappings_file_unprepared, read_mappings_file_with_templates, + read_sources_file, read_sources_file_from_path_with_templates, read_templates, + }, }; use log::{debug, error, info}; -use shared::error::TuliproxError; -use shared::model::{ConfigPaths, ConfigType, PatternTemplate}; -use std::path::{Path, PathBuf}; -use std::sync::Arc; +use shared::{ + error::TuliproxError, + model::{ConfigPaths, ConfigType, PatternTemplate}, +}; +use std::{ + path::{Path, PathBuf}, + sync::Arc, +}; #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub enum ConfigFile { @@ -48,7 +54,9 @@ impl ConfigFile { // ----------------------------------------------------------------- /// Load and merge global templates using the current app state. - fn load_prepared_global_templates(app_state: &Arc) -> Result>, TuliproxError> { + fn load_prepared_global_templates( + app_state: &Arc, + ) -> Result>, TuliproxError> { let paths = app_state.app_config.paths.load(); let config = app_state.app_config.config.load(); Self::load_prepared_global_templates_with_config(&paths, &config) @@ -107,14 +115,9 @@ impl ConfigFile { } } - fn apply_mapping_reload( - app_state: &Arc, - prepared_mapping: Option, - ) { + fn apply_mapping_reload(app_state: &Arc, prepared_mapping: Option) { if let Some(prepared_mapping) = prepared_mapping { - app_state - .app_config - .set_mappings(prepared_mapping.mapping_file_path.as_str(), &prepared_mapping.mappings); + app_state.app_config.set_mappings(prepared_mapping.mapping_file_path.as_str(), &prepared_mapping.mappings); for mapping_file in prepared_mapping.mapping_files { info!("Loaded mapping file {}", mapping_file.display()); } @@ -126,8 +129,7 @@ impl ConfigFile { prepared_templates: Option<&[PatternTemplate]>, ) -> Result<(), TuliproxError> { let paths = app_state.app_config.paths.load(); - let prepared_mapping = - Self::prepare_mapping_reload(paths.mapping_file_path.as_deref(), prepared_templates)?; + let prepared_mapping = Self::prepare_mapping_reload(paths.mapping_file_path.as_deref(), prepared_templates)?; Self::apply_mapping_reload(app_state, prepared_mapping); Ok(()) } @@ -219,27 +221,15 @@ impl ConfigFile { let config_dto = read_config_file(config_file.as_str(), true, true)?; let default_mapping_path = utils::get_default_mappings_path(paths.config_path.as_str()); - let current_mapping_path = paths - .mapping_file_path - .clone() - .unwrap_or_else(|| default_mapping_path.clone()); - let next_mapping_path = config_dto - .mapping_path - .clone() - .filter(|path| !path.trim().is_empty()) - .unwrap_or(default_mapping_path); + let current_mapping_path = paths.mapping_file_path.clone().unwrap_or_else(|| default_mapping_path.clone()); + let next_mapping_path = + config_dto.mapping_path.clone().filter(|path| !path.trim().is_empty()).unwrap_or(default_mapping_path); let mapping_changed = current_mapping_path != next_mapping_path; let default_template_path = utils::get_default_templates_path(paths.config_path.as_str()); - let current_template_path = paths - .template_file_path - .clone() - .unwrap_or_else(|| default_template_path.clone()); - let next_template_path = config_dto - .template_path - .clone() - .filter(|path| !path.trim().is_empty()) - .unwrap_or(default_template_path); + let current_template_path = paths.template_file_path.clone().unwrap_or_else(|| default_template_path.clone()); + let next_template_path = + config_dto.template_path.clone().filter(|path| !path.trim().is_empty()).unwrap_or(default_template_path); let template_changed = current_template_path != next_template_path; let mut config: Config = Config::from(config_dto); @@ -261,8 +251,7 @@ impl ConfigFile { PreparedFollowUp::Sources(prepared) } else if mapping_changed { // Only mapping path changed; templates are the same → load templates once. - let prepared_templates = - Self::load_prepared_global_templates_with_config(&effective_paths, &config)?; + let prepared_templates = Self::load_prepared_global_templates_with_config(&effective_paths, &config)?; let prepared = Self::prepare_mapping_reload( effective_paths.mapping_file_path.as_deref(), prepared_templates.as_deref(), @@ -345,11 +334,8 @@ impl ConfigFile { } ConfigFile::Template | ConfigFile::Sources => { ConfigFile::load_sources(app_state).await?; - let event_type = if matches!(self, ConfigFile::Template) { - ConfigType::Template - } else { - ConfigType::Sources - }; + let event_type = + if matches!(self, ConfigFile::Template) { ConfigType::Template } else { ConfigType::Sources }; app_state.event_manager.send_event(EventMessage::ConfigChange(event_type)); } ConfigFile::Config => { diff --git a/backend/src/api/config_watch.rs b/backend/src/api/config_watch.rs index 375de528a..36988e477 100644 --- a/backend/src/api/config_watch.rs +++ b/backend/src/api/config_watch.rs @@ -1,51 +1,44 @@ -use crate::api::config_file::ConfigFile; -use crate::api::model::{AppState, EventMessage}; -use crate::model::{Config, SourcesConfig}; -use crate::utils; -use crate::utils::is_directory; -use arc_swap::access::Access; -use arc_swap::ArcSwap; +use crate::{ + api::{ + config_file::ConfigFile, + model::{AppState, EventMessage}, + }, + model::{Config, SourcesConfig}, + utils, + utils::is_directory, +}; +use arc_swap::{access::Access, ArcSwap}; use log::{error, info}; -use notify::event::{AccessKind, AccessMode}; -use notify::{recommended_watcher, EventKind, RecursiveMode, Watcher}; -use shared::error::{TuliproxError, TuliproxErrorKind}; -use shared::model::ConfigPaths; -use std::collections::HashMap; -use std::path::{Path, PathBuf}; -use std::sync::Arc; +use notify::{ + event::{AccessKind, AccessMode}, + recommended_watcher, EventKind, RecursiveMode, Watcher, +}; +use shared::{ + error::{TuliproxError, TuliproxErrorKind}, + model::ConfigPaths, +}; +use std::{ + collections::HashMap, + path::{Path, PathBuf}, + sync::Arc, +}; use tokio_util::sync::CancellationToken; #[allow(clippy::too_many_lines)] -fn start_config_watch( - app_state: &Arc, - cancel_token: &CancellationToken, -) -> Result<(), TuliproxError> { +fn start_config_watch(app_state: &Arc, cancel_token: &CancellationToken) -> Result<(), TuliproxError> { // let (tx, mut rx) = mpsc::channel::>(100); let (std_tx, std_rx) = std::sync::mpsc::channel(); let (tx, mut rx) = tokio::sync::mpsc::channel(100); - let paths = - > as Access>::load(&app_state.app_config.paths); - let mapping_file_path = paths - .mapping_file_path - .as_ref() - .map_or_else(String::new, ToString::to_string); - let template_file_path = paths - .template_file_path - .as_ref() - .map_or_else(String::new, ToString::to_string); - let files = get_watch_files( - app_state, - &paths, - mapping_file_path.as_str(), - template_file_path.as_str(), - ); + let paths = > as Access>::load(&app_state.app_config.paths); + let mapping_file_path = paths.mapping_file_path.as_ref().map_or_else(String::new, ToString::to_string); + let template_file_path = paths.template_file_path.as_ref().map_or_else(String::new, ToString::to_string); + let files = get_watch_files(app_state, &paths, mapping_file_path.as_str(), template_file_path.as_str()); // // // Add a path to be watched. All files and directories at that path and // // below will be monitored for changes. let path = Path::new(paths.config_path.as_str()); - let recursive_mode = if (!mapping_file_path.is_empty() - && utils::is_directory(&mapping_file_path)) + let recursive_mode = if (!mapping_file_path.is_empty() && utils::is_directory(&mapping_file_path)) || (!template_file_path.is_empty() && utils::is_directory(&template_file_path)) { RecursiveMode::Recursive @@ -64,16 +57,10 @@ fn start_config_watch( }); let mut watcher = recommended_watcher(std_tx).map_err(|err| { - TuliproxError::new( - TuliproxErrorKind::Info, - format!("Failed to init config file watcher {err}"), - ) + TuliproxError::new(TuliproxErrorKind::Info, format!("Failed to init config file watcher {err}")) })?; watcher.watch(path, recursive_mode).map_err(|err| { - TuliproxError::new( - TuliproxErrorKind::Info, - format!("Failed to start config file watcher {err}"), - ) + TuliproxError::new(TuliproxErrorKind::Info, format!("Failed to start config file watcher {err}")) })?; info!("Watching config file changes {}", path.display()); @@ -92,8 +79,7 @@ fn start_config_watch( let _keep_watcher_alive = watcher; - let mut debounce_timer = - Box::pin(tokio::time::sleep(tokio::time::Duration::from_millis(0))); + let mut debounce_timer = Box::pin(tokio::time::sleep(tokio::time::Duration::from_millis(0))); let mut timer_active = false; let mut pending_configs: HashMap = HashMap::new(); @@ -178,8 +164,7 @@ fn get_watch_files( mapping_file_path: &str, template_file_path: &str, ) -> HashMap { - let sources = - > as Access>::load(&app_state.app_config.sources); + let sources = > as Access>::load(&app_state.app_config.sources); let input_files_paths = sources.get_input_files(); let mut files = HashMap::new(); [ diff --git a/backend/src/api/endpoints/api_playlist_utils.rs b/backend/src/api/endpoints/api_playlist_utils.rs index 78023ef5c..1800a3bda 100644 --- a/backend/src/api/endpoints/api_playlist_utils.rs +++ b/backend/src/api/endpoints/api_playlist_utils.rs @@ -1,23 +1,39 @@ -use crate::model::{AppConfig, ConfigInput, ConfigTarget}; -use crate::utils::{m3u, xtream}; +use crate::{ + api::api_utils::{empty_json_list_response, json_or_bin_response, stream_json_or_bin_response_stream}, + model::{AppConfig, ConfigInput, ConfigTarget}, + repository::{ + iter_raw_m3u_input_playlist, iter_raw_m3u_target_playlist, iter_raw_xtream_input_playlist, + iter_raw_xtream_target_playlist, + }, + utils::{m3u, xtream}, +}; use axum::response::IntoResponse; use log::warn; -use serde_json::{json}; -use shared::model::{InputType, M3uPlaylistItem, PlaylistItemType, TargetType, UiPlaylistItem, XtreamCluster, XtreamPlaylistItem}; +use serde_json::json; +use shared::{ + model::{ + InputType, M3uPlaylistItem, PlaylistItemType, TargetType, UiPlaylistItem, XtreamCluster, XtreamPlaylistItem, + }, + utils::interner_gc, +}; use std::sync::Arc; -use crate::api::api_utils::{empty_json_list_response, json_or_bin_response, stream_json_or_bin_response_stream}; -use shared::utils::interner_gc; -use crate::repository::{iter_raw_m3u_input_playlist, iter_raw_m3u_target_playlist, iter_raw_xtream_input_playlist, iter_raw_xtream_target_playlist}; use tokio_stream::StreamExt; -pub(in crate::api::endpoints) async fn get_playlist_for_target(cfg_target: Option<&ConfigTarget>, cfg: &AppConfig, cluster: XtreamCluster, accept: Option<&str>) -> impl IntoResponse + Send { +pub(in crate::api::endpoints) async fn get_playlist_for_target( + cfg_target: Option<&ConfigTarget>, + cfg: &AppConfig, + cluster: XtreamCluster, + accept: Option<&str>, +) -> impl IntoResponse + Send { if let Some(target) = cfg_target { if target.has_output(TargetType::Xtream) { let Some(channel_iterator) = iter_raw_xtream_target_playlist(cfg, target, cluster).await else { - return empty_json_list_response(); + return empty_json_list_response(); }; let item_filter = if cluster == XtreamCluster::Series { - |pli: &XtreamPlaylistItem| !matches!(pli.item_type, PlaylistItemType::Series | PlaylistItemType::LocalSeries) + |pli: &XtreamPlaylistItem| { + !matches!(pli.item_type, PlaylistItemType::Series | PlaylistItemType::LocalSeries) + } } else { |_pli: &XtreamPlaylistItem| true }; @@ -28,7 +44,9 @@ pub(in crate::api::endpoints) async fn get_playlist_for_target(cfg_target: Optio return empty_json_list_response(); }; let item_filter = if cluster == XtreamCluster::Series { - |pli: &M3uPlaylistItem| !matches!(pli.item_type, PlaylistItemType::Series | PlaylistItemType::LocalSeries) + |pli: &M3uPlaylistItem| { + !matches!(pli.item_type, PlaylistItemType::Series | PlaylistItemType::LocalSeries) + } } else { |_pli: &M3uPlaylistItem| true }; @@ -52,18 +70,22 @@ pub(in crate::api::endpoints) async fn get_playlist_for_target(cfg_target: Optio (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid Arguments"}))).into_response() } - -pub(in crate::api::endpoints) async fn get_playlist_for_input(cfg_input: Option<&Arc>, cfg: &AppConfig, cluster: XtreamCluster, accept: Option<&str>) -> impl IntoResponse + Send { +pub(in crate::api::endpoints) async fn get_playlist_for_input( + cfg_input: Option<&Arc>, + cfg: &AppConfig, + cluster: XtreamCluster, + accept: Option<&str>, +) -> impl IntoResponse + Send { if let Some(input) = cfg_input { if matches!(input.input_type, InputType::Xtream | InputType::XtreamBatch) { let Some(channel_iterator) = iter_raw_xtream_input_playlist(cfg, input, cluster).await else { - return empty_json_list_response(); + return empty_json_list_response(); }; let converted_stream = channel_iterator.map(UiPlaylistItem::from); return stream_json_or_bin_response_stream(accept, converted_stream).into_response(); } else if matches!(input.input_type, InputType::M3u | InputType::M3uBatch) { let Some(channels) = iter_raw_m3u_input_playlist(cfg, input, Some(cluster)).await else { - return empty_json_list_response(); + return empty_json_list_response(); }; let converted_stream = channels.filter_map(|res| match res { Ok(pli) => Some(UiPlaylistItem::from(pli)), @@ -78,30 +100,46 @@ pub(in crate::api::endpoints) async fn get_playlist_for_input(cfg_input: Option< (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid Arguments"}))).into_response() } -pub(in crate::api::endpoints) async fn get_playlist_for_custom_provider(client: &reqwest::Client, cfg_input: Option<&Arc>, app_config: &Arc, cluster: XtreamCluster, accept: Option<&str>) -> impl IntoResponse + Send { +pub(in crate::api::endpoints) async fn get_playlist_for_custom_provider( + client: &reqwest::Client, + cfg_input: Option<&Arc>, + app_config: &Arc, + cluster: XtreamCluster, + accept: Option<&str>, +) -> impl IntoResponse + Send { let cfg = app_config.config.load(); match cfg_input { Some(input) => { - let (result, errors) = - match input.input_type { - InputType::M3u | InputType::M3uBatch => m3u::download_m3u_playlist(app_config, client, &cfg, input).await, - InputType::Xtream | InputType::XtreamBatch => { - let (pl, err, _) = xtream::download_xtream_playlist(app_config, client, input, Some(&[cluster])).await; - (pl, err) - } - InputType::Library => { - return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({ "error": "Library inputs are not supported on this endpoint"}))).into_response(); - } - }; + let (result, errors) = match input.input_type { + InputType::M3u | InputType::M3uBatch => { + m3u::download_m3u_playlist(app_config, client, &cfg, input).await + } + InputType::Xtream | InputType::XtreamBatch => { + let (pl, err, _) = + xtream::download_xtream_playlist(app_config, client, input, Some(&[cluster])).await; + (pl, err) + } + InputType::Library => { + return ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(json!({ "error": "Library inputs are not supported on this endpoint"})), + ) + .into_response(); + } + }; if result.is_empty() { let error_strings: Vec = errors.iter().map(ToString::to_string).collect(); - (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": error_strings.join(", ")}))).into_response() + (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": error_strings.join(", ")}))) + .into_response() } else { - let channels: Vec = result.iter().flat_map(|g| g.channels.iter()).map(UiPlaylistItem::from).collect(); + let channels: Vec = + result.iter().flat_map(|g| g.channels.iter()).map(UiPlaylistItem::from).collect(); interner_gc(); json_or_bin_response(accept, &channels).into_response() } } - None => (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid Arguments"}))).into_response(), + None => { + (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid Arguments"}))).into_response() + } } } diff --git a/backend/src/api/endpoints/custom_video_stream_api.rs b/backend/src/api/endpoints/custom_video_stream_api.rs index 68a34f623..1705d74e6 100644 --- a/backend/src/api/endpoints/custom_video_stream_api.rs +++ b/backend/src/api/endpoints/custom_video_stream_api.rs @@ -1,26 +1,22 @@ +use crate::{ + api::model::{create_custom_video_stream_response, AppState, CustomVideoStreamType}, + auth::Fingerprint, +}; use axum::response::IntoResponse; -use std::str::FromStr; -use std::sync::Arc; -use crate::api::model::{create_custom_video_stream_response, AppState, CustomVideoStreamType}; -use crate::auth::Fingerprint; +use std::{str::FromStr, sync::Arc}; async fn cvs_api( fingerprint: Fingerprint, - axum::extract::Path((username, password, stream_type)): axum::extract::Path<( - String, - String, - String, - )>, + axum::extract::Path((username, password, stream_type)): axum::extract::Path<(String, String, String)>, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse + Send { - let cvs_type = stream_type.strip_suffix(".ts").unwrap_or(&stream_type); let Ok(custom_video_type) = CustomVideoStreamType::from_str(cvs_type) else { return axum::http::StatusCode::NOT_FOUND.into_response(); }; - let Some((user, _target)) = app_state.app_config.get_target_for_user(&username, &password) else { + let Some((user, _target)) = app_state.app_config.get_target_for_user(&username, &password) else { return axum::http::StatusCode::FORBIDDEN.into_response(); }; @@ -28,14 +24,9 @@ async fn cvs_api( return axum::http::StatusCode::FORBIDDEN.into_response(); } - create_custom_video_stream_response( - &app_state, - &fingerprint.addr, - custom_video_type - ).await.into_response() + create_custom_video_stream_response(&app_state, &fingerprint.addr, custom_video_type).await.into_response() } pub fn cvs_api_register() -> axum::Router> { - axum::Router::new() - .route("/cvs/{username}/{password}/{stream_type}", axum::routing::get(cvs_api)) + axum::Router::new().route("/cvs/{username}/{password}/{stream_type}", axum::routing::get(cvs_api)) } diff --git a/backend/src/api/endpoints/download_api.rs b/backend/src/api/endpoints/download_api.rs index c33ca69c8..987a75b3e 100644 --- a/backend/src/api/endpoints/download_api.rs +++ b/backend/src/api/endpoints/download_api.rs @@ -1,89 +1,98 @@ -use crate::api::model::AppState; -use crate::api::model::{DownloadQueue, FileDownload, FileDownloadRequest}; -use crate::model::{AppConfig, VideoDownloadConfig}; -use crate::utils::{async_file_writer, request, IO_BUFFER_SIZE}; -use tokio::sync::RwLock; +use crate::{ + api::model::{AppState, DownloadQueue, FileDownload, FileDownloadRequest}, + model::{AppConfig, VideoDownloadConfig}, + utils::{async_file_writer, request, request::create_client, IO_BUFFER_SIZE}, +}; +use axum::response::IntoResponse; use futures::stream::TryStreamExt; use log::info; use serde_json::{json, Value}; -use tokio::fs::File; -use tokio::io::AsyncWriteExt; -use std::ops::Deref; -use std::sync::Arc; -use tokio::fs; -use axum::response::IntoResponse; -use shared::utils::bytes_to_megabytes; -use shared::error::to_io_error; -use crate::utils::request::create_client; +use shared::{error::to_io_error, utils::bytes_to_megabytes}; +use std::{ops::Deref, sync::Arc}; +use tokio::{fs, fs::File, io::AsyncWriteExt, sync::RwLock}; async fn download_file(active: Arc>>, client: &reqwest::Client) -> Result<(), String> { if let Some(file_download) = active.read().await.as_ref().as_ref() { match client.get(file_download.url.clone()).send().await { - Ok(response) => { - match fs::create_dir_all(&file_download.file_dir).await { - Ok(()) => { - if let Some(file_path_str) = file_download.file_path.to_str() { - info!("Downloading {file_path_str}"); - match File::create(&file_download.file_path).await { - Ok(file) => { - let mut buf_writer = async_file_writer(file); - let mut downloaded: u64 = 0; - let mut stream = response.bytes_stream().map_err(to_io_error); - let mut write_counter = 0; - loop { - match stream.try_next().await { - Ok(item) => { - if let Some(chunk) = item { - match buf_writer.write_all(&chunk).await { - Ok(()) => { - write_counter += chunk.len(); - if write_counter >= IO_BUFFER_SIZE { - buf_writer.flush().await.map_err(|err| err.to_string())?; - write_counter = 0; - } - - downloaded += chunk.len() as u64; - if let Some(lock) = active.write().await.as_mut() { - lock.size = downloaded; - } + Ok(response) => match fs::create_dir_all(&file_download.file_dir).await { + Ok(()) => { + if let Some(file_path_str) = file_download.file_path.to_str() { + info!("Downloading {file_path_str}"); + match File::create(&file_download.file_path).await { + Ok(file) => { + let mut buf_writer = async_file_writer(file); + let mut downloaded: u64 = 0; + let mut stream = response.bytes_stream().map_err(to_io_error); + let mut write_counter = 0; + loop { + match stream.try_next().await { + Ok(item) => { + if let Some(chunk) = item { + match buf_writer.write_all(&chunk).await { + Ok(()) => { + write_counter += chunk.len(); + if write_counter >= IO_BUFFER_SIZE { + buf_writer.flush().await.map_err(|err| err.to_string())?; + write_counter = 0; + } + + downloaded += chunk.len() as u64; + if let Some(lock) = active.write().await.as_mut() { + lock.size = downloaded; } - Err(err) => return Err(format!("Error while writing to file: {file_path_str} {err}")) } - } else { - let megabytes = bytes_to_megabytes(downloaded); - info!("Downloaded {file_path_str}, filesize: {megabytes}MB"); - if let Some(lock) = active.write().await.as_mut() { - lock.size = downloaded; + Err(err) => { + return Err(format!( + "Error while writing to file: {file_path_str} {err}" + )) } - buf_writer.flush().await.map_err(|err| err.to_string())?; - buf_writer.shutdown().await.map_err(|err| err.to_string())?; - return Ok(()); } + } else { + let megabytes = bytes_to_megabytes(downloaded); + info!("Downloaded {file_path_str}, filesize: {megabytes}MB"); + if let Some(lock) = active.write().await.as_mut() { + lock.size = downloaded; + } + buf_writer.flush().await.map_err(|err| err.to_string())?; + buf_writer.shutdown().await.map_err(|err| err.to_string())?; + return Ok(()); } - Err(err) => return Err(format!("Error while writing to file: {file_path_str} {err}")) + } + Err(err) => { + return Err(format!("Error while writing to file: {file_path_str} {err}")) } } } - Err(err) => Err(format!("Error while writing to file: {file_path_str} {err}")) } - } else { - Err("Error file-download file-path unknown".to_string()) + Err(err) => Err(format!("Error while writing to file: {file_path_str} {err}")), } + } else { + Err("Error file-download file-path unknown".to_string()) } - Err(err) => Err(format!("Error while creating directory for file: {} {}", &file_download.file_dir.to_str().unwrap_or("?"), err)) } - } - Err(err) => Err(format!("Error while opening url: {} {}", &file_download.url, err)) + Err(err) => Err(format!( + "Error while creating directory for file: {} {}", + &file_download.file_dir.to_str().unwrap_or("?"), + err + )), + }, + Err(err) => Err(format!("Error while opening url: {} {}", &file_download.url, err)), } } else { Err("No active file download".to_string()) } } -async fn run_download_queue(cfg: &AppConfig, download_cfg: &VideoDownloadConfig, download_queue: &Arc) -> Result<(), String> { +async fn run_download_queue( + cfg: &AppConfig, + download_cfg: &VideoDownloadConfig, + download_queue: &Arc, +) -> Result<(), String> { let next_download = download_queue.as_ref().queue.lock().await.pop_front(); if next_download.is_some() { - { *download_queue.as_ref().active.write().await = next_download; } + { + *download_queue.as_ref().active.write().await = next_download; + } let config = cfg.config.load(); let disabled_headers = cfg.get_disabled_headers(); let headers = request::get_request_headers( @@ -127,7 +136,6 @@ async fn run_download_queue(cfg: &AppConfig, download_cfg: &VideoDownloadConfig, Ok(()) } - macro_rules! download_info { ($file_download:expr) => { json!({"uuid": $file_download.uuid, "filename": $file_download.filename, @@ -139,14 +147,18 @@ macro_rules! download_info { pub async fn queue_download_file( axum::extract::State(app_state): axum::extract::State>, axum::extract::Json(req): axum::extract::Json, -) -> impl axum::response::IntoResponse + Send { +) -> impl axum::response::IntoResponse + Send { let app_config = &*app_state.app_config; let config = app_config.config.load(); if let Some(video_cfg) = config.video.as_ref() { if let Some(download_cfg) = video_cfg.download.as_ref() { if download_cfg.directory.is_empty() { - return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Server config missing video.download.directory configuration"}))).into_response(); + return ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(json!({"error": "Server config missing video.download.directory configuration"})), + ) + .into_response(); } match FileDownload::new(req.url.as_str(), req.filename.as_str(), download_cfg) { Some(file_download) => { @@ -154,30 +166,49 @@ pub async fn queue_download_file( if app_state.downloads.active.read().await.is_none() { match run_download_queue(&app_state.app_config, download_cfg, &app_state.downloads).await { Ok(()) => {} - Err(err) => return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err}))).into_response(), + Err(err) => { + return ( + axum::http::StatusCode::INTERNAL_SERVER_ERROR, + axum::Json(json!({"error": err})), + ) + .into_response() + } } } axum::Json(download_info!(&file_download)).into_response() } - None => (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid Arguments"}))).into_response(), + None => (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid Arguments"}))) + .into_response(), } } else { - (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Server config missing video.download configuration"}))).into_response() + ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(json!({"error": "Server config missing video.download configuration"})), + ) + .into_response() } } else { - (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Server config missing video configuration"}))).into_response() + (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Server config missing video configuration"}))) + .into_response() } } pub async fn download_file_info( axum::extract::State(app_state): axum::extract::State>, -) -> impl axum::response::IntoResponse + Send { - let finished_list: &[Value] = &app_state.downloads.finished.write().await.drain(..) - .map(|fd| download_info!(fd)).collect::>(); +) -> impl axum::response::IntoResponse + Send { + let finished_list: &[Value] = + &app_state.downloads.finished.write().await.drain(..).map(|fd| download_info!(fd)).collect::>(); - (*app_state.downloads.active.read().await).as_ref().map_or_else(|| axum::Json(json!({ - "completed": true, "downloads": finished_list - })), |file_download| axum::Json(json!({ - "completed": false, "downloads": finished_list, "active": download_info!(file_download) - }))) + (*app_state.downloads.active.read().await).as_ref().map_or_else( + || { + axum::Json(json!({ + "completed": true, "downloads": finished_list + })) + }, + |file_download| { + axum::Json(json!({ + "completed": false, "downloads": finished_list, "active": download_info!(file_download) + })) + }, + ) } diff --git a/backend/src/api/endpoints/extract_accept_header.rs b/backend/src/api/endpoints/extract_accept_header.rs index 4b3c6ce32..dcee01b21 100644 --- a/backend/src/api/endpoints/extract_accept_header.rs +++ b/backend/src/api/endpoints/extract_accept_header.rs @@ -1,11 +1,11 @@ -use axum::extract::FromRequestParts; -use axum::http::request::Parts; -use axum::http::StatusCode; +use axum::{ + extract::FromRequestParts, + http::{request::Parts, StatusCode}, +}; #[derive(Debug, PartialEq, Eq, Clone)] pub struct ExtractAcceptHeader(pub Option); - impl FromRequestParts for ExtractAcceptHeader where B: Send + Sync, diff --git a/backend/src/api/endpoints/hdhomerun_api.rs b/backend/src/api/endpoints/hdhomerun_api.rs index 6561b56aa..3d3a96671 100644 --- a/backend/src/api/endpoints/hdhomerun_api.rs +++ b/backend/src/api/endpoints/hdhomerun_api.rs @@ -1,21 +1,24 @@ -use crate::api::api_utils::{try_unwrap_body, internal_server_error}; -use crate::api::model::HdHomerunAppState; -use crate::auth::AuthBasic; -use crate::model::{AppConfig, ConfigInputFlags, ConfigTarget, ProxyUserCredentials}; -use crate::utils::arc_str_serde; -use crate::processing::parser::xtream::get_xtream_url; -use crate::repository::{iter_raw_m3u_target_playlist, M3uPlaylistIterator}; -use crate::repository::XtreamPlaylistIterator; +use crate::{ + api::{ + api_utils::{internal_server_error, try_unwrap_body}, + model::HdHomerunAppState, + }, + auth::AuthBasic, + model::{AppConfig, ConfigInputFlags, ConfigTarget, ProxyUserCredentials}, + processing::parser::xtream::get_xtream_url, + repository::{iter_raw_m3u_target_playlist, M3uPlaylistIterator, XtreamPlaylistIterator}, + utils::arc_str_serde, +}; use axum::response::IntoResponse; use bytes::Bytes; use futures::{stream, Stream, StreamExt}; use log::{error, warn}; use serde::{Deserialize, Serialize}; use serde_json::json; -use shared::model::{ - M3uPlaylistItem, PlaylistItemType, TargetType, XtreamCluster, XtreamPlaylistItem, +use shared::{ + model::{M3uPlaylistItem, PlaylistItemType, TargetType, XtreamCluster, XtreamPlaylistItem}, + utils::concat_path, }; -use shared::utils::{concat_path}; use std::sync::Arc; #[derive(Serialize, Deserialize, Clone)] @@ -84,8 +87,8 @@ impl Device { self.model_name, self.model_number, self.tuner_count, - self.id, // Correct: 8-digit hex device ID - self.udn // Correct: Application/Device UUID for UDN + self.id, // Correct: 8-digit hex device ID + self.udn // Correct: Application/Device UUID for UDN ) } } @@ -96,17 +99,16 @@ fn xtream_item_to_lineup_stream( credentials: Arc, base_url: Option, channels: Option, -) -> impl Stream> +) -> impl Stream> where - I: Stream + Send + Unpin + 'static, + I: Stream + Send + Unpin + 'static, { match channels { Some(chans) => { let mapped = chans.map(move |(item, has_next)| { let input = cfg.get_input_by_name(&item.input_name); - let (live_stream_use_prefix, live_stream_without_extension) = input - .as_ref() - .map_or((true, false), |i| { + let (live_stream_use_prefix, live_stream_without_extension) = + input.as_ref().map_or((true, false), |i| { ( i.has_flag(ConfigInputFlags::XtreamLiveStreamUsePrefix), i.has_flag(ConfigInputFlags::XtreamLiveStreamWithoutExtension), @@ -124,7 +126,8 @@ where container_extension.as_deref(), live_stream_use_prefix, live_stream_without_extension, - ).into(), + ) + .into(), }; let lineup = Lineup { @@ -138,7 +141,7 @@ where content.push(','); } Ok(Bytes::from(content)) - }, + } Err(_) => Ok(Bytes::from("")), } }); @@ -148,9 +151,9 @@ where } } -fn m3u_item_to_lineup_stream(channels: Option) -> impl Stream> +fn m3u_item_to_lineup_stream(channels: Option) -> impl Stream> where - I: Stream + Send + Unpin + 'static, + I: Stream + Send + Unpin + 'static, { match channels { Some(chans) => { @@ -158,11 +161,7 @@ where let lineup = Lineup { guide_number: item.epg_channel_id.clone().unwrap_or(item.name.clone()), guide_name: item.title.clone(), - url: if item.t_stream_url.is_empty() { - item.url.clone() - } else { - item.t_stream_url.clone() - } + url: if item.t_stream_url.is_empty() { item.url.clone() } else { item.t_stream_url.clone() }, }; match serde_json::to_string(&lineup) { Ok(mut content) => { @@ -170,7 +169,7 @@ where content.push(','); } Ok(Bytes::from(content)) - }, + } Err(_) => Ok(Bytes::from("")), } }); @@ -181,20 +180,10 @@ where } fn create_device(app_state: &Arc) -> Option { - if let Some(credentials) = app_state - .app_state - .app_config - .get_user_credentials(&app_state.device.t_username) - { - let server_info = app_state - .app_state - .app_config - .get_user_server_info(&credentials); + if let Some(credentials) = app_state.app_state.app_config.get_user_credentials(&app_state.device.t_username) { + let server_info = app_state.app_state.app_config.get_user_server_info(&credentials); let device = &app_state.device; - let device_url = format!( - "{}://{}:{}", - server_info.protocol, server_info.host, device.port - ); + let device_url = format!("{}://{}:{}", server_info.protocol, server_info.host, device.port); Some(Device { friendly_name: device.friendly_name.clone(), @@ -256,9 +245,7 @@ async fn discover_json( async fn lineup_status( axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse { - let current_state = app_state - .hd_scan_state - .load(std::sync::atomic::Ordering::Acquire); + let current_state = app_state.hd_scan_state.load(std::sync::atomic::Ordering::Acquire); if current_state < 0 { axum::Json(json!({ "ScanInProgress": 0, @@ -266,32 +253,31 @@ async fn lineup_status( "Source": "Cable", "SourceList": ["Cable"], })) - .into_response() + .into_response() } else { let new_state = current_state.saturating_add(20); let final_state = if new_state > 100 { 100 } else { new_state }; let cfg = Arc::clone(&app_state.app_state.app_config); - let num_of_channels = if let Some((user, target)) = - cfg.get_target_for_username(&app_state.device.t_username) - { + let num_of_channels = if let Some((user, target)) = cfg.get_target_for_username(&app_state.device.t_username) { if target.has_output(TargetType::M3u) { - if let Some(iter) = iter_raw_m3u_target_playlist(&cfg, &target, None).await - { + if let Some(iter) = iter_raw_m3u_target_playlist(&cfg, &target, None).await { iter.filter_map(|res| async move { res.ok() }).count().await } else { 0 } } else if target.has_output(TargetType::Xtream) { let credentials = Arc::new(user); - let live = match XtreamPlaylistIterator::new(XtreamCluster::Live, &cfg, &target, None, &credentials).await { - Ok(stream) => stream.count().await, - Err(_) => 0, - }; - let vod = match XtreamPlaylistIterator::new(XtreamCluster::Video, &cfg, &target, None, &credentials).await { - Ok(stream) => stream.count().await, - Err(_) => 0, - }; + let live = + match XtreamPlaylistIterator::new(XtreamCluster::Live, &cfg, &target, None, &credentials).await { + Ok(stream) => stream.count().await, + Err(_) => 0, + }; + let vod = + match XtreamPlaylistIterator::new(XtreamCluster::Video, &cfg, &target, None, &credentials).await { + Ok(stream) => stream.count().await, + Err(_) => 0, + }; live + vod } else { 0 @@ -301,13 +287,9 @@ async fn lineup_status( }; if final_state >= 100 { - app_state - .hd_scan_state - .store(-1, std::sync::atomic::Ordering::Release); + app_state.hd_scan_state.store(-1, std::sync::atomic::Ordering::Release); } else { - app_state - .hd_scan_state - .store(final_state, std::sync::atomic::Ordering::Release); + app_state.hd_scan_state.store(final_state, std::sync::atomic::Ordering::Release); } let found = (num_of_channels * usize::try_from(final_state).unwrap_or(1)) / 100; axum::Json(json!({ @@ -315,7 +297,7 @@ async fn lineup_status( "Progress": final_state, "Found": found, })) - .into_response() + .into_response() } } @@ -330,15 +312,11 @@ async fn lineup_post( ) -> impl IntoResponse { match query.scan.as_str() { "start" => { - app_state - .hd_scan_state - .store(0, std::sync::atomic::Ordering::Release); + app_state.hd_scan_state.store(0, std::sync::atomic::Ordering::Release); axum::http::StatusCode::OK.into_response() } "abort" => { - app_state - .hd_scan_state - .store(-1, std::sync::atomic::Ordering::Release); + app_state.hd_scan_state.store(-1, std::sync::atomic::Ordering::Release); axum::http::StatusCode::OK.into_response() } _ => axum::http::StatusCode::BAD_REQUEST.into_response(), @@ -351,33 +329,22 @@ async fn lineup( credentials: &Arc, target: &ConfigTarget, ) -> impl IntoResponse { - let use_output = target - .get_hdhomerun_output() - .as_ref() - .and_then(|o| o.use_output); + let use_output = target.get_hdhomerun_output().as_ref().and_then(|o| o.use_output); let use_all = use_output.is_none(); let use_m3u = use_output.as_ref() == Some(&TargetType::M3u); let use_xtream = use_output.as_ref() == Some(&TargetType::Xtream); if (use_all || use_m3u) && target.has_output(TargetType::M3u) { - let iterator = M3uPlaylistIterator::new(cfg, target, credentials) - .await - .ok(); + let iterator = M3uPlaylistIterator::new(cfg, target, credentials).await.ok(); let stream = m3u_item_to_lineup_stream(iterator); let body_stream = stream::once(async { Ok(Bytes::from("[")) }) .chain(stream) .chain(stream::once(async { Ok(Bytes::from("]")) })); return try_unwrap_body!(axum::response::Response::builder() .status(axum::http::StatusCode::OK) - .header( - axum::http::header::CONTENT_TYPE, - mime::APPLICATION_JSON.to_string() - ) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) .body(axum::body::Body::from_stream(body_stream))); } else if (use_all || use_xtream) && target.has_output(TargetType::Xtream) { - let server_info = app_state - .app_state - .app_config - .get_user_server_info(credentials); + let server_info = app_state.app_state.app_config.get_user_server_info(credentials); let base_url = server_info.get_base_url(); let base_url_live = if credentials.proxy.is_redirect(PlaylistItemType::Live) @@ -395,14 +362,8 @@ async fn lineup( Some(base_url) }; - let live_channels = - XtreamPlaylistIterator::new(XtreamCluster::Live, cfg, target, None, credentials) - .await - .ok(); - let vod_channels = - XtreamPlaylistIterator::new(XtreamCluster::Video, cfg, target, None, credentials) - .await - .ok(); + let live_channels = XtreamPlaylistIterator::new(XtreamCluster::Live, cfg, target, None, credentials).await.ok(); + let vod_channels = XtreamPlaylistIterator::new(XtreamCluster::Video, cfg, target, None, credentials).await.ok(); let live_stream = xtream_item_to_lineup_stream( Arc::clone(cfg), XtreamCluster::Live, @@ -420,7 +381,8 @@ async fn lineup( let mut live_stream_peek = Box::pin(live_stream.peekable()); let mut vod_stream_peek = Box::pin(vod_stream.peekable()); - let both_non_empty = live_stream_peek.as_mut().peek().await.is_some() && vod_stream_peek.as_mut().peek().await.is_some(); + let both_non_empty = + live_stream_peek.as_mut().peek().await.is_some() && vod_stream_peek.as_mut().peek().await.is_some(); let comma_stream = if both_non_empty { stream::once(async { Ok(Bytes::from(",")) }).left_stream() } else { @@ -435,10 +397,7 @@ async fn lineup( return try_unwrap_body!(axum::response::Response::builder() .status(axum::http::StatusCode::OK) - .header( - axum::http::header::CONTENT_TYPE, - mime::APPLICATION_JSON.to_string() - ) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) .body(axum::body::Body::from_stream(body_stream))); } axum::http::StatusCode::NOT_FOUND.into_response() @@ -449,15 +408,12 @@ async fn auth_lineup_json( axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse { let cfg = Arc::clone(&app_state.app_state.app_config); - if let Some((credentials, target)) = cfg.get_target_for_username(&app_state.device.t_username) - { + if let Some((credentials, target)) = cfg.get_target_for_username(&app_state.device.t_username) { if !username.eq(&credentials.username) || !password.eq(&credentials.password) { return axum::http::StatusCode::UNAUTHORIZED.into_response(); } let user_credentials = Arc::new(credentials); - return lineup(&app_state, &cfg, &user_credentials, &target) - .await - .into_response(); + return lineup(&app_state, &cfg, &user_credentials, &target).await.into_response(); } axum::http::StatusCode::NOT_FOUND.into_response() } @@ -466,12 +422,9 @@ async fn lineup_json( axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse { let cfg = Arc::clone(&app_state.app_state.app_config); - if let Some((credentials, target)) = cfg.get_target_for_username(&app_state.device.t_username) - { + if let Some((credentials, target)) = cfg.get_target_for_username(&app_state.device.t_username) { let user_credentials = Arc::new(credentials); - return lineup(&app_state, &cfg, &user_credentials, &target) - .await - .into_response(); + return lineup(&app_state, &cfg, &user_credentials, &target).await.into_response(); } axum::http::StatusCode::NOT_FOUND.into_response() } @@ -492,16 +445,9 @@ pub fn hdhr_api_register(basic_auth: bool) -> axum::Router (url, None), } } - None => { - match app_state - .active_provider - .get_next_provider(&input.name) - .await - { - Some(provider_cfg) => { - let stream_url = get_stream_alternative_url(&url, input, &provider_cfg); - debug_if_enabled!( + None => match app_state.active_provider.get_next_provider(&input.name).await { + Some(provider_cfg) => { + let stream_url = get_stream_alternative_url(&url, input, &provider_cfg); + debug_if_enabled!( "API endpoint [HLS] create_session_fingerprint user={} virtual_id={virtual_id} provider={} stream_url={}", sanitize_sensitive_info(&user.username), provider_cfg.name, sanitize_sensitive_info(&stream_url) ); - let user_session_token = create_session_fingerprint(fingerprint, &user.username, virtual_id); - let session_token = app_state.active_users.create_user_session( + let user_session_token = create_session_fingerprint(fingerprint, &user.username, virtual_id); + let session_token = app_state + .active_users + .create_user_session( user, &user_session_token, virtual_id, @@ -169,37 +163,46 @@ pub(in crate::api) async fn handle_hls_stream_request( &stream_url, &fingerprint.addr, connection_permission, - ).await; - (stream_url, Some(session_token)) - } - None => (url, None), + ) + .await; + (stream_url, Some(session_token)) } - } + None => (url, None), + }, }; - // Don't forward Range on playlist fetch; segments use original headers in provider path let filter_header: HeaderFilter = Some(Box::new(|name: &str| !name.eq_ignore_ascii_case("range"))); let forwarded = get_headers_from_request(req_headers, &filter_header); let disabled_headers = app_state.get_disabled_headers(); let default_user_agent = app_state.app_config.config.load().default_user_agent.clone(); - let headers = request::get_request_headers( - None, - Some(&forwarded), - disabled_headers.as_ref(), - default_user_agent.as_deref(), - ); + let headers = + request::get_request_headers(None, Some(&forwarded), disabled_headers.as_ref(), default_user_agent.as_deref()); let input_source = InputSource::from(input).with_url(request_url); - match request::download_text_content( - &app_state.app_config, - &app_state.http_client.load(), - &input_source, - Some(&headers), - None, - false, - ) + let use_manual_redirects = app_state.should_use_manual_redirects(); + let download_result = if use_manual_redirects { + request::download_text_content_with_manual_redirects( + &app_state.app_config, + &app_state.http_client_no_redirect.load(), + &input_source, + Some(&headers), + None, + false, + MAX_MANUAL_REDIRECTS, + ) .await - { + } else { + request::download_text_content( + &app_state.app_config, + &app_state.http_client.load(), + &input_source, + Some(&headers), + None, + false, + ) + .await + }; + match download_result { Ok((content, response_url)) => { let rewrite_hls_props = RewriteHlsProps { secret: &app_state.app_config.encrypt_secret, @@ -214,10 +217,7 @@ pub(in crate::api) async fn handle_hls_stream_request( hls_response(hls_content).into_response() } Err(err) => { - error!( - "Failed to download m3u8 {}", - sanitize_sensitive_info(err.to_string().as_str()) - ); + error!("Failed to download m3u8 {}", sanitize_sensitive_info(err.to_string().as_str())); let custom_stream_response = app_state.app_config.custom_stream_response.load(); if custom_stream_response.as_ref().and_then(|c| c.channel_unavailable.as_ref()).is_some() { @@ -226,7 +226,8 @@ pub(in crate::api) async fn handle_hls_stream_request( &server_info.get_base_url(), user.username, user.password, - CustomVideoStreamType::ChannelUnavailable); + CustomVideoStreamType::ChannelUnavailable + ); let playlist = PLAYLIST_TEMPLATE.replace("{url}", &url); hls_response(playlist).into_response() @@ -237,7 +238,11 @@ pub(in crate::api) async fn handle_hls_stream_request( } } -async fn get_stream_channel(app_state: &Arc, target: &Arc, virtual_id: u32) -> Option { +async fn get_stream_channel( + app_state: &Arc, + target: &Arc, + virtual_id: u32, +) -> Option { if target.has_output(TargetType::Xtream) { if let Ok(pli) = xtream_get_item_for_stream_id(virtual_id, app_state, target, None).await { return Some(pli.to_stream_channel(target.id)); @@ -258,7 +263,7 @@ async fn resolve_stream_channel( Some(mut channel) => { channel.url = hls_url.clone(); channel - }, + } None => StreamChannel { target_id: target.id, virtual_id, @@ -284,9 +289,7 @@ async fn hls_api_stream( axum::extract::State(app_state): axum::extract::State>, ) -> impl axum::response::IntoResponse + Send { let (user, target) = try_option_bad_request!( - app_state - .app_config - .get_target_for_user(¶ms.username, ¶ms.password), + app_state.app_config.get_target_for_user(¶ms.username, ¶ms.password), false, format!("Could not find any user for hls stream {}", params.username) ); @@ -295,8 +298,9 @@ async fn hls_api_stream( &app_state, &fingerprint.addr, CustomVideoStreamType::UserAccountExpired, - ).await - .into_response(); + ) + .await + .into_response(); } let target_name = &target.name; @@ -309,37 +313,35 @@ async fn hls_api_stream( debug_if_enabled!("ID chain for hls endpoint: request_stream_id={} -> virtual_id={virtual_id}", params.stream_id); let user_session_token = create_session_fingerprint(&fingerprint, &user.username, virtual_id); - let mut user_session = app_state - .active_users - .get_and_update_user_session(&user.username, &user_session_token).await; + let mut user_session = + app_state.active_users.get_and_update_user_session(&user.username, &user_session_token).await; if let Some(session) = &mut user_session { if session.permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response( - &app_state, &fingerprint.addr, + &app_state, + &fingerprint.addr, CustomVideoStreamType::UserConnectionsExhausted, - ).await.into_response(); - } - - if app_state - .active_provider - .is_over_limit(&session.provider) + ) .await - { - return create_custom_video_stream_response( - &app_state, &fingerprint.addr, - CustomVideoStreamType::ProviderConnectionsExhausted, - ).await - .into_response(); + .into_response(); } - let hls_url = match get_hls_session_token_and_url_from_token( - &app_state.app_config.encrypt_secret, - ¶ms.token, - ) { - Some((Some(session_token), hls_url)) if session.token.eq(&session_token) => hls_url, - _ => return axum::http::StatusCode::BAD_REQUEST.into_response(), - }; + if app_state.active_provider.is_over_limit(&session.provider).await { + return create_custom_video_stream_response( + &app_state, + &fingerprint.addr, + CustomVideoStreamType::ProviderConnectionsExhausted, + ) + .await + .into_response(); + } + + let hls_url = + match get_hls_session_token_and_url_from_token(&app_state.app_config.encrypt_secret, ¶ms.token) { + Some((Some(session_token), hls_url)) if session.token.eq(&session_token) => hls_url, + _ => return axum::http::StatusCode::BAD_REQUEST.into_response(), + }; let hls_url = hls_url.intern(); session.stream_url = hls_url.clone(); if session.virtual_id == virtual_id { @@ -355,8 +357,8 @@ async fn hls_api_stream( &input, &user, ) - .await - .into_response(); + .await + .into_response(); } } else { return axum::http::StatusCode::BAD_REQUEST.into_response(); @@ -365,9 +367,12 @@ async fn hls_api_stream( let connection_permission = user.connection_permission(&app_state).await; if connection_permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response( - &app_state, &fingerprint.addr, + &app_state, + &fingerprint.addr, CustomVideoStreamType::UserConnectionsExhausted, - ).await.into_response(); + ) + .await + .into_response(); } if is_hls_url(&session.stream_url) { @@ -387,8 +392,7 @@ async fn hls_api_stream( } if is_file_url(&session.stream_url) { - let stream_channel = - resolve_stream_channel(&app_state, &target, virtual_id, &hls_url).await; + let stream_channel = resolve_stream_channel(&app_state, &target, virtual_id, &hls_url).await; return local_stream_response( &fingerprint, &app_state, @@ -405,15 +409,7 @@ async fn hls_api_stream( } let stream_channel = resolve_stream_channel(&app_state, &target, virtual_id, &hls_url).await; - force_provider_stream_response( - &fingerprint, - &app_state, - session, - stream_channel, - &req_headers, - &input, - &user, - ) + force_provider_stream_response(&fingerprint, &app_state, session, stream_channel, &req_headers, &input, &user) .await .into_response() } else { @@ -421,12 +417,9 @@ async fn hls_api_stream( } } - pub fn hls_api_register() -> axum::Router> { - axum::Router::new().route( - "/hls/{username}/{password}/{input_id}/{stream_id}/{token}", - axum::routing::get(hls_api_stream), - ) + axum::Router::new() + .route("/hls/{username}/{password}/{input_id}/{stream_id}/{token}", axum::routing::get(hls_api_stream)) //cfg.service(web::resource("/hls/{token}/{stream}").route(web::get().to(xtream_player_api_hls_stream))); //cfg.service(web::resource("/play/{token}/{type}").route(web::get().to(xtream_player_api_play_stream))); } diff --git a/backend/src/api/endpoints/library_api.rs b/backend/src/api/endpoints/library_api.rs index 4336a2adc..33973310b 100644 --- a/backend/src/api/endpoints/library_api.rs +++ b/backend/src/api/endpoints/library_api.rs @@ -1,11 +1,15 @@ -use crate::api::model::{AppState, EventMessage}; +use crate::{ + api::{ + library_scan::spawn_library_scan, + model::{AppState, EventMessage}, + }, + library::LibraryProcessor, +}; use axum::response::IntoResponse; use log::{debug, warn}; -use std::sync::Arc; use serde_json::json; use shared::model::{LibraryScanRequest, LibraryScanSummary, LibraryScanSummaryStatus, LibraryStatus}; -use crate::api::library_scan::spawn_library_scan; -use crate::library::LibraryProcessor; +use std::sync::Arc; // Triggers a library scan async fn scan_library( @@ -22,7 +26,11 @@ async fn scan_library( result: None, }; let _ = app_state.event_manager.send_event(EventMessage::LibraryScanProgress(response)); - return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Library update already in progress.".to_string()}))).into_response(); + return ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(json!({"error": "Library update already in progress.".to_string()})), + ) + .into_response(); }; // Check if Library is enabled @@ -37,32 +45,26 @@ async fn scan_library( result: None, }; let _ = app_state.event_manager.send_event(EventMessage::LibraryScanProgress(response)); - return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Library is not enabled".to_string()}))).into_response(); + return ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(json!({"error": "Library is not enabled".to_string()})), + ) + .into_response(); } } }; let client = app_state.http_client.load_full().as_ref().clone(); let event_manager = Arc::clone(&app_state.event_manager); - spawn_library_scan( - event_manager, - lib_config, - metadata_update_config, - client, - request.force_rescan, - "", - permit, - ); + spawn_library_scan(event_manager, lib_config, metadata_update_config, client, request.force_rescan, "", permit); axum::http::StatusCode::ACCEPTED.into_response() - } /// Gets Library status async fn get_library_status( axum::extract::State(app_state): axum::extract::State>, ) -> axum::response::Response { - let config_snapshot = app_state.app_config.config.load(); if let Some(config) = config_snapshot.library.as_ref() { if config.enabled { @@ -71,14 +73,8 @@ async fn get_library_status( let processor = LibraryProcessor::new(config.clone(), config_snapshot.metadata_update.as_ref(), client); let entries = processor.get_all_entries().await; - let movies = entries - .iter() - .filter(|e| e.metadata.is_movie()) - .count(); - let series = entries - .iter() - .filter(|e| e.metadata.is_series()) - .count(); + let movies = entries.iter().filter(|e| e.metadata.is_movie()).count(); + let series = entries.iter().filter(|e| e.metadata.is_series()).count(); let response = LibraryStatus { enabled: true, @@ -94,10 +90,8 @@ async fn get_library_status( let response = LibraryStatus::default(); axum::Json(response).into_response() - } - /// Registers Library API routes pub fn library_api_register(router: axum::Router>) -> axum::Router> { router diff --git a/backend/src/api/endpoints/m3u_api.rs b/backend/src/api/endpoints/m3u_api.rs index 139dbf788..1273b848e 100644 --- a/backend/src/api/endpoints/m3u_api.rs +++ b/backend/src/api/endpoints/m3u_api.rs @@ -1,24 +1,29 @@ -use crate::api::api_utils::{create_session_fingerprint, local_stream_response, try_unwrap_body}; -use crate::api::api_utils::{ - force_provider_stream_response, get_user_target, get_user_target_by_credentials, - is_seek_request, redirect, redirect_response, resource_response, separate_number_and_remainder, - stream_response, try_result_not_found, try_option_bad_request, try_result_bad_request, RedirectParams, +use crate::{ + api::{ + api_utils::{ + create_session_fingerprint, force_provider_stream_response, get_user_target, + get_user_target_by_credentials, is_seek_request, local_stream_response, redirect, redirect_response, + resource_response, separate_number_and_remainder, stream_response, try_option_bad_request, + try_result_bad_request, try_result_not_found, try_unwrap_body, RedirectParams, + }, + endpoints::{ + hls_api::handle_hls_stream_request, + xtream_api::{ApiStreamContext, ApiStreamRequest}, + }, + model::{create_custom_video_stream_response, AppState, CustomVideoStreamType, UserApiRequest}, + }, + auth::Fingerprint, + repository::{m3u_get_item_for_stream_id, m3u_load_rewrite_playlist, storage_const}, + utils::debug_if_enabled, }; -use crate::api::endpoints::hls_api::handle_hls_stream_request; -use crate::api::endpoints::xtream_api::{ApiStreamContext, ApiStreamRequest}; -use crate::api::model::AppState; -use crate::api::model::UserApiRequest; -use crate::api::model::{create_custom_video_stream_response, CustomVideoStreamType}; -use crate::auth::Fingerprint; -use crate::repository::{m3u_get_item_for_stream_id, m3u_load_rewrite_playlist}; -use crate::repository::storage_const; -use crate::utils::debug_if_enabled; use axum::response::IntoResponse; use bytes::Bytes; use futures::StreamExt; use log::{debug, error}; -use shared::model::{FieldGetAccessor, PlaylistEntry, PlaylistItemType, TargetType, UserConnectionPermission, XtreamCluster}; -use shared::utils::{concat_path, extract_extension_from_url, sanitize_sensitive_info, HLS_EXT}; +use shared::{ + model::{FieldGetAccessor, PlaylistEntry, PlaylistItemType, TargetType, UserConnectionPermission, XtreamCluster}, + utils::{concat_path, extract_extension_from_url, sanitize_sensitive_info, HLS_EXT}, +}; use std::sync::Arc; async fn m3u_api(api_req: &UserApiRequest, app_state: &AppState) -> impl IntoResponse + Send { @@ -34,15 +39,9 @@ async fn m3u_api(api_req: &UserApiRequest, app_state: &AppState) -> impl IntoRes let mut builder = axum::response::Response::builder() .status(axum::http::StatusCode::OK) - .header( - axum::http::header::CONTENT_TYPE, - mime::TEXT_PLAIN_UTF_8.to_string(), - ); + .header(axum::http::header::CONTENT_TYPE, mime::TEXT_PLAIN_UTF_8.to_string()); if api_req.content_type == "m3u_plus" { - builder = builder.header( - "Content-Disposition", - "attachment; filename=\"playlist.m3u\"", - ); + builder = builder.header("Content-Disposition", "attachment; filename=\"playlist.m3u\""); } try_unwrap_body!(builder.body(axum::body::Body::from_stream(content_stream))) } @@ -80,26 +79,20 @@ async fn m3u_api_stream( // _addr: &std::net::SocketAddr, ) -> impl IntoResponse + Send { let (user, target) = try_option_bad_request!( - get_user_target_by_credentials( - stream_req.username, - stream_req.password, - api_req, - app_state - ), + get_user_target_by_credentials(stream_req.username, stream_req.password, api_req, app_state), false, - format!( - "Could not find any user for m3u stream {}", - stream_req.username - ) + format!("Could not find any user for m3u stream {}", stream_req.username) ); - let _guard = app_state.app_config.file_locks.write_lock_str(&user.username).await; + let _guard = app_state.app_config.file_locks.write_lock_str(&user.username).await; if user.permission_denied(app_state) { return create_custom_video_stream_response( - app_state, &fingerprint.addr, + app_state, + &fingerprint.addr, CustomVideoStreamType::UserAccountExpired, - ).await + ) + .await .into_response(); } @@ -123,11 +116,9 @@ async fn m3u_api_stream( } let input = try_option_bad_request!( - app_state - .app_config - .get_input_by_name(&pli.input_name), - true, - format!("Can't find input {} for target {target_name}, stream_id {virtual_id}", pli.input_name) + app_state.app_config.get_input_by_name(&pli.input_name), + true, + format!("Can't find input {} for target {target_name}, stream_id {virtual_id}", pli.input_name) ); if pli.item_type.is_local() { @@ -142,37 +133,37 @@ async fn m3u_api_stream( &user, connection_permission, true, - ).await.into_response(); + ) + .await + .into_response(); } let cluster = XtreamCluster::try_from(pli.item_type).unwrap_or(XtreamCluster::Live); - + debug_if_enabled!( "ID chain for m3u endpoint: request_stream_id={} -> action_stream_id={action_stream_id} -> req_virtual_id={req_virtual_id} -> virtual_id={virtual_id}", stream_req.stream_id); let session_key = create_session_fingerprint(fingerprint, &user.username, virtual_id); - let user_session = app_state - .active_users - .get_and_update_user_session(&user.username, &session_key).await; + let user_session = app_state.active_users.get_and_update_user_session(&user.username, &session_key).await; let session_url = if let Some(session) = &user_session { if session.permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response( - app_state, &fingerprint.addr, + app_state, + &fingerprint.addr, CustomVideoStreamType::UserConnectionsExhausted, - ).await + ) + .await .into_response(); } - if app_state - .active_provider - .is_over_limit(&session.provider) - .await - { + if app_state.active_provider.is_over_limit(&session.provider).await { return create_custom_video_stream_response( - app_state, &fingerprint.addr, + app_state, + &fingerprint.addr, CustomVideoStreamType::ProviderConnectionsExhausted, - ).await + ) + .await .into_response(); } if session.virtual_id == virtual_id && is_seek_request(cluster, req_headers).await { @@ -197,9 +188,11 @@ async fn m3u_api_stream( let connection_permission = user.connection_permission(app_state).await; if connection_permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response( - app_state, &fingerprint.addr, + app_state, + &fingerprint.addr, CustomVideoStreamType::UserConnectionsExhausted, - ).await + ) + .await .into_response(); } @@ -222,9 +215,7 @@ async fn m3u_api_stream( return response.into_response(); } - let extension = stream_ext.unwrap_or_else(|| { - extract_extension_from_url(&pli.url).unwrap_or_default() - }); + let extension = stream_ext.unwrap_or_else(|| extract_extension_from_url(&pli.url).unwrap_or_default()); let is_hls_request = pli.item_type == PlaylistItemType::LiveHls || pli.item_type == PlaylistItemType::LiveDash @@ -276,9 +267,7 @@ async fn m3u_api_resource( let Ok(m3u_stream_id) = stream_id.parse::() else { return axum::http::StatusCode::BAD_REQUEST.into_response(); }; - let Some((user, target)) = - get_user_target_by_credentials(&username, &password, &api_req, &app_state) - else { + let Some((user, target)) = get_user_target_by_credentials(&username, &password, &api_req, &app_state) else { return axum::http::StatusCode::BAD_REQUEST.into_response(); }; if user.permission_denied(&app_state) { @@ -290,34 +279,23 @@ async fn m3u_api_resource( debug!("Target has no m3u playlist {target_name}"); return axum::http::StatusCode::BAD_REQUEST.into_response(); } - let m3u_item = - match m3u_get_item_for_stream_id(m3u_stream_id, &app_state, &target).await { - Ok(item) => item, - Err(err) => { - error!( - "Failed to get m3u url: {}", - sanitize_sensitive_info(err.to_string().as_str()) - ); - return axum::http::StatusCode::NOT_FOUND.into_response(); - } - }; + let m3u_item = match m3u_get_item_for_stream_id(m3u_stream_id, &app_state, &target).await { + Ok(item) => item, + Err(err) => { + error!("Failed to get m3u url: {}", sanitize_sensitive_info(err.to_string().as_str())); + return axum::http::StatusCode::NOT_FOUND.into_response(); + } + }; let stream_url = m3u_item.get_field(resource.as_str()); match stream_url { None => axum::http::StatusCode::NOT_FOUND.into_response(), Some(url) => { - if user.proxy.is_redirect(m3u_item.item_type) - || target.is_force_redirect(m3u_item.item_type) - { - debug!( - "Redirecting stream request to {}", - sanitize_sensitive_info(&url) - ); + if user.proxy.is_redirect(m3u_item.item_type) || target.is_force_redirect(m3u_item.item_type) { + debug!("Redirecting stream request to {}", sanitize_sensitive_info(&url)); redirect(&url).into_response() } else { - resource_response(&app_state, &url, &req_headers, None) - .await - .into_response() + resource_response(&app_state, &url, &req_headers, None).await.into_response() } } } @@ -329,11 +307,7 @@ macro_rules! create_m3u_api_stream { fingerprint: Fingerprint, req_headers: axum::http::HeaderMap, axum::extract::Query(api_req): axum::extract::Query, - axum::extract::Path((username, password, stream_id)): axum::extract::Path<( - String, - String, - String, - )>, + axum::extract::Path((username, password, stream_id)): axum::extract::Path<(String, String, String)>, axum::extract::State(app_state): axum::extract::State>, // axum::extract::ConnectInfo(addr): axum::extract::ConnectInfo, ) -> impl IntoResponse + Send { @@ -383,26 +357,14 @@ pub fn m3u_api_register() -> axum::Router> { router, [ (storage_const::M3U_STREAM_PATH, m3u_api_live_stream_alt), - ( - concat_path(storage_const::M3U_STREAM_PATH, "live"), - m3u_api_live_stream - ), - ( - concat_path(storage_const::M3U_STREAM_PATH, "movie"), - m3u_api_movie_stream - ), - ( - concat_path(storage_const::M3U_STREAM_PATH, "series"), - m3u_api_series_stream - ) + (concat_path(storage_const::M3U_STREAM_PATH, "live"), m3u_api_live_stream), + (concat_path(storage_const::M3U_STREAM_PATH, "movie"), m3u_api_movie_stream), + (concat_path(storage_const::M3U_STREAM_PATH, "series"), m3u_api_series_stream) ] ); router.route( - &format!( - "/{}/{{username}}/{{password}}/{{stream_id}}/{{resource}}", - storage_const::M3U_RESOURCE_PATH - ), + &format!("/{}/{{username}}/{{password}}/{{stream_id}}/{{resource}}", storage_const::M3U_RESOURCE_PATH), axum::routing::get(m3u_api_resource), ) } diff --git a/backend/src/api/endpoints/mod.rs b/backend/src/api/endpoints/mod.rs index 919d786b8..da7091bf9 100644 --- a/backend/src/api/endpoints/mod.rs +++ b/backend/src/api/endpoints/mod.rs @@ -1,17 +1,17 @@ -pub(in crate::api) mod download_api; -pub(in crate::api) mod v1_api; -pub(in crate::api) mod xtream_api; -pub(in crate::api) mod m3u_api; -pub(in crate::api) mod xmltv_api; -pub(in crate::api) mod web_index; -pub(in crate::api) mod hls_api; -mod user_api; -pub(in crate::api) mod hdhomerun_api; mod api_playlist_utils; -pub(in crate::api) mod websocket_api; pub(in crate::api) mod custom_video_stream_api; +pub(in crate::api) mod download_api; +mod extract_accept_header; +pub(in crate::api) mod hdhomerun_api; +pub(in crate::api) mod hls_api; +mod library_api; +pub(in crate::api) mod m3u_api; +mod user_api; +pub(in crate::api) mod v1_api; +mod v1_api_config; mod v1_api_playlist; mod v1_api_user; -mod v1_api_config; -mod extract_accept_header; -mod library_api; \ No newline at end of file +pub(in crate::api) mod web_index; +pub(in crate::api) mod websocket_api; +pub(in crate::api) mod xmltv_api; +pub(in crate::api) mod xtream_api; diff --git a/backend/src/api/endpoints/user_api.rs b/backend/src/api/endpoints/user_api.rs index f6e60edd7..7cef565b9 100644 --- a/backend/src/api/endpoints/user_api.rs +++ b/backend/src/api/endpoints/user_api.rs @@ -1,20 +1,23 @@ -use crate::api::api_utils::try_unwrap_body; -use crate::api::api_utils::{get_user_target_by_username, get_username_from_auth_header}; -use crate::api::model::AppState; -use crate::auth::validator_user; -use crate::auth::AuthBearer; -use crate::model::PlaylistXtreamCategory; -use crate::model::{AppConfig, ConfigTarget}; -use crate::repository::{iter_raw_m3u_target_playlist, load_user_bouquet_as_json, save_user_bouquet}; -use crate::repository::xtream_get_playlist_categories; +use crate::{ + api::{ + api_utils::{get_user_target_by_username, get_username_from_auth_header, try_unwrap_body}, + model::AppState, + }, + auth::{validator_user, AuthBearer}, + model::{AppConfig, ConfigTarget, PlaylistXtreamCategory}, + repository::{ + iter_raw_m3u_target_playlist, load_user_bouquet_as_json, save_user_bouquet, xtream_get_playlist_categories, + }, +}; use axum::response::IntoResponse; use bytes::Bytes; use futures::{stream, StreamExt}; use log::error; -use shared::model::{PlaylistBouquetDto, TargetType, XtreamCluster}; -use std::collections::HashSet; -use std::sync::Arc; -use shared::utils::concat_path_leading_slash; +use shared::{ + model::{PlaylistBouquetDto, TargetType, XtreamCluster}, + utils::concat_path_leading_slash, +}; +use std::{collections::HashSet, sync::Arc}; fn get_categories_from_xtream(categories: Option>) -> Vec { let mut groups: Vec = Vec::new(); @@ -26,7 +29,6 @@ fn get_categories_from_xtream(categories: Option>) - groups } - async fn get_categories_from_m3u_playlist(target: &ConfigTarget, config: &AppConfig) -> Vec> { let mut groups = Vec::new(); if let Some(mut iter) = iter_raw_m3u_target_playlist(config, target, None).await { @@ -59,16 +61,28 @@ async fn playlist_categories( let target_name = &target.name; let xtream_stream = if target.has_output(TargetType::Xtream) { let config = &app_state.app_config.config.load(); - let live_categories = get_categories_from_xtream(xtream_get_playlist_categories(config, target_name, XtreamCluster::Live).await); - let vod_categories = get_categories_from_xtream(xtream_get_playlist_categories(config, target_name, XtreamCluster::Video).await); - let series_categories = get_categories_from_xtream(xtream_get_playlist_categories(config, target_name, XtreamCluster::Series).await); + let live_categories = get_categories_from_xtream( + xtream_get_playlist_categories(config, target_name, XtreamCluster::Live).await, + ); + let vod_categories = get_categories_from_xtream( + xtream_get_playlist_categories(config, target_name, XtreamCluster::Video).await, + ); + let series_categories = get_categories_from_xtream( + xtream_get_playlist_categories(config, target_name, XtreamCluster::Series).await, + ); stream::iter(vec![ Ok::(Bytes::from(r#"{"live": "#)), - Ok::(Bytes::from(serde_json::to_string(&live_categories).unwrap_or("[]".to_string()))), + Ok::(Bytes::from( + serde_json::to_string(&live_categories).unwrap_or("[]".to_string()), + )), Ok::(Bytes::from(r#", "vod": "#.to_string())), - Ok::(Bytes::from(serde_json::to_string(&vod_categories).unwrap_or("[]".to_string()))), + Ok::(Bytes::from( + serde_json::to_string(&vod_categories).unwrap_or("[]".to_string()), + )), Ok::(Bytes::from(r#", "series": "#)), - Ok::(Bytes::from(serde_json::to_string(&series_categories).unwrap_or("[]".to_string()))), + Ok::(Bytes::from( + serde_json::to_string(&series_categories).unwrap_or("[]".to_string()), + )), Ok::(Bytes::from(r"}")), ]) } else { @@ -79,7 +93,9 @@ async fn playlist_categories( let live_categories = get_categories_from_m3u_playlist(&target, &app_state.app_config).await; stream::iter(vec![ Ok::(Bytes::from(r#"{"live": "#)), - Ok::(Bytes::from(serde_json::to_string(&live_categories).unwrap_or("[]".to_string()))), + Ok::(Bytes::from( + serde_json::to_string(&live_categories).unwrap_or("[]".to_string()), + )), Ok::(Bytes::from(r#","vod":[],"series":[]}"#)), ]) } else { @@ -140,9 +156,11 @@ async fn playlist_bouquet( return try_unwrap_body!(axum::response::Response::builder() .status(axum::http::StatusCode::OK) .header("Content-Type", mime::APPLICATION_JSON.to_string()) - .body(axum::body::Body::from(format!(r#"{{"xtream": {}, "m3u": {} }}"#, + .body(axum::body::Body::from(format!( + r#"{{"xtream": {}, "m3u": {} }}"#, xtream.unwrap_or("null".to_string()), - m3u.unwrap_or("null".to_string()))))); + m3u.unwrap_or("null".to_string()) + )))); } } try_unwrap_body!(axum::response::Response::builder() @@ -152,16 +170,14 @@ async fn playlist_bouquet( } pub fn user_api_register(app_state: Arc, web_ui_path: &str) -> axum::Router> { - axum::Router::new() - .nest( - &concat_path_leading_slash(web_ui_path, "/api/v1/user"), - axum::Router::new() - .route("/playlist/categories", axum::routing::get(playlist_categories)) - .route("/playlist/bouquet", axum::routing::get(playlist_bouquet)) - .route("/playlist/bouquet", axum::routing::post(save_playlist_bouquet)) - .route_layer(axum::middleware::from_fn_with_state(app_state, validator_user)), - ) - + axum::Router::new().nest( + &concat_path_leading_slash(web_ui_path, "/api/v1/user"), + axum::Router::new() + .route("/playlist/categories", axum::routing::get(playlist_categories)) + .route("/playlist/bouquet", axum::routing::get(playlist_bouquet)) + .route("/playlist/bouquet", axum::routing::post(save_playlist_bouquet)) + .route_layer(axum::middleware::from_fn_with_state(app_state, validator_user)), + ) // cfg.service(web::scope("/api/v1/user") // .wrap(HttpAuthentication::with_fn(validator_user)) diff --git a/backend/src/api/endpoints/v1_api.rs b/backend/src/api/endpoints/v1_api.rs index 40affdce4..19a0e4230 100644 --- a/backend/src/api/endpoints/v1_api.rs +++ b/backend/src/api/endpoints/v1_api.rs @@ -1,26 +1,30 @@ -use crate::api::api_utils::{json_or_bin_response, try_unwrap_body, internal_server_error}; -use crate::api::endpoints::download_api; -use crate::api::endpoints::user_api::user_api_register; -use crate::api::endpoints::v1_api_playlist::v1_api_playlist_register; -use crate::api::endpoints::v1_api_user::v1_api_user_register; -use crate::api::model::AppState; -use crate::auth::validator_admin; -use crate::utils::ip_checker::get_ips; -use crate::{VERSION}; +use crate::{ + api::{ + api_utils::{internal_server_error, json_or_bin_response, try_unwrap_body}, + endpoints::{ + download_api, extract_accept_header::ExtractAcceptHeader, library_api::library_api_register, + user_api::user_api_register, v1_api_config::v1_api_config_register, + v1_api_playlist::v1_api_playlist_register, v1_api_user::v1_api_user_register, + }, + model::AppState, + }, + auth::validator_admin, + model::InputSource, + repository::get_geoip_path, + utils::{ip_checker::get_ips, request::download_text_content, GeoIp}, + VERSION, +}; use axum::response::IntoResponse; -use shared::model::{default_geoip_url, InputFetchMethod, IpCheckDto, StatusCheck}; -use shared::utils::{concat_path_leading_slash, Internable}; -use std::collections::{BTreeMap, HashMap}; -use std::io::{Cursor}; -use std::sync::Arc; use log::{error, info}; -use crate::api::endpoints::extract_accept_header::ExtractAcceptHeader; -use crate::api::endpoints::v1_api_config::v1_api_config_register; -use crate::api::endpoints::library_api::library_api_register; -use crate::model::InputSource; -use crate::repository::get_geoip_path; -use crate::utils::GeoIp; -use crate::utils::request::download_text_content; +use shared::{ + model::{default_geoip_url, InputFetchMethod, IpCheckDto, StatusCheck}, + utils::{concat_path_leading_slash, Internable}, +}; +use std::{ + collections::{BTreeMap, HashMap}, + io::Cursor, + sync::Arc, +}; async fn create_ipinfo_check(app_state: &Arc) -> Option<(Option, Option)> { let config = app_state.app_config.config.load(); @@ -35,9 +39,7 @@ async fn create_ipinfo_check(app_state: &Arc) -> Option<(Option) -> StatusCheck { let cache = match app_state.cache.load().as_ref().as_ref() { None => None, - Some(lock) => { - Some(lock.lock().await.get_size_text()) - } + Some(lock) => Some(lock.lock().await.get_size_text()), }; let (active_users, active_user_connections, active_user_streams) = { let active_user = &app_state.active_users; @@ -45,7 +47,8 @@ pub async fn create_status_check(app_state: &Arc) -> StatusCheck { (user_count, connection_count, active_user.active_streams().await) }; - let active_provider_connections = app_state.active_provider.active_connections().await.map(|c| c.into_iter().collect::>()); + let active_provider_connections = + app_state.active_provider.active_connections().await.map(|c| c.into_iter().collect::>()); StatusCheck { status: "ok".to_string(), @@ -62,19 +65,25 @@ pub async fn create_status_check(app_state: &Arc) -> StatusCheck { async fn status(axum::extract::State(app_state): axum::extract::State>) -> axum::response::Response { let status = create_status_check(&app_state).await; match serde_json::to_string_pretty(&status) { - Ok(pretty_json) => try_unwrap_body!(axum::response::Response::builder().status(axum::http::StatusCode::OK) - .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()).body(pretty_json)), + Ok(pretty_json) => try_unwrap_body!(axum::response::Response::builder() + .status(axum::http::StatusCode::OK) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) + .body(pretty_json)), Err(_) => axum::Json(status).into_response(), } } -async fn streams(ExtractAcceptHeader(accept): ExtractAcceptHeader, - axum::extract::State(app_state): axum::extract::State>) -> axum::response::Response { +async fn streams( + ExtractAcceptHeader(accept): ExtractAcceptHeader, + axum::extract::State(app_state): axum::extract::State>, +) -> axum::response::Response { let streams = app_state.active_users.active_streams().await; json_or_bin_response(accept.as_deref(), &streams).into_response() } -async fn geoip_update(axum::extract::State(app_state): axum::extract::State>) -> axum::response::Response { +async fn geoip_update( + axum::extract::State(app_state): axum::extract::State>, +) -> axum::response::Response { let config = app_state.app_config.config.load(); if let Some(geoip) = config.reverse_proxy.as_ref().and_then(|r| r.geoip.as_ref()) { if geoip.enabled { @@ -82,7 +91,7 @@ async fn geoip_update(axum::extract::State(app_state): axum::extract::State { - let reader = Cursor::new(content); - let mut geoip = GeoIp::new(); - let result = { - match geoip.import_ipv4_from_csv(reader, geoip_db_path) { - Ok(size) => { - (Some(size), None) - } - Err(err) => (None, Some(err)) + Ok((content, _)) => { + let reader = Cursor::new(content); + let mut geoip = GeoIp::new(); + let result = { + match geoip.import_ipv4_from_csv(reader, geoip_db_path) { + Ok(size) => (Some(size), None), + Err(err) => (None, Some(err)), } - }; + }; - return match result { - (Some(_), None) => { - info!("GeoIp db updated"); - app_state.geoip.store(Some(Arc::new(geoip))); - axum::http::StatusCode::OK.into_response() - }, - (None, Some(err)) => { - error!("Failed to process geoip db: {err}"); - internal_server_error!() - }, - _ => internal_server_error!() - } - } - Err(err) => { - error!("Failed to download geoip db: {err}"); - axum::http::StatusCode::BAD_REQUEST.into_response() - } - } + return match result { + (Some(_), None) => { + info!("GeoIp db updated"); + app_state.geoip.store(Some(Arc::new(geoip))); + axum::http::StatusCode::OK.into_response() + } + (None, Some(err)) => { + error!("Failed to process geoip db: {err}"); + internal_server_error!() + } + _ => internal_server_error!(), + }; + } + Err(err) => { + error!("Failed to download geoip db: {err}"); + axum::http::StatusCode::BAD_REQUEST.into_response() + } + }; } } axum::http::StatusCode::BAD_REQUEST.into_response() @@ -138,20 +145,23 @@ async fn geoip_update(axum::extract::State(app_state): axum::extract::State>) -> axum::response::Response { if let Some((ipv4, ipv6)) = create_ipinfo_check(&app_state).await { - let ipcheck = IpCheckDto { - ipv4, - ipv6, - }; + let ipcheck = IpCheckDto { ipv4, ipv6 }; return match serde_json::to_string(&ipcheck) { - Ok(json) => try_unwrap_body!(axum::response::Response::builder().status(axum::http::StatusCode::OK) - .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()).body(json)), + Ok(json) => try_unwrap_body!(axum::response::Response::builder() + .status(axum::http::StatusCode::OK) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) + .body(json)), Err(_) => axum::Json(ipcheck).into_response(), }; } axum::http::StatusCode::BAD_REQUEST.into_response() } -pub fn v1_api_register(web_auth_enabled: bool, app_state: Arc, web_ui_path: &str) -> axum::Router> { +pub fn v1_api_register( + web_auth_enabled: bool, + app_state: Arc, + web_ui_path: &str, +) -> axum::Router> { let mut router = axum::Router::new(); router = router .route("/status", axum::routing::get(status)) diff --git a/backend/src/api/endpoints/v1_api_config.rs b/backend/src/api/endpoints/v1_api_config.rs index 9f519d706..538201037 100644 --- a/backend/src/api/endpoints/v1_api_config.rs +++ b/backend/src/api/endpoints/v1_api_config.rs @@ -1,19 +1,30 @@ -use crate::api::api_utils::{internal_server_error, try_unwrap_body}; -use crate::api::config_file::ConfigFile; -use crate::api::model::AppState; -use crate::model::{ApiProxyConfig, InputSource, validate_library_paths_from_dto}; -use crate::utils; -use crate::utils::request::download_text_content; -use crate::utils::{persist_messaging_templates, prepare_sources_batch, prepare_users, read_api_proxy_file}; -use axum::response::IntoResponse; -use axum::Router; +use crate::{ + api::{ + api_utils::{internal_server_error, try_unwrap_body}, + config_file::ConfigFile, + model::AppState, + }, + model::{validate_library_paths_from_dto, ApiProxyConfig, InputSource}, + utils, + utils::{ + persist_messaging_templates, prepare_sources_batch, prepare_users, read_api_proxy_file, + request::download_text_content, + }, +}; +use axum::{response::IntoResponse, Router}; use log::error; use serde_json::json; -use shared::error::TuliproxError; -use shared::model::{ApiProxyConfigDto, ConfigDto, SourcesConfigDto}; +use shared::{ + error::TuliproxError, + model::{ApiProxyConfigDto, ConfigDto, SourcesConfigDto}, +}; use std::sync::Arc; -pub(in crate::api::endpoints) async fn intern_save_config_api_proxy(backup_dir: &str, api_proxy: &ApiProxyConfigDto, file_path: &str) -> Option { +pub(in crate::api::endpoints) async fn intern_save_config_api_proxy( + backup_dir: &str, + api_proxy: &ApiProxyConfigDto, + file_path: &str, +) -> Option { match utils::save_api_proxy(file_path, backup_dir, api_proxy).await { Ok(()) => {} Err(err) => { @@ -49,7 +60,8 @@ async fn save_config_main( } else { if let Err(err) = persist_messaging_templates(&app_state, &mut cfg).await { error!("Failed to persist messaging templates: {err}"); - return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))).into_response(); + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) + .into_response(); } let paths = app_state.app_config.paths.load(); @@ -57,7 +69,8 @@ async fn save_config_main( let config = app_state.app_config.config.load(); let backup_dir = config.get_backup_dir(); if let Some(err) = intern_save_config_main(file_path, backup_dir.as_ref(), &cfg).await { - return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))).into_response(); + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) + .into_response(); } axum::http::StatusCode::OK.into_response() } @@ -71,17 +84,15 @@ async fn save_config_sources( Ok(value) => value, Err(err) => { error!("Failed to validate source.yml {err}"); - return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": err.to_string()}))).into_response(); + return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": err.to_string()}))) + .into_response(); } }; if let Some(template_definition) = templates_to_persist.as_ref() { if let Err(err) = utils::persist_templates_config(&app_state, template_definition).await { error!("Failed to save template config {err}"); - return ( - axum::http::StatusCode::INTERNAL_SERVER_ERROR, - axum::Json(json!({"error": err.to_string()})), - ) + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) .into_response(); } } @@ -90,10 +101,7 @@ async fn save_config_sources( Ok(_) => {} Err(err) => { error!("Failed to persist source.yml {err}"); - return ( - axum::http::StatusCode::INTERNAL_SERVER_ERROR, - axum::Json(json!({"error": err.to_string()})), - ) + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) .into_response(); } } @@ -101,10 +109,7 @@ async fn save_config_sources( // Reload from disk so runtime always uses fully prepared sources/mappings/templates. if let Err(err) = ConfigFile::load_sources(&app_state).await { error!("Failed to reload prepared sources after save {err}"); - return ( - axum::http::StatusCode::INTERNAL_SERVER_ERROR, - axum::Json(json!({"error": err.to_string()})), - ) + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) .into_response(); } @@ -112,9 +117,8 @@ async fn save_config_sources( axum::http::StatusCode::OK.into_response() } - async fn get_config_api_proxy_config( - axum::extract::State(app_state): axum::extract::State> + axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse + Send { let paths = app_state.app_config.paths.load(); let api_proxy_file_path = paths.api_proxy_file_path.as_str(); @@ -139,14 +143,14 @@ async fn save_config_api_proxy_config( ) -> impl IntoResponse + Send { for server_info in &mut req_api_proxy.server { if !server_info.validate() { - return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid content"}))).into_response(); + return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid content"}))) + .into_response(); } } // TODO if hot reload is on, it is loaded twice, avoid this // Build the updated config without mutating global state yet - let base = app_state.app_config.api_proxy.load() - .as_deref().cloned().unwrap_or_default(); + let base = app_state.app_config.api_proxy.load().as_deref().cloned().unwrap_or_default(); let updated_api_proxy = ApiProxyConfig { use_user_db: req_api_proxy.use_user_db, server: req_api_proxy.server.iter().map(Into::into).collect(), @@ -157,21 +161,23 @@ async fn save_config_api_proxy_config( let backup_dir = config.get_backup_dir(); let paths = app_state.app_config.paths.load(); - if let Some(err) = intern_save_config_api_proxy(backup_dir.as_ref(), &ApiProxyConfigDto::from(&updated_api_proxy), paths.api_proxy_file_path.as_str()).await { - return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))).into_response(); + if let Some(err) = intern_save_config_api_proxy( + backup_dir.as_ref(), + &ApiProxyConfigDto::from(&updated_api_proxy), + paths.api_proxy_file_path.as_str(), + ) + .await + { + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) + .into_response(); } // Persist succeeded — now update in‑memory state - app_state - .app_config - .api_proxy - .store(Some(Arc::new(updated_api_proxy))); + app_state.app_config.api_proxy.store(Some(Arc::new(updated_api_proxy))); axum::http::StatusCode::OK.into_response() } -async fn config( - axum::extract::State(app_state): axum::extract::State>, -) -> impl IntoResponse + Send { +async fn config(axum::extract::State(app_state): axum::extract::State>) -> impl IntoResponse + Send { let paths = app_state.app_config.paths.load(); match utils::read_app_config_dto(&paths, true, false) { Ok(mut app_config) => { @@ -208,7 +214,7 @@ async fn config_batch_content( None, false, ) - .await + .await { Ok((content, _path)) => { // Return CSV with explicit content-type @@ -224,10 +230,10 @@ async fn config_batch_content( }; } } - (axum::http::StatusCode::NOT_FOUND, axum::Json(json!({"error": "Input not found or batch URL missing"}))).into_response() + (axum::http::StatusCode::NOT_FOUND, axum::Json(json!({"error": "Input not found or batch URL missing"}))) + .into_response() } - pub fn v1_api_config_register(router: Router>) -> axum::Router> { router .route("/config", axum::routing::get(config)) diff --git a/backend/src/api/endpoints/v1_api_playlist.rs b/backend/src/api/endpoints/v1_api_playlist.rs index 7fc7faeee..76d828899 100644 --- a/backend/src/api/endpoints/v1_api_playlist.rs +++ b/backend/src/api/endpoints/v1_api_playlist.rs @@ -1,21 +1,31 @@ -use crate::api::api_utils::{create_api_proxy_user, json_or_bin_response}; -use crate::api::endpoints::api_playlist_utils::{get_playlist_for_custom_provider, get_playlist_for_input, get_playlist_for_target}; -use crate::api::endpoints::extract_accept_header::ExtractAcceptHeader; -use crate::api::model::AppState; -use crate::auth::create_access_token; -use crate::model::{parse_xmltv_for_web_ui_from_url, ConfigInput, ConfigInputFlags, ConfigInputFlagsSet, ConfigInputOptions}; -use axum::response::IntoResponse; -use axum::{Router}; +use crate::{ + api::{ + api_utils::{create_api_proxy_user, json_or_bin_response}, + endpoints::{ + api_playlist_utils::{get_playlist_for_custom_provider, get_playlist_for_input, get_playlist_for_target}, + extract_accept_header::ExtractAcceptHeader, + xmltv_api::serve_epg_web_ui, + xtream_api::xtream_get_stream_info_response, + }, + model::AppState, + }, + auth::create_access_token, + model::{parse_xmltv_for_web_ui_from_url, ConfigInput, ConfigInputFlags, ConfigInputFlagsSet, ConfigInputOptions}, + processing::processor::exec_processing, + repository::xtream_get_item_for_stream_id, +}; +use axum::{response::IntoResponse, Router}; use log::{debug, error}; use serde_json::json; -use shared::model::{InputType, PlaylistEpgRequest, PlaylistRequest, ProxyType, TargetType, UiPlaylistItem, WebplayerUrlRequest, XtreamCluster}; -use shared::utils::{sanitize_sensitive_info, Internable}; +use shared::{ + model::{ + InputType, PlaylistEpgRequest, PlaylistRequest, ProxyType, TargetType, UiPlaylistItem, WebplayerUrlRequest, + XtreamCluster, + }, + utils::{sanitize_sensitive_info, Internable}, +}; use std::sync::Arc; use url::Url; -use crate::api::endpoints::xmltv_api::{serve_epg_web_ui}; -use crate::api::endpoints::xtream_api::xtream_get_stream_info_response; -use crate::processing::processor::exec_processing; -use crate::repository::xtream_get_item_for_stream_id; fn create_config_input_for_m3u(url: &str) -> ConfigInput { ConfigInput { @@ -75,11 +85,24 @@ async fn playlist_update( let provider_manager = Arc::clone(&app_state.active_provider); let disabled_headers = app_state.get_disabled_headers(); let metadata_manager = Arc::clone(&app_state.metadata_manager); + let update_guard = app_state.update_guard.clone(); tokio::spawn({ async move { - exec_processing(&http_client, app_config, valid_targets, Some(event_manager), - Some(playlist_state), Some(app_state.update_guard.clone()), - disabled_headers, Some(provider_manager), Some(metadata_manager), None, None).await; + exec_processing( + &http_client, + app_config, + valid_targets, + Some(event_manager), + Some(app_state.clone()), + Some(playlist_state), + Some(update_guard), + disabled_headers, + Some(provider_manager), + Some(metadata_manager), + None, + None, + ) + .await; } }); axum::http::StatusCode::ACCEPTED.into_response() @@ -100,34 +123,60 @@ async fn playlist_content( let _config = app_state.app_config.config.load(); let client = app_state.http_client.load(); match playlist_req { - PlaylistRequest::Target(target_id) => { - get_playlist_for_target(app_state.app_config.get_target_by_id(*target_id).as_deref(), &app_state.app_config, cluster, accept.as_deref()).await.into_response() - } - PlaylistRequest::Input(input_id) => { - get_playlist_for_input(app_state.app_config.get_input_by_id(*input_id).as_ref(), &app_state.app_config, cluster, accept.as_deref()).await.into_response() - } - PlaylistRequest::CustomXtream(xtream) => { - match Url::parse(&xtream.url) { - Ok(parsed) if parsed.scheme() == "http" || parsed.scheme() == "https" => { - let input = Arc::new(create_config_input_for_xtream(&xtream.username, &xtream.password, &xtream.url)); - get_playlist_for_custom_provider(client.as_ref(), Some(&input), &app_state.app_config, cluster, accept.as_deref()).await.into_response() - } - _ => { - (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid url scheme; only http/https are allowed"}))).into_response() - } + PlaylistRequest::Target(target_id) => get_playlist_for_target( + app_state.app_config.get_target_by_id(*target_id).as_deref(), + &app_state.app_config, + cluster, + accept.as_deref(), + ) + .await + .into_response(), + PlaylistRequest::Input(input_id) => get_playlist_for_input( + app_state.app_config.get_input_by_id(*input_id).as_ref(), + &app_state.app_config, + cluster, + accept.as_deref(), + ) + .await + .into_response(), + PlaylistRequest::CustomXtream(xtream) => match Url::parse(&xtream.url) { + Ok(parsed) if parsed.scheme() == "http" || parsed.scheme() == "https" => { + let input = Arc::new(create_config_input_for_xtream(&xtream.username, &xtream.password, &xtream.url)); + get_playlist_for_custom_provider( + client.as_ref(), + Some(&input), + &app_state.app_config, + cluster, + accept.as_deref(), + ) + .await + .into_response() } - } - PlaylistRequest::CustomM3u(m3u) => { - match Url::parse(&m3u.url) { - Ok(parsed) if parsed.scheme() == "http" || parsed.scheme() == "https" => { - let input = Arc::new(create_config_input_for_m3u(&m3u.url)); - get_playlist_for_custom_provider(client.as_ref(), Some(&input), &app_state.app_config, cluster, accept.as_deref()).await.into_response() - } - _ => { - (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": "Invalid url scheme; only http/https are allowed"}))).into_response() - } + _ => ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(json!({"error": "Invalid url scheme; only http/https are allowed"})), + ) + .into_response(), + }, + PlaylistRequest::CustomM3u(m3u) => match Url::parse(&m3u.url) { + Ok(parsed) if parsed.scheme() == "http" || parsed.scheme() == "https" => { + let input = Arc::new(create_config_input_for_m3u(&m3u.url)); + get_playlist_for_custom_provider( + client.as_ref(), + Some(&input), + &app_state.app_config, + cluster, + accept.as_deref(), + ) + .await + .into_response() } - } + _ => ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(json!({"error": "Invalid url scheme; only http/https are allowed"})), + ) + .into_response(), + }, } } @@ -138,14 +187,7 @@ macro_rules! create_player_api_for_cluster { axum::extract::State(app_state): axum::extract::State>, axum::extract::Json(playlist_req): axum::extract::Json, ) -> impl IntoResponse + Send { - playlist_content( - accept.clone(), - &app_state, - &playlist_req, - $cluster - ) - .await - .into_response() + playlist_content(accept.clone(), &app_state, &playlist_req, $cluster).await.into_response() } }; } @@ -155,10 +197,7 @@ create_player_api_for_cluster!(playlist_content_vod, XtreamCluster::Video); create_player_api_for_cluster!(playlist_content_series, XtreamCluster::Series); async fn playlist_series_info( - axum::extract::Path((virtual_id, _provider_id)): axum::extract::Path<( - String, - String, - )>, + axum::extract::Path((virtual_id, _provider_id)): axum::extract::Path<(String, String)>, axum::extract::State(app_state): axum::extract::State>, axum::extract::Json(playlist_req): axum::extract::Json, ) -> impl IntoResponse + Send { @@ -168,7 +207,15 @@ async fn playlist_series_info( if target.has_output(TargetType::Xtream) { let mut user = create_api_proxy_user(&app_state); user.proxy = ProxyType::Redirect; - return xtream_get_stream_info_response(&app_state, &user, &target, &virtual_id, XtreamCluster::Series).await.into_response(); + return xtream_get_stream_info_response( + &app_state, + &user, + &target, + &virtual_id, + XtreamCluster::Series, + ) + .await + .into_response(); } } } @@ -184,7 +231,7 @@ async fn playlist_series_info( PlaylistRequest::CustomXtream(_xtream) => { // TODO: Implement series info retrieval for custom Xtream requests debug!("TODO: Implement series info retrieval for custom Xtream requests"); - }, + } PlaylistRequest::CustomM3u(_) => {} } axum::http::StatusCode::NO_CONTENT.into_response() @@ -196,10 +243,20 @@ async fn playlist_webplayer( ) -> impl axum::response::IntoResponse + Send { let access_token = create_access_token(&app_state.app_config.access_token_secret, 30); let config = app_state.app_config.config.load(); - let server_name = config.web_ui.as_ref().and_then(|web_ui| web_ui.player_server.as_ref()).map_or("default", |server_name| server_name.as_str()); + let server_name = config + .web_ui + .as_ref() + .and_then(|web_ui| web_ui.player_server.as_ref()) + .map_or("default", |server_name| server_name.as_str()); let server_info = app_state.app_config.get_server_info(server_name); let base_url = server_info.get_base_url(); - format!("{base_url}/token/{access_token}/{}/{}/{}", playlist_item.target_id, playlist_item.cluster.as_stream_type(), playlist_item.virtual_id).into_response() + format!( + "{base_url}/token/{access_token}/{}/{}/{}", + playlist_item.target_id, + playlist_item.cluster.as_stream_type(), + playlist_item.virtual_id + ) + .into_response() } async fn playlist_epg( @@ -257,10 +314,10 @@ async fn playlist_episode_item( if let Some(target) = app_state.app_config.get_target_by_id(target_id) { if target.has_output(TargetType::Xtream) { if let Ok(vid) = virtual_id.parse::() { - if let Ok(pli) = xtream_get_item_for_stream_id( - vid, &app_state, &target, Some(XtreamCluster::Series) - ).await { - return axum::Json(json!(UiPlaylistItem::from(pli))).into_response(); + if let Ok(pli) = + xtream_get_item_for_stream_id(vid, &app_state, &target, Some(XtreamCluster::Series)).await + { + return axum::Json(json!(UiPlaylistItem::from(pli))).into_response(); } } } diff --git a/backend/src/api/endpoints/v1_api_user.rs b/backend/src/api/endpoints/v1_api_user.rs index 362c09e5a..c44532a7a 100644 --- a/backend/src/api/endpoints/v1_api_user.rs +++ b/backend/src/api/endpoints/v1_api_user.rs @@ -1,14 +1,18 @@ -use crate::api::model::AppState; -use crate::api::panel_api::{sync_panel_api_alias_pool_for_target, target_has_alias_pool_min}; -use crate::model::{ApiProxyConfig, ProxyUserCredentials, TargetUser}; -use crate::repository::store_api_user; -use axum::response::IntoResponse; -use axum::Router; +use crate::{ + api::{ + model::AppState, + panel_api::{sync_panel_api_alias_pool_for_target, target_has_alias_pool_min}, + }, + model::{ApiProxyConfig, ProxyUserCredentials, TargetUser}, + repository::store_api_user, +}; +use axum::{response::IntoResponse, Router}; use serde_json::json; -use shared::model::{ApiProxyConfigDto, ProxyUserCredentialsDto}; -use shared::utils::mask_credentials; -use std::path::PathBuf; -use std::sync::Arc; +use shared::{ + model::{ApiProxyConfigDto, ProxyUserCredentialsDto}, + utils::mask_credentials, +}; +use std::{path::PathBuf, sync::Arc}; #[allow(clippy::too_many_lines)] async fn save_config_api_proxy_user( @@ -18,19 +22,11 @@ async fn save_config_api_proxy_user( axum::extract::Json(mut credential): axum::extract::Json, ) -> impl axum::response::IntoResponse + Send { let virtual_file = PathBuf::from("api_proxy"); - let _lock = app_state - .app_config - .file_locks - .write_lock(&virtual_file) - .await; + let _lock = app_state.app_config.file_locks.write_lock(&virtual_file).await; credential.prepare(); if let Err(err) = credential.validate() { - return ( - axum::http::StatusCode::BAD_REQUEST, - axum::Json(json!({"error": err.to_string()})), - ) - .into_response(); + return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": err.to_string()}))).into_response(); } let is_update = method == axum::http::Method::PUT; @@ -55,9 +51,7 @@ async fn save_config_api_proxy_user( if u == c && user.username != credential.username { return ( axum::http::StatusCode::BAD_REQUEST, - axum::Json( - json!({"error": format!("Duplicate token {}", mask_credentials(c))}), - ), + axum::Json(json!({"error": format!("Duplicate token {}", mask_credentials(c))})), ) .into_response(); } @@ -67,8 +61,9 @@ async fn save_config_api_proxy_user( if !is_update { return ( axum::http::StatusCode::BAD_REQUEST, - axum::Json(json!({"error": format!("Duplicate username {}", &credential.username)})) - ).into_response(); + axum::Json(json!({"error": format!("Duplicate username {}", &credential.username)})), + ) + .into_response(); } // mark position of user (for update / move) @@ -91,10 +86,7 @@ async fn save_config_api_proxy_user( let target_idx = if let Some(idx) = existing_target_index { idx } else { - api_proxy.user.push(TargetUser { - target: target_name.clone(), - credentials: vec![], - }); + api_proxy.user.push(TargetUser { target: target_name.clone(), credentials: vec![] }); api_proxy.user.len() - 1 }; @@ -106,14 +98,11 @@ async fn save_config_api_proxy_user( if user_target_idx == target_idx { // Update - api_proxy.user[user_target_idx].credentials[user_idx] = - ProxyUserCredentials::from(&credential); + api_proxy.user[user_target_idx].credentials[user_idx] = ProxyUserCredentials::from(&credential); } else { // Move: remove from old target and insert into new target api_proxy.user[user_target_idx].credentials.remove(user_idx); - api_proxy.user[target_idx] - .credentials - .push(ProxyUserCredentials::from(&credential)); + api_proxy.user[target_idx].credentials.push(ProxyUserCredentials::from(&credential)); remove_empty_target = api_proxy.user[user_target_idx].credentials.is_empty(); } @@ -122,19 +111,14 @@ async fn save_config_api_proxy_user( } } else { // new user - api_proxy.user[target_idx] - .credentials - .push(ProxyUserCredentials::from(&credential)); + api_proxy.user[target_idx].credentials.push(ProxyUserCredentials::from(&credential)); } let new_api_proxy = Arc::new(api_proxy); if new_api_proxy.use_user_db { if let Err(err) = store_api_user(&app_state.app_config, &new_api_proxy.user).await { - return ( - axum::http::StatusCode::INTERNAL_SERVER_ERROR, - axum::Json(json!({"error": err.to_string()})), - ) + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) .into_response(); } } else { @@ -148,19 +132,13 @@ async fn save_config_api_proxy_user( ) .await { - return ( - axum::http::StatusCode::INTERNAL_SERVER_ERROR, - axum::Json(json!({"error": err.to_string()})), - ) + return (axum::http::StatusCode::INTERNAL_SERVER_ERROR, axum::Json(json!({"error": err.to_string()}))) .into_response(); } } // Update state after successful save - app_state - .app_config - .api_proxy - .store(Some(Arc::clone(&new_api_proxy))); + app_state.app_config.api_proxy.store(Some(Arc::clone(&new_api_proxy))); if target_has_alias_pool_min(&app_state, &target_name) { let app_state_clone = Arc::clone(&app_state); @@ -184,9 +162,7 @@ async fn delete_config_api_proxy_user( for target_user in &mut api_proxy.user { if target_user.target == target_name { let count = target_user.credentials.len(); - target_user - .credentials - .retain(|user| user.username != username); + target_user.credentials.retain(|user| user.username != username); modified = count != target_user.credentials.len(); break; } @@ -205,13 +181,12 @@ async fn delete_config_api_proxy_user( let config = app_state.app_config.config.load(); let backup_dir = config.get_backup_dir(); let paths = app_state.app_config.paths.load(); - if let Some(err) = - crate::api::endpoints::v1_api_config::intern_save_config_api_proxy( - backup_dir.as_ref(), - &ApiProxyConfigDto::from(&*new_api_proxy), - paths.api_proxy_file_path.as_str(), - ) - .await + if let Some(err) = crate::api::endpoints::v1_api_config::intern_save_config_api_proxy( + backup_dir.as_ref(), + &ApiProxyConfigDto::from(&*new_api_proxy), + paths.api_proxy_file_path.as_str(), + ) + .await { return ( axum::http::StatusCode::INTERNAL_SERVER_ERROR, @@ -220,16 +195,11 @@ async fn delete_config_api_proxy_user( .into_response(); } } - app_state - .app_config - .api_proxy - .store(Some(Arc::clone(&new_api_proxy))); + app_state.app_config.api_proxy.store(Some(Arc::clone(&new_api_proxy))); } else { return ( axum::http::StatusCode::BAD_REQUEST, - axum::Json( - json!({"error": format!("User not found {username} in target {target_name}")}), - ), + axum::Json(json!({"error": format!("User not found {username} in target {target_name}")})), ) .into_response(); } @@ -239,16 +209,7 @@ async fn delete_config_api_proxy_user( pub fn v1_api_user_register(router: Router>) -> axum::Router> { router - .route( - "/user/{target}", - axum::routing::post(save_config_api_proxy_user), - ) - .route( - "/user/{target}", - axum::routing::put(save_config_api_proxy_user), - ) - .route( - "/user/{target}/{username}", - axum::routing::delete(delete_config_api_proxy_user), - ) + .route("/user/{target}", axum::routing::post(save_config_api_proxy_user)) + .route("/user/{target}", axum::routing::put(save_config_api_proxy_user)) + .route("/user/{target}/{username}", axum::routing::delete(delete_config_api_proxy_user)) } diff --git a/backend/src/api/endpoints/web_index.rs b/backend/src/api/endpoints/web_index.rs index f7b8fe977..2c70d0759 100644 --- a/backend/src/api/endpoints/web_index.rs +++ b/backend/src/api/endpoints/web_index.rs @@ -1,29 +1,31 @@ -use crate::api::api_utils::serve_file; -use crate::api::api_utils::try_unwrap_body; -use crate::api::model::AppState; -use crate::auth::{create_jwt_admin, create_jwt_user, is_admin, verify_password, verify_token, AuthBearer}; -use axum::response::IntoResponse; -use log::{error}; -use serde_json::json; -use shared::model::{TokenResponse, UserCredential, TOKEN_NO_AUTH}; -use shared::utils::{concat_path_leading_slash, CONSTANTS}; -use std::path::{Path, PathBuf}; -use std::sync::Arc; -use axum::body::Body; -use axum::http::Request; +use crate::{ + api::{ + api_utils::{serve_file, try_unwrap_body}, + model::AppState, + }, + auth::{create_jwt_admin, create_jwt_user, is_admin, verify_password, verify_token, AuthBearer}, +}; +use axum::{body::Body, http::Request, response::IntoResponse}; use base64::Engine; +use log::error; +use lol_html::{element, RewriteStrSettings}; //use base64::engine::general_purpose; use openssl::rand::rand_bytes; +use serde_json::json; +use shared::{ + model::{TokenResponse, UserCredential, TOKEN_NO_AUTH}, + utils::{concat_path_leading_slash, CONSTANTS}, +}; +use std::{ + path::{Path, PathBuf}, + sync::Arc, +}; //use openssl::sha::{sha256}; use tower::{Service, ServiceExt}; use tower_http::services::ServeFile; -use lol_html::{element, RewriteStrSettings}; fn no_web_auth_token() -> impl axum::response::IntoResponse + Send { - axum::Json(TokenResponse { - token: TOKEN_NO_AUTH.to_string(), - username: "admin".to_string(), - }).into_response() + axum::Json(TokenResponse { token: TOKEN_NO_AUTH.to_string(), username: "admin".to_string() }).into_response() } async fn token( @@ -45,11 +47,7 @@ async fn token( if verify_password(hash, password.as_bytes()) { if let Ok(token) = create_jwt_admin(web_auth, username) { req.zeroize(); - return axum::Json( - TokenResponse { - token, - username: req.username.clone(), - }).into_response(); + return axum::Json(TokenResponse { token, username: req.username.clone() }).into_response(); } } } @@ -57,11 +55,7 @@ async fn token( if credentials.password == password { if let Ok(token) = create_jwt_user(web_auth, username) { req.zeroize(); - return axum::Json( - TokenResponse { - token, - username: req.username.clone(), - }).into_response(); + return axum::Json(TokenResponse { token, username: req.username.clone() }).into_response(); } } } @@ -94,11 +88,7 @@ async fn token_refresh( create_jwt_user(web_auth, &username) }; if let Ok(token) = new_token { - return axum::Json( - TokenResponse { - token, - username: username.clone(), - }).into_response(); + return axum::Json(TokenResponse { token, username: username.clone() }).into_response(); } } axum::http::StatusCode::UNAUTHORIZED.into_response() @@ -133,7 +123,6 @@ fn inject_nonce_with_parser(html: String, nonce_b64: &str) -> String { lol_html::rewrite_str(&html, settings).unwrap_or(html) } - async fn index( axum::extract::State(app_state): axum::extract::State>, ) -> impl axum::response::IntoResponse + Send { @@ -144,14 +133,20 @@ async fn index( let mut new_content = { if let Some(web_ui_path) = &config.web_ui.as_ref().and_then(|c| c.path.as_ref()) { // modify all url or src attributes in the html file - let mut the_content = CONSTANTS.re_base_href.replace_all(&content, |caps: ®ex::Captures| { - format!(r#"{}="{}""#, &caps[1], concat_path_leading_slash(web_ui_path, &caps[2])) - }).to_string(); + let mut the_content = CONSTANTS + .re_base_href + .replace_all(&content, |caps: ®ex::Captures| { + format!(r#"{}="{}""#, &caps[1], concat_path_leading_slash(web_ui_path, &caps[2])) + }) + .to_string(); // replace wasm paths - the_content = CONSTANTS.re_base_href_wasm.replace_all(&the_content, |caps: ®ex::Captures| { - format!("'{}", concat_path_leading_slash(web_ui_path, &caps[1])) - }).to_string(); + the_content = CONSTANTS + .re_base_href_wasm + .replace_all(&the_content, |caps: ®ex::Captures| { + format!("'{}", concat_path_leading_slash(web_ui_path, &caps[1])) + }) + .to_string(); let new_base = format!(r#""#); @@ -197,11 +192,8 @@ async fn index( let mut builder = axum::response::Response::builder() .header(axum::http::header::CONTENT_TYPE, mime::TEXT_HTML_UTF_8.as_ref()); - if let Some(csp) = config - .web_ui - .as_ref() - .and_then(|w| w.content_security_policy.as_ref()) - .filter(|c| c.enabled) + if let Some(csp) = + config.web_ui.as_ref().and_then(|w| w.content_security_policy.as_ref()).filter(|c| c.enabled) { let mut attrs = vec![ "default-src 'self'".to_string(), @@ -252,7 +244,7 @@ async fn index_config( } if let Some(app_logo) = json_data.get_mut("appLogo") { if let Some(url) = app_logo.as_str() { - let new_url = concat_path_leading_slash(web_ui_path, url); + let new_url = concat_path_leading_slash(web_ui_path, url); *app_logo = json!(new_url); } } @@ -265,7 +257,7 @@ async fn index_config( if let Some(web_path) = json_data.get_mut("webPath") { if let Some(_path) = web_path.as_str() { - let new_url = format!("/{web_ui_path}"); + let new_url = format!("/{web_ui_path}"); *web_path = json!(new_url); } } else { @@ -289,54 +281,57 @@ async fn index_config( pub fn index_register_without_path(web_dir_path: &Path) -> axum::Router> { axum::Router::new() - .nest("/auth", axum::Router::new() - .route("/token", axum::routing::post(token)) - .route("/refresh", axum::routing::post(token_refresh))) - .merge(axum::Router::new() - .route("/", axum::routing::get(index)) - .fallback(axum::routing::get_service(tower_http::services::ServeDir::new(web_dir_path)))) + .nest( + "/auth", + axum::Router::new() + .route("/token", axum::routing::post(token)) + .route("/refresh", axum::routing::post(token_refresh)), + ) + .merge( + axum::Router::new() + .route("/", axum::routing::get(index)) + .fallback(axum::routing::get_service(tower_http::services::ServeDir::new(web_dir_path))), + ) } pub fn index_register_with_path(web_dir_path: &Path, web_ui_path: &str) -> axum::Router> { let web_dir_path_clone = PathBuf::from(web_dir_path); let web_ui_router = axum::Router::new() - .route("/", axum::routing::get(index)) - .route("/config.json", axum::routing::get(index_config)) - .route("/{filename}", axum::routing::get(async move - |axum::extract::Path(filename): axum::extract::Path| { - let full_path = web_dir_path_clone.join(&filename); - let svc = ServeFile::new(full_path); - svc.oneshot(Request::new(Body::empty())).await - })) - .fallback({ - let mut serve_dir = tower_http::services::ServeDir::new(web_dir_path); - let path_prefix = format!("/{web_ui_path}"); - move |req: axum::http::Request<_>| { - let mut path = req.uri().path().to_string(); + .route("/", axum::routing::get(index)) + .route("/config.json", axum::routing::get(index_config)) + .route( + "/{filename}", + axum::routing::get(async move |axum::extract::Path(filename): axum::extract::Path| { + let full_path = web_dir_path_clone.join(&filename); + let svc = ServeFile::new(full_path); + svc.oneshot(Request::new(Body::empty())).await + }), + ) + .fallback({ + let mut serve_dir = tower_http::services::ServeDir::new(web_dir_path); + let path_prefix = format!("/{web_ui_path}"); + move |req: axum::http::Request<_>| { + let mut path = req.uri().path().to_string(); - if path.starts_with(&path_prefix) { - path = path[path_prefix.len()..].to_string(); - } - - let mut builder = axum::http::Uri::builder(); - if let Some(scheme) = req.uri().scheme() { - builder = builder.scheme(scheme.clone()); - } - if let Some(authority) = req.uri().authority() { - builder = builder.authority(authority.clone()); - } - let new_uri = builder.path_and_query(path) - .build() - .unwrap(); - - let new_req = axum::http::Request::builder() - .method(req.method()) - .uri(new_uri) - .body(req.into_body()).unwrap(); - - serve_dir.call(new_req) + if path.starts_with(&path_prefix) { + path = path[path_prefix.len()..].to_string(); } - }); + + let mut builder = axum::http::Uri::builder(); + if let Some(scheme) = req.uri().scheme() { + builder = builder.scheme(scheme.clone()); + } + if let Some(authority) = req.uri().authority() { + builder = builder.authority(authority.clone()); + } + let new_uri = builder.path_and_query(path).build().unwrap(); + + let new_req = + axum::http::Request::builder().method(req.method()).uri(new_uri).body(req.into_body()).unwrap(); + + serve_dir.call(new_req) + } + }); let auth_router = axum::Router::new() .route("/token", axum::routing::post(token)) @@ -345,8 +340,11 @@ pub fn index_register_with_path(web_dir_path: &Path, web_ui_path: &str) -> axum: let web_ui_path_clone = web_ui_path.to_string(); axum::Router::new() .nest(&concat_path_leading_slash(web_ui_path, "auth"), auth_router) - .route(&format!("/{web_ui_path}"), axum::routing::get(|| async move { - axum::response::Redirect::permanent(&format!("/{web_ui_path_clone}/")) - })) + .route( + &format!("/{web_ui_path}"), + axum::routing::get( + || async move { axum::response::Redirect::permanent(&format!("/{web_ui_path_clone}/")) }, + ), + ) .nest(&format!("/{web_ui_path}/"), web_ui_router) } diff --git a/backend/src/api/endpoints/websocket_api.rs b/backend/src/api/endpoints/websocket_api.rs index 592ea9d13..817ea553d 100644 --- a/backend/src/api/endpoints/websocket_api.rs +++ b/backend/src/api/endpoints/websocket_api.rs @@ -1,13 +1,22 @@ -use crate::api::endpoints::v1_api::create_status_check; -use crate::api::model::AppState; -use crate::api::model::EventMessage; -use crate::auth::{verify_token_admin, verify_token_user}; -use axum::extract::ws::CloseFrame; -use axum::{extract::ws::{Message, WebSocket, WebSocketUpgrade},response::IntoResponse}; +use crate::{ + api::{ + endpoints::v1_api::create_status_check, + model::{AppState, EventMessage}, + }, + auth::{verify_token_admin, verify_token_user}, +}; +use axum::{ + extract::ws::{CloseFrame, Message, WebSocket, WebSocketUpgrade}, + response::IntoResponse, +}; use log::{error, trace}; -use shared::model::{ProtocolHandler, ProtocolHandlerMemory, ProtocolMessage, UserCommand, UserRole, WsCloseCode, PROTOCOL_VERSION}; +use shared::{ + model::{ + ProtocolHandler, ProtocolHandlerMemory, ProtocolMessage, UserCommand, UserRole, WsCloseCode, PROTOCOL_VERSION, + }, + utils::{concat_path_leading_slash, default_kick_secs}, +}; use std::sync::Arc; -use shared::utils::{concat_path_leading_slash, default_kick_secs}; // WebSocket upgrade handler async fn websocket_handler( @@ -29,15 +38,10 @@ async fn websocket_handler_auth( pub fn ws_api_register(web_auth_enabled: bool, web_ui_path: &str) -> axum::Router> { if web_auth_enabled { - axum::Router::new().route( - &concat_path_leading_slash(web_ui_path, "ws"), - axum::routing::get(websocket_handler_auth), - ) + axum::Router::new() + .route(&concat_path_leading_slash(web_ui_path, "ws"), axum::routing::get(websocket_handler_auth)) } else { - axum::Router::new().route( - &concat_path_leading_slash(web_ui_path, "ws"), - axum::routing::get(websocket_handler), - ) + axum::Router::new().route(&concat_path_leading_slash(web_ui_path, "ws"), axum::routing::get(websocket_handler)) } } @@ -62,17 +66,10 @@ fn get_secret_key(app_state: &AppState, auth: bool) -> Option> { return None; } - app_state - .app_config - .config - .load() - .web_ui - .as_ref() - .and_then(|c| c.auth.as_ref()) - .map(|c| { - let secret_key: &[u8] = c.secret.as_ref(); - secret_key.to_vec() - }) + app_state.app_config.config.load().web_ui.as_ref().and_then(|c| c.auth.as_ref()).map(|c| { + let secret_key: &[u8] = c.secret.as_ref(); + secret_key.to_vec() + }) } async fn handle_handshake(msg: Message, socket: &mut WebSocket, version: u8) -> Result<(), String> { @@ -80,10 +77,7 @@ async fn handle_handshake(msg: Message, socket: &mut WebSocket, version: u8) -> if bytes.len() == 1 { let client_version = bytes[0]; if client_version == version { - socket - .send(Message::binary(bytes)) - .await - .map_err(|e| e.to_string())?; + socket.send(Message::binary(bytes)).await.map_err(|e| e.to_string())?; return Ok(()); } error!("Protocol Version mismatch: server={version}, client={client_version}"); @@ -110,19 +104,19 @@ async fn handle_protocol_message( if let Message::Binary(bytes) = msg { match ProtocolMessage::from_bytes(bytes) { Ok(ProtocolMessage::Auth(auth_token)) => { - mem.token = None; - if !auth_required || verify_auth_admin_token(&auth_token, secret_key) { - mem.role = UserRole::Admin; - mem.token = Some(auth_token); - Some(ProtocolMessage::Authorized) - } else if verify_auth_user_token(&auth_token, secret_key) { - mem.role = UserRole::User; - mem.token = Some(auth_token); - Some(ProtocolMessage::Authorized) - } else { - Some(ProtocolMessage::Unauthorized) - } - }, + mem.token = None; + if !auth_required || verify_auth_admin_token(&auth_token, secret_key) { + mem.role = UserRole::Admin; + mem.token = Some(auth_token); + Some(ProtocolMessage::Authorized) + } else if verify_auth_user_token(&auth_token, secret_key) { + mem.role = UserRole::User; + mem.token = Some(auth_token); + Some(ProtocolMessage::Authorized) + } else { + Some(ProtocolMessage::Unauthorized) + } + } Ok(ProtocolMessage::StatusRequest(auth_token)) => { if !auth_required || verify_auth_admin_token(&auth_token, secret_key) { mem.role = UserRole::Admin; @@ -132,7 +126,7 @@ async fn handle_protocol_message( } else { Some(ProtocolMessage::Unauthorized) } - }, + } Ok(ProtocolMessage::UserAction(cmd)) => { if let Some(token) = mem.token.as_ref() { if !auth_required || verify_auth_admin_token(token, secret_key) { @@ -143,7 +137,7 @@ async fn handle_protocol_message( } else { Some(ProtocolMessage::UserActionResponse(false)) } - }, + } Ok(ProtocolMessage::ActiveProviderCountRequest(auth_token)) => { if !auth_required || verify_auth_admin_token(&auth_token, secret_key) { mem.role = UserRole::Admin; @@ -153,16 +147,14 @@ async fn handle_protocol_message( } else { Some(ProtocolMessage::Unauthorized) } - }, + } Ok(_) => { trace!("Unexpected protocol message after handshake"); None } Err(e) => { error!("Invalid websocket message: {e}"); - Some(ProtocolMessage::Error(format!( - "Invalid websocket message: {e}" - ))) + Some(ProtocolMessage::Error(format!("Invalid websocket message: {e}"))) } } } else { @@ -193,39 +185,31 @@ async fn handle_incoming_message( Some(protocol_msg) => { let bytes = match protocol_msg.to_bytes() { Ok(bytes) => bytes, - Err(err) => ProtocolMessage::Error(err.to_string()) - .to_bytes() - .map_err(|e| e.to_string())?, + Err(err) => ProtocolMessage::Error(err.to_string()).to_bytes().map_err(|e| e.to_string())?, }; - Ok(socket - .send(Message::Binary(bytes)) - .await - .map_err(|e| e.to_string())?) + Ok(socket.send(Message::Binary(bytes)).await.map_err(|e| e.to_string())?) } } } } } -async fn handle_event_message(socket: &mut WebSocket, event: EventMessage, handler: &ProtocolHandler) -> Result<(), String> { +async fn handle_event_message( + socket: &mut WebSocket, + event: EventMessage, + handler: &ProtocolHandler, +) -> Result<(), String> { match handler { - ProtocolHandler::Version(_) => {}, + ProtocolHandler::Version(_) => {} ProtocolHandler::Default(mem) => { if mem.role.is_admin() { match event { EventMessage::ServerError(error) => { - let msg = ProtocolMessage::ServerError(error) - .to_bytes() - .map_err(|e| e.to_string())?; - socket - .send(Message::Binary(msg)) - .await - .map_err(|e| format!("Server Error event: {e} "))?; + let msg = ProtocolMessage::ServerError(error).to_bytes().map_err(|e| e.to_string())?; + socket.send(Message::Binary(msg)).await.map_err(|e| format!("Server Error event: {e} "))?; } EventMessage::ActiveUser(event) => { - let msg = ProtocolMessage::ActiveUserResponse(event) - .to_bytes() - .map_err(|e| e.to_string())?; + let msg = ProtocolMessage::ActiveUserResponse(event).to_bytes().map_err(|e| e.to_string())?; socket .send(Message::Binary(msg)) .await @@ -241,22 +225,17 @@ async fn handle_event_message(socket: &mut WebSocket, event: EventMessage, handl .map_err(|e| format!("Provider connection change event: {e} "))?; } EventMessage::ConfigChange(config) => { - let msg = ProtocolMessage::ConfigChangeResponse(config) - .to_bytes() - .map_err(|e| e.to_string())?; + let msg = + ProtocolMessage::ConfigChangeResponse(config).to_bytes().map_err(|e| e.to_string())?; socket .send(Message::Binary(msg)) .await .map_err(|e| format!("Configuration files change event: {e} "))?; } EventMessage::PlaylistUpdate(state) => { - let msg = ProtocolMessage::PlaylistUpdateResponse(state) - .to_bytes() - .map_err(|e| e.to_string())?; - socket - .send(Message::Binary(msg)) - .await - .map_err(|e| format!("Playlist update event: {e} "))?; + let msg = + ProtocolMessage::PlaylistUpdateResponse(state).to_bytes().map_err(|e| e.to_string())?; + socket.send(Message::Binary(msg)).await.map_err(|e| format!("Playlist update event: {e} "))?; } EventMessage::PlaylistUpdateProgress(target, msg) => { let msg = ProtocolMessage::PlaylistUpdateProgressResponse(target, msg) @@ -268,13 +247,9 @@ async fn handle_event_message(socket: &mut WebSocket, event: EventMessage, handl .map_err(|e| format!("Playlist update progress event: {e} "))?; } EventMessage::SystemInfoUpdate(system_info) => { - let msg = ProtocolMessage::SystemInfoResponse(system_info) - .to_bytes() - .map_err(|e| e.to_string())?; - socket - .send(Message::Binary(msg)) - .await - .map_err(|e| format!("System info event: {e} "))?; + let msg = + ProtocolMessage::SystemInfoResponse(system_info).to_bytes().map_err(|e| e.to_string())?; + socket.send(Message::Binary(msg)).await.map_err(|e| format!("System info event: {e} "))?; } EventMessage::LibraryScanProgress(summary) => { let msg = ProtocolMessage::LibraryScanProgressResponse(summary) @@ -330,8 +305,9 @@ async fn handle_user_action(app_state: &Arc, cmd: UserCommand) -> bool match cmd { UserCommand::Kick(addr, virtual_id, _secs) => { // secs could be later used for different kick configurations. Currently, we only have 1. - let kick_secs = app_state.app_config.config.load().web_ui.as_ref().map_or_else(default_kick_secs, |wc| wc.kick_secs); + let kick_secs = + app_state.app_config.config.load().web_ui.as_ref().map_or_else(default_kick_secs, |wc| wc.kick_secs); app_state.connection_manager.kick_connection(&addr, virtual_id, kick_secs).await } } -} \ No newline at end of file +} diff --git a/backend/src/api/endpoints/xmltv_api.rs b/backend/src/api/endpoints/xmltv_api.rs index 1beaa3573..78c8fdc7d 100644 --- a/backend/src/api/endpoints/xmltv_api.rs +++ b/backend/src/api/endpoints/xmltv_api.rs @@ -1,26 +1,35 @@ -use crate::api::api_utils::{empty_json_response_as_array, get_user_target, get_user_target_by_credentials, internal_server_error, resource_response, stream_json_or_bin_response_stream, try_unwrap_body}; -use crate::api::model::AppState; -use crate::api::model::UserApiRequest; -use crate::model::{Config, EPG_ATTRIB_ID, EPG_TAG_CHANNEL}; -use crate::model::{ConfigTarget, ProxyUserCredentials, TargetOutput}; -use crate::repository::m3u_get_epg_file_path_for_target; -use crate::repository::storage_const; -use crate::repository::XML_PREAMBLE; -use crate::repository::{get_target_storage_path, BPlusTreeQuery, LockedReceiverStream}; -use crate::repository::{xtream_get_epg_file_path_for_target, xtream_get_storage_path}; -use crate::utils; -use crate::utils::{deobscure_text, file_exists_async, format_xmltv_time_utc, get_epg_processing_options, obscure_text, EpgProcessingOptions, EpgTimeShift}; +use crate::{ + api::{ + api_utils::{ + empty_json_response_as_array, get_user_target, get_user_target_by_credentials, internal_server_error, + resource_response, stream_json_or_bin_response_stream, try_unwrap_body, + }, + model::{AppState, UserApiRequest}, + }, + model::{Config, ConfigTarget, ProxyUserCredentials, TargetOutput, EPG_ATTRIB_ID, EPG_TAG_CHANNEL}, + repository::{ + get_target_storage_path, m3u_get_epg_file_path_for_target, storage_const, xtream_get_epg_file_path_for_target, + xtream_get_storage_path, BPlusTreeQuery, LockedReceiverStream, XML_PREAMBLE, + }, + utils, + utils::{ + deobscure_text, file_exists_async, format_xmltv_time_utc, get_epg_processing_options, obscure_text, + EpgProcessingOptions, EpgTimeShift, + }, +}; use axum::response::IntoResponse; use chrono::{DateTime, TimeZone}; use log::{error, trace}; use quick_xml::events::{BytesEnd, BytesStart, BytesText, Event}; -use shared::concat_string; -use shared::model::{EpgChannel, EpgProgramme, ShortEpgDto, ShortEpgResultDto}; -use std::path::{Path, PathBuf}; -use std::sync::Arc; -use tokio::io::AsyncWriteExt; -use tokio::sync::mpsc; -use tokio::task; +use shared::{ + concat_string, + model::{EpgChannel, EpgProgramme, ShortEpgDto, ShortEpgResultDto}, +}; +use std::{ + path::{Path, PathBuf}, + sync::Arc, +}; +use tokio::{io::AsyncWriteExt, sync::mpsc, task}; use tokio_util::io::ReaderStream; pub fn get_empty_epg_response() -> axum::response::Response { @@ -30,15 +39,11 @@ pub fn get_empty_epg_response() -> axum::response::Response { .body(axum::body::Body::from(r#""#))) } - fn get_epg_path_for_target_of_type(target_name: &str, epg_path: PathBuf) -> Option { if utils::path_exists(&epg_path) { return Some(epg_path); } - trace!( - "Can't find epg file for {target_name} target: {}", - epg_path.to_str().unwrap_or("?") - ); + trace!("Can't find epg file for {target_name} target: {}", epg_path.to_str().unwrap_or("?")); None } @@ -150,12 +155,21 @@ async fn serve_epg_with_rewrites( let epg_processing_options = get_epg_processing_options(app_state, user, target); - let base_url = if !matches!(epg_processing_options.time_shift, EpgTimeShift::None) || epg_processing_options.rewrite_urls { - let server_info = app_state.app_config.get_user_server_info(user); - Some(concat_string!(&server_info.get_base_url(), "/", storage_const::EPG_RESOURCE_PATH, "/", &user.username, "/", &user.password)) - } else { - None - }; + let base_url = + if !matches!(epg_processing_options.time_shift, EpgTimeShift::None) || epg_processing_options.rewrite_urls { + let server_info = app_state.app_config.get_user_server_info(user); + Some(concat_string!( + &server_info.get_base_url(), + "/", + storage_const::EPG_RESOURCE_PATH, + "/", + &user.username, + "/", + &user.password + )) + } else { + None + }; let limit = limit.unwrap_or_default(); @@ -179,10 +193,7 @@ async fn serve_epg_with_rewrites( }); tokio::spawn(async move { if let Err(err) = spawn_handle.await { - error!( - "EPG rewrite producer task failed for {}: {err}", - epg_path_for_log.display() - ); + error!("EPG rewrite producer task failed for {}: {err}", epg_path_for_log.display()); } }); @@ -193,7 +204,8 @@ async fn serve_epg_with_rewrites( error!("EPG: Failed to write xml header {err}"); return; } - if let Err(err) = tx.write_all(r#""#.as_bytes()).await { + if let Err(err) = tx.write_all(r#""#.as_bytes()).await + { error!("EPG: Failed to write xml tv header {err}"); return; } @@ -220,8 +232,11 @@ async fn serve_epg_with_rewrites( continue_on_err!(writer.write_event_async(Event::End(elem)).await); if let Some(icon_url) = &channel.icon { - let icon = match (epg_processing_options.rewrite_urls, base_url.as_ref(), - obscure_text(&epg_processing_options.encrypt_secret, icon_url)) { + let icon = match ( + epg_processing_options.rewrite_urls, + base_url.as_ref(), + obscure_text(&epg_processing_options.encrypt_secret, icon_url), + ) { (true, Some(base), Ok(enc)) => concat_string!(base, "/", &enc), _ => icon_url.to_string(), }; @@ -239,8 +254,14 @@ async fn serve_epg_with_rewrites( for programme in programmes { let mut elem = BytesStart::new("programme"); let (user_start, user_stop) = (programme.start, programme.stop); - elem.push_attribute(("start", format_xmltv_time_utc(user_start, &epg_processing_options.time_shift).as_str())); - elem.push_attribute(("stop", format_xmltv_time_utc(user_stop, &epg_processing_options.time_shift).as_str())); + elem.push_attribute(( + "start", + format_xmltv_time_utc(user_start, &epg_processing_options.time_shift).as_str(), + )); + elem.push_attribute(( + "stop", + format_xmltv_time_utc(user_stop, &epg_processing_options.time_shift).as_str(), + )); elem.push_attribute(("channel", channel.id.as_ref())); continue_on_err!(writer.write_event_async(Event::Start(elem)).await); @@ -276,8 +297,8 @@ async fn serve_epg_with_rewrites( let body_stream = ReaderStream::new(rx); try_unwrap_body!(axum::response::Response::builder() - .header(axum::http::header::CONTENT_TYPE, mime::TEXT_XML.to_string()) - .body(axum::body::Body::from_stream(body_stream))) + .header(axum::http::header::CONTENT_TYPE, mime::TEXT_XML.to_string()) + .body(axum::body::Body::from_stream(body_stream))) } async fn get_epg_channel(app_state: &Arc, channel_id: &Arc, epg_path: &Path) -> Option { @@ -301,7 +322,9 @@ async fn get_epg_channel(app_state: &Arc, channel_id: &Arc, epg_p None } } - }).await { + }) + .await + { Ok(result) => result, Err(err) => { error!("Failed to run epg query task: {err}"); @@ -318,9 +341,14 @@ fn format_xmltv_time(ts: i64) -> String { } } -fn get_applied_epg_timeshift(programme: &EpgProgramme, epg_processing_options: &EpgProcessingOptions) -> (String, String, i64, i64) { +fn get_applied_epg_timeshift( + programme: &EpgProgramme, + epg_processing_options: &EpgProcessingOptions, +) -> (String, String, i64, i64) { match &epg_processing_options.time_shift { - EpgTimeShift::None => (format_xmltv_time(programme.start), format_xmltv_time(programme.stop), programme.start, programme.stop), + EpgTimeShift::None => { + (format_xmltv_time(programme.start), format_xmltv_time(programme.stop), programme.start, programme.stop) + } EpgTimeShift::Fixed(m) => { let off = i64::from(*m) * 60; let s = programme.start + off; @@ -333,12 +361,22 @@ fn get_applied_epg_timeshift(programme: &EpgProgramme, epg_processing_options: & // We use the original timestamps (programme.start/stop) here because TimeZone adjustment // is only for the visual string representation. The absolute event time (UTC) remains unchanged. // Unlike 'Fixed' offset which artificially shifts the event time. - (s_dt.format("%Y-%m-%d %H:%M:%S").to_string(), e_dt.format("%Y-%m-%d %H:%M:%S").to_string(), programme.start, programme.stop) + ( + s_dt.format("%Y-%m-%d %H:%M:%S").to_string(), + e_dt.format("%Y-%m-%d %H:%M:%S").to_string(), + programme.start, + programme.stop, + ) } } } -fn from_programme(stream_id: &Arc, epg_id: &Arc, programme: &EpgProgramme, epg_processing_options: &EpgProcessingOptions) -> ShortEpgDto { +fn from_programme( + stream_id: &Arc, + epg_id: &Arc, + programme: &EpgProgramme, + epg_processing_options: &EpgProcessingOptions, +) -> ShortEpgDto { let (start_str, end_str, start_ts, stop_ts) = get_applied_epg_timeshift(programme, epg_processing_options); ShortEpgDto { @@ -371,15 +409,23 @@ pub async fn serve_short_epg( ) -> axum::response::Response { let short_epg = { // It seems provider set limit to 4 if it is undefined oor 0. - let limit = if limit > 0 { limit} else { DEFAULT_SHORT_EPG_LIMIT }; + let limit = if limit > 0 { limit } else { DEFAULT_SHORT_EPG_LIMIT }; if file_exists_async(epg_path).await { if let Some(epg_channel) = get_epg_channel(app_state, channel_id, epg_path).await { let epg_processing_options = get_epg_processing_options(app_state, user, target); ShortEpgResultDto { epg_listings: if limit > 0 { - epg_channel.get_programme_with_limit(limit).iter().map(|p| from_programme(&stream_id, channel_id, p, &epg_processing_options)).collect() + epg_channel + .get_programme_with_limit(limit) + .iter() + .map(|p| from_programme(&stream_id, channel_id, p, &epg_processing_options)) + .collect() } else { - epg_channel.programmes.iter().map(|p| from_programme(&stream_id, channel_id, p, &epg_processing_options)).collect() + epg_channel + .programmes + .iter() + .map(|p| from_programme(&stream_id, channel_id, p, &epg_processing_options)) + .collect() }, } } else { @@ -391,11 +437,10 @@ pub async fn serve_short_epg( }; match serde_json::to_string(&short_epg) { - Ok(json) => ( - axum::http::StatusCode::OK, - [(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string())], - json - ).into_response(), + Ok(json) => { + (axum::http::StatusCode::OK, [(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string())], json) + .into_response() + } Err(_) => internal_server_error!(), } } @@ -436,23 +481,18 @@ async fn xmltv_api( async fn epg_api_resource( req_headers: axum::http::HeaderMap, axum::extract::Query(api_req): axum::extract::Query, - axum::extract::Path((username, password, resource)): axum::extract::Path<( - String, - String, - String, - )>, + axum::extract::Path((username, password, resource)): axum::extract::Path<(String, String, String)>, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse + Send { - let Some((user, _target)) = - get_user_target_by_credentials(&username, &password, &api_req, &app_state) - else { + let Some((user, _target)) = get_user_target_by_credentials(&username, &password, &api_req, &app_state) else { return axum::http::StatusCode::BAD_REQUEST.into_response(); }; if user.permission_denied(&app_state) { return axum::http::StatusCode::FORBIDDEN.into_response(); } - let encrypt_secret = app_state.app_config.get_reverse_proxy_rewrite_secret().unwrap_or_else(|| app_state.app_config.encrypt_secret); + let encrypt_secret = + app_state.app_config.get_reverse_proxy_rewrite_secret().unwrap_or_else(|| app_state.app_config.encrypt_secret); if let Ok(resource_url) = deobscure_text(&encrypt_secret, &resource) { resource_response(&app_state, &resource_url, &req_headers, None).await.into_response() } else { @@ -475,7 +515,8 @@ pub fn xmltv_api_register() -> axum::Router> { .route("/xmltv.php", axum::routing::get(xmltv_api)) .route("/update/epg.php", axum::routing::get(xmltv_api)) .route("/epg", axum::routing::get(xmltv_api)) - .route(&format!("/{}/{{username}}/{{password}}/{{resource}}", storage_const::EPG_RESOURCE_PATH), - axum::routing::get(epg_api_resource), + .route( + &format!("/{}/{{username}}/{{password}}/{{resource}}", storage_const::EPG_RESOURCE_PATH), + axum::routing::get(epg_api_resource), ) } diff --git a/backend/src/api/endpoints/xtream_api.rs b/backend/src/api/endpoints/xtream_api.rs index 8e088b3db..9feac8f5c 100644 --- a/backend/src/api/endpoints/xtream_api.rs +++ b/backend/src/api/endpoints/xtream_api.rs @@ -1,40 +1,64 @@ // https://github.com/tellytv/go.xtream-codes/blob/master/structs.go // Xtream api -> https://9tzx6f0ozj.apidog.io/ -use crate::api::api_utils; -use crate::api::api_utils::{create_api_proxy_user, create_session_fingerprint, empty_json_response_as_array, empty_json_response_as_object, force_provider_stream_response, get_user_target, get_user_target_by_credentials, internal_server_error, is_seek_request, local_stream_response, redirect, redirect_response, resource_response, separate_number_and_remainder, stream_response, try_option_bad_request, try_result_bad_request, try_result_not_found, try_unwrap_body, RedirectParams}; -use crate::api::endpoints::hls_api::handle_hls_stream_request; -use crate::api::endpoints::xmltv_api::{get_empty_epg_response, get_epg_path_for_target, serve_short_epg}; -use crate::api::model::AppState; -use crate::api::model::UserApiRequest; -use crate::api::model::XtreamAuthorizationResponse; -use crate::api::model::{create_custom_video_stream_response, CustomVideoStreamType}; -use crate::auth::Fingerprint; -use crate::model::{xtream_mapping_option_from_target_options, ConfigTarget}; -use crate::model::{Config, ConfigInput, ConfigInputFlags}; -use crate::model::{InputSource, ProxyUserCredentials}; -use crate::repository::get_target_storage_path; -use crate::repository::storage_const; -use crate::repository::VirtualIdRecord; -use crate::repository::{get_target_id_mapping, user_get_bouquet_filter, xtream_get_collection_path, xtream_get_item_for_stream_id, xtream_load_rewrite_playlist}; -use crate::utils::xtream::create_vod_info_from_item; -use crate::utils::{apply_timeshift, debug_if_enabled, file_exists_async, parse_timeshift, trace_if_enabled}; -use crate::utils::{request, xtream}; -use axum::http::HeaderMap; -use axum::response::IntoResponse; +use crate::{ + api::{ + api_utils, + api_utils::{ + create_api_proxy_user, create_session_fingerprint, empty_json_response_as_array, + empty_json_response_as_object, force_provider_stream_response, get_user_target, + get_user_target_by_credentials, internal_server_error, is_seek_request, local_stream_response, redirect, + redirect_response, resource_response, separate_number_and_remainder, stream_response, + try_option_bad_request, try_result_bad_request, try_result_not_found, try_unwrap_body, RedirectParams, + }, + endpoints::{ + hls_api::handle_hls_stream_request, + xmltv_api::{get_empty_epg_response, get_epg_path_for_target, serve_short_epg}, + }, + model::{ + create_custom_video_stream_response, AppState, CustomVideoStreamType, UserApiRequest, + XtreamAuthorizationResponse, + }, + }, + auth::Fingerprint, + model::{ + xtream_mapping_option_from_target_options, Config, ConfigInput, ConfigInputFlags, ConfigTarget, InputSource, + ProxyUserCredentials, + }, + repository::{ + get_target_id_mapping, get_target_storage_path, storage_const, user_get_bouquet_filter, + xtream_get_collection_path, xtream_get_item_for_stream_id, xtream_load_rewrite_playlist, VirtualIdRecord, + }, + utils::{ + apply_timeshift, debug_if_enabled, file_exists_async, parse_timeshift, request, trace_if_enabled, xtream, + xtream::create_vod_info_from_item, + }, +}; +use axum::{http::HeaderMap, response::IntoResponse}; use bytes::Bytes; -use futures::stream::{self, StreamExt}; -use futures::Stream; +use futures::{ + stream::{self, StreamExt}, + Stream, +}; use log::{debug, error, warn}; use serde::{Deserialize, Serialize}; use serde_json::{json, Map, Value}; -use shared::concat_string; -use shared::error::{info_err, info_err_res, TuliproxError}; -use shared::model::{create_stream_channel_with_type, PlaylistEntry, PlaylistItemType, ProxyType, ShortEpgResultDto, TargetType, UserConnectionPermission, XtreamCluster, XtreamPlaylistItem}; -use shared::utils::{deserialize_as_string, extract_extension_from_url, generate_playlist_uuid, sanitize_sensitive_info, trim_slash, Internable, HLS_EXT}; -use std::fmt::Write; -use std::fmt::{Display, Formatter}; -use std::str::FromStr; -use std::sync::Arc; +use shared::{ + concat_string, + error::{info_err, info_err_res, TuliproxError}, + model::{ + create_stream_channel_with_type, PlaylistEntry, PlaylistItemType, ProxyType, ShortEpgResultDto, TargetType, + UserConnectionPermission, XtreamCluster, XtreamPlaylistItem, + }, + utils::{ + deserialize_as_string, extract_extension_from_url, generate_playlist_uuid, sanitize_sensitive_info, trim_slash, + Internable, HLS_EXT, + }, +}; +use std::{ + fmt::{Display, Formatter, Write}, + str::FromStr, + sync::Arc, +}; #[derive(Serialize, Deserialize, Debug, Copy, Clone, Eq, PartialEq)] pub enum ApiStreamContext { @@ -54,13 +78,15 @@ impl ApiStreamContext { impl Display for ApiStreamContext { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", - match self { - Self::Live | Self::LiveAlt => Self::LIVE, - Self::Movie => Self::MOVIE, - Self::Series => Self::SERIES, - Self::Timeshift => Self::TIMESHIFT, - } + write!( + f, + "{}", + match self { + Self::Live | Self::LiveAlt => Self::LIVE, + Self::Movie => Self::MOVIE, + Self::Series => Self::SERIES, + Self::Timeshift => Self::TIMESHIFT, + } ) } } @@ -108,14 +134,7 @@ impl<'a> ApiStreamRequest<'a> { stream_id: &'a str, action_path: &'a str, ) -> Self { - Self { - context, - access_token: false, - username, - password, - stream_id, - action_path, - } + Self { context, access_token: false, username, password, stream_id, action_path } } pub const fn from_access_token( context: ApiStreamContext, @@ -123,18 +142,10 @@ impl<'a> ApiStreamRequest<'a> { stream_id: &'a str, action_path: &'a str, ) -> Self { - Self { - context, - access_token: false, - username: "", - password, - stream_id, - action_path, - } + Self { context, access_token: false, username: "", password, stream_id, action_path } } } - #[derive(Serialize, Deserialize)] struct XtreamCategoryEntry { #[serde(deserialize_with = "deserialize_as_string")] @@ -156,9 +167,7 @@ pub(in crate::api) fn get_xtream_player_api_stream_url( let use_prefix = input.has_flag(ConfigInputFlags::XtreamLiveStreamUsePrefix); String::from(if use_prefix { "live" } else { "" }) } - ApiStreamContext::Movie | ApiStreamContext::Series | ApiStreamContext::Timeshift => { - context.to_string() - } + ApiStreamContext::Movie | ApiStreamContext::Series | ApiStreamContext::Timeshift => context.to_string(), }; let mut parts = vec![ trim_slash(&input_user_info.base_url), @@ -195,9 +204,8 @@ async fn xtream_player_api_stream( app_state: &Arc, api_req: &UserApiRequest, stream_req: ApiStreamRequest<'_>, - user_target: Option<(ProxyUserCredentials, Arc)> + user_target: Option<(ProxyUserCredentials, Arc)>, ) -> impl IntoResponse + Send { - // if log::log_enabled!(log::Level::Debug) { // debug!( // "Stream request ctx={} user={} stream_id={} action_path={}", @@ -214,22 +222,35 @@ async fn xtream_player_api_stream( let (user, target) = match user_target { None => try_option_bad_request!( - get_user_target_by_credentials( stream_req.username, stream_req.password, api_req, app_state), - false, - format!("Could not find any user for xc stream {}", stream_req.username)), - Some((user, target)) => (user, target) + get_user_target_by_credentials(stream_req.username, stream_req.password, api_req, app_state), + false, + format!("Could not find any user for xc stream {}", stream_req.username) + ), + Some((user, target)) => (user, target), }; let _guard = app_state.app_config.file_locks.write_lock_str(&user.username).await; if user.permission_denied(app_state) { - return create_custom_video_stream_response(app_state, &fingerprint.addr, CustomVideoStreamType::UserAccountExpired).await.into_response(); + return create_custom_video_stream_response( + app_state, + &fingerprint.addr, + CustomVideoStreamType::UserAccountExpired, + ) + .await + .into_response(); } let target_name = &target.name; if !target.has_output(TargetType::Xtream) { debug!("Target has no xtream codes playlist {target_name}"); - return create_custom_video_stream_response(app_state, &fingerprint.addr, CustomVideoStreamType::ChannelUnavailable).await.into_response(); + return create_custom_video_stream_response( + app_state, + &fingerprint.addr, + CustomVideoStreamType::ChannelUnavailable, + ) + .await + .into_response(); } let (action_stream_id, stream_ext) = separate_number_and_remainder(stream_req.stream_id); @@ -245,9 +266,12 @@ async fn xtream_player_api_stream( } let input = try_option_bad_request!( - app_state.app_config.get_input_by_name(&pli.input_name), - true, - format!( "Can't find input {} for target {target_name}, context {}, stream_id {virtual_id}", pli.input_name, stream_req.context) + app_state.app_config.get_input_by_name(&pli.input_name), + true, + format!( + "Can't find input {} for target {target_name}, context {}, stream_id {virtual_id}", + pli.input_name, stream_req.context + ) ); if pli.item_type.is_local() { @@ -262,7 +286,9 @@ async fn xtream_player_api_stream( &user, connection_permission, true, - ).await.into_response(); + ) + .await + .into_response(); } let (cluster, item_type) = if stream_req.context == ApiStreamContext::Timeshift { @@ -275,29 +301,27 @@ async fn xtream_player_api_stream( "ID chain for xtream endpoint: request_stream_id={} -> action_stream_id={action_stream_id} -> req_virtual_id={req_virtual_id} -> virtual_id={virtual_id}", stream_req.stream_id); let session_key = create_session_fingerprint(fingerprint, &user.username, virtual_id); - let user_session = app_state - .active_users - .get_and_update_user_session(&user.username, &session_key).await; + let user_session = app_state.active_users.get_and_update_user_session(&user.username, &session_key).await; let session_url = if let Some(session) = &user_session { if session.permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response( - app_state, &fingerprint.addr, + app_state, + &fingerprint.addr, CustomVideoStreamType::UserConnectionsExhausted, - ).await - .into_response(); + ) + .await + .into_response(); } - if app_state - .active_provider - .is_over_limit(&session.provider) - .await - { + if app_state.active_provider.is_over_limit(&session.provider).await { return create_custom_video_stream_response( - app_state, &fingerprint.addr, + app_state, + &fingerprint.addr, CustomVideoStreamType::ProviderConnectionsExhausted, - ).await - .into_response(); + ) + .await + .into_response(); } let stream_channel = create_stream_channel_with_type(target.id, &pli, item_type); @@ -313,8 +337,8 @@ async fn xtream_player_api_stream( &input, &user, ) - .await - .into_response(); + .await + .into_response(); } session.stream_url.clone() @@ -325,10 +349,12 @@ async fn xtream_player_api_stream( let connection_permission = user.connection_permission(app_state).await; if connection_permission == UserConnectionPermission::Exhausted { return create_custom_video_stream_response( - app_state, &fingerprint.addr, + app_state, + &fingerprint.addr, CustomVideoStreamType::UserConnectionsExhausted, - ).await - .into_response(); + ) + .await + .into_response(); } let context = stream_req.context; @@ -360,9 +386,8 @@ async fn xtream_player_api_stream( ) ); - let is_hls_request = item_type == PlaylistItemType::LiveHls - || item_type == PlaylistItemType::LiveDash - || extension == HLS_EXT; + let is_hls_request = + item_type == PlaylistItemType::LiveHls || item_type == PlaylistItemType::LiveDash || extension == HLS_EXT; // Reverse proxy mode if is_hls_request { return handle_hls_stream_request( @@ -376,8 +401,8 @@ async fn xtream_player_api_stream( req_headers, connection_permission, ) - .await - .into_response(); + .await + .into_response(); } let stream_channel = create_stream_channel_with_type(target.id, &pli, item_type); @@ -394,13 +419,22 @@ async fn xtream_player_api_stream( &user, connection_permission, ) - .await - .into_response() + .await + .into_response() } -fn get_query_path(action_path: &str, stream_ext: Option<&String>, pli: &XtreamPlaylistItem, app_state: &Arc) -> (String, String) { +fn get_query_path( + action_path: &str, + stream_ext: Option<&String>, + pli: &XtreamPlaylistItem, + app_state: &Arc, +) -> (String, String) { let discard_extension = if pli.item_type.is_live() { - app_state.app_config.sources.load().get_input_by_name(&pli.input_name) + app_state + .app_config + .sources + .load() + .get_input_by_name(&pli.input_name) .as_ref() .is_some_and(|i| i.has_flag(ConfigInputFlags::XtreamLiveStreamWithoutExtension)) } else { @@ -444,12 +478,7 @@ async fn xtream_player_api_stream_with_token( let (action_stream_id, stream_ext) = separate_number_and_remainder(stream_req.stream_id); let req_virtual_id: u32 = try_result_bad_request!(action_stream_id.trim().parse()); let pli = try_result_bad_request!( - xtream_get_item_for_stream_id( - req_virtual_id, - app_state, - &target, - None - ).await, + xtream_get_item_for_stream_id(req_virtual_id, app_state, &target, None).await, true, format!("Failed to read xtream item for stream id {req_virtual_id}") ); @@ -476,13 +505,14 @@ async fn xtream_player_api_stream_with_token( &user, UserConnectionPermission::Allowed, true, - ).await.into_response(); + ) + .await + .into_response(); } let session_key = create_session_fingerprint(fingerprint, "webui", virtual_id); - let is_hls_request = - pli.item_type == PlaylistItemType::LiveHls || stream_ext.as_deref() == Some(HLS_EXT); + let is_hls_request = pli.item_type == PlaylistItemType::LiveHls || stream_ext.as_deref() == Some(HLS_EXT); // TODO how should we use fixed provider for hls in multi provider config? @@ -499,19 +529,14 @@ async fn xtream_player_api_stream_with_token( req_headers, UserConnectionPermission::Allowed, ) - .await - .into_response(); + .await + .into_response(); } let (query_path, _extension) = get_query_path(stream_req.action_path, stream_ext.as_ref(), &pli, app_state); let stream_url = try_option_bad_request!( - get_xtream_player_api_stream_url( - &input, - stream_req.context, - &query_path, - &pli.url - ), + get_xtream_player_api_stream_url(&input, stream_req.context, &query_path, &pli.url), true, format!( "Can't find stream url for target {target_name}, context {}, stream_id {}", @@ -519,10 +544,7 @@ async fn xtream_player_api_stream_with_token( ) ); - trace_if_enabled!( - "Streaming stream request from {}", - sanitize_sensitive_info(&stream_url) - ); + trace_if_enabled!("Streaming stream request from {}", sanitize_sensitive_info(&stream_url)); stream_response( fingerprint, app_state, @@ -535,8 +557,8 @@ async fn xtream_player_api_stream_with_token( &user, UserConnectionPermission::Allowed, ) - .await - .into_response() + .await + .into_response() } else { axum::http::StatusCode::BAD_REQUEST.into_response() } @@ -549,17 +571,9 @@ async fn xtream_player_api_resource( resource_req: ApiStreamRequest<'_>, ) -> impl IntoResponse { let (user, target) = try_option_bad_request!( - get_user_target_by_credentials( - resource_req.username, - resource_req.password, - api_req, - app_state - ), + get_user_target_by_credentials(resource_req.username, resource_req.password, api_req, app_state), false, - format!( - "Could not find any user xc resource {}", - resource_req.username - ) + format!("Could not find any user xc resource {}", resource_req.username) ); if user.permission_denied(app_state) { return axum::http::StatusCode::FORBIDDEN.into_response(); @@ -572,12 +586,7 @@ async fn xtream_player_api_resource( let req_virtual_id: u32 = try_result_bad_request!(resource_req.stream_id.trim().parse()); let resource = resource_req.action_path.trim(); let pli = try_result_bad_request!( - xtream_get_item_for_stream_id( - req_virtual_id, - app_state, - &target, - None - ).await, + xtream_get_item_for_stream_id(req_virtual_id, app_state, &target, None).await, true, format!("Failed to read xtream item for stream id {req_virtual_id}") ); @@ -588,10 +597,7 @@ async fn xtream_player_api_resource( None => axum::http::StatusCode::NOT_FOUND.into_response(), Some(url) => { if user.proxy.is_redirect(pli.item_type) || target.is_force_redirect(pli.item_type) { - trace_if_enabled!( - "Redirecting resource request to {}", - sanitize_sensitive_info(&url) - ); + trace_if_enabled!("Redirecting resource request to {}", sanitize_sensitive_info(&url)); redirect(&url).into_response() } else { trace_if_enabled!("Resource request to {}", sanitize_sensitive_info(&url)); @@ -606,11 +612,7 @@ macro_rules! create_xtream_player_api_stream { async fn $fn_name( fingerprint: Fingerprint, req_headers: HeaderMap, - axum::extract::Path((username, password, stream_id)): axum::extract::Path<( - String, - String, - String, - )>, + axum::extract::Path((username, password, stream_id)): axum::extract::Path<(String, String, String)>, axum::extract::State(app_state): axum::extract::State>, axum::extract::Query(api_req): axum::extract::Query, ) -> impl IntoResponse + Send { @@ -691,23 +693,20 @@ async fn xtream_player_api_timeshift_stream( ) -> impl IntoResponse + Send { let username = get_non_empty(×hift_request.username, &api_req.username, &api_form_req.username).to_string(); let password = get_non_empty(×hift_request.password, &api_req.password, &api_form_req.password).to_string(); - let stream_id = get_non_empty(×hift_request.stream_id, &api_req.stream_id, &api_form_req.stream_id).to_string(); + let stream_id = + get_non_empty(×hift_request.stream_id, &api_req.stream_id, &api_form_req.stream_id).to_string(); let duration = get_non_empty(×hift_request.duration, &api_req.duration, &api_form_req.duration); let start_time = get_non_empty(×hift_request.start, &api_req.start, &api_form_req.start); let (user, target) = try_option_bad_request!( - get_user_target_by_credentials( &username, &password, &api_form_req, &app_state), - false, - format!("Could not find any user {username}") - ); + get_user_target_by_credentials(&username, &password, &api_form_req, &app_state), + false, + format!("Could not find any user {username}") + ); let epg_timeshift = parse_timeshift(user.epg_request_timeshift.as_deref()); let start = apply_timeshift(start_time, &epg_timeshift); - let action_path = if start.is_empty() { - format!("{duration}/{start_time}") - } else { - format!("{duration}/{start}") - }; + let action_path = if start.is_empty() { format!("{duration}/{start_time}") } else { format!("{duration}/{start}") }; api_req.username.clone_from(&username); api_req.password.clone_from(&password); @@ -718,17 +717,11 @@ async fn xtream_player_api_timeshift_stream( &req_headers, &app_state, &api_req, - ApiStreamRequest::from( - ApiStreamContext::Timeshift, - &username, - &password, - &stream_id, - &action_path, - ), + ApiStreamRequest::from(ApiStreamContext::Timeshift, &username, &password, &stream_id, &action_path), Some((user, target)), ) - .await - .into_response() + .await + .into_response() } async fn xtream_player_api_timeshift_query_stream( @@ -754,35 +747,25 @@ async fn xtream_player_api_timeshift_query_stream( } let (user, target) = try_option_bad_request!( - get_user_target_by_credentials( username, password, &api_query_req, &app_state), - false, - format!("Could not find any user {username}") - ); + get_user_target_by_credentials(username, password, &api_query_req, &app_state), + false, + format!("Could not find any user {username}") + ); let epg_timeshift = parse_timeshift(user.epg_request_timeshift.as_deref()); let start = apply_timeshift(start_time, &epg_timeshift); - let action_path = if start.is_empty() { - format!("{duration}/{start_time}") - } else { - format!("{duration}/{start}") - }; + let action_path = if start.is_empty() { format!("{duration}/{start_time}") } else { format!("{duration}/{start}") }; xtream_player_api_stream( &fingerprint, &req_headers, &app_state, &api_query_req, - ApiStreamRequest::from( - ApiStreamContext::Timeshift, - username, - password, - stream_id, - &action_path, - ), + ApiStreamRequest::from(ApiStreamContext::Timeshift, username, password, stream_id, &action_path), Some((user, target)), ) - .await - .into_response() + .await + .into_response() } pub async fn xtream_get_stream_info_response( @@ -797,19 +780,23 @@ pub async fn xtream_get_stream_info_response( Err(_) => return try_unwrap_body!(empty_json_response_as_array()), }; - if let Ok(pli) = xtream_get_item_for_stream_id( - virtual_id, - app_state, - target, - Some(cluster), - ).await { + if let Ok(pli) = xtream_get_item_for_stream_id(virtual_id, app_state, target, Some(cluster)).await { if pli.item_type.is_local() { - let Ok(xtream_output) = target.get_xtream_output().ok_or_else(|| info_err!("Unexpected: xtream output required for target {}", target.name)) else { + let Ok(xtream_output) = target + .get_xtream_output() + .ok_or_else(|| info_err!("Unexpected: xtream output required for target {}", target.name)) + else { return try_unwrap_body!(empty_json_response_as_array()); }; let server_info = app_state.app_config.get_user_server_info(user); - let options = xtream_mapping_option_from_target_options(target, xtream_output, &app_state.app_config, user, Some(server_info.get_base_url().as_str())); + let options = xtream_mapping_option_from_target_options( + target, + xtream_output, + &app_state.app_config, + user, + Some(server_info.get_base_url().as_str()), + ); return axum::Json(pli.to_info_document(&options)).into_response(); } @@ -829,14 +816,12 @@ pub async fn xtream_get_stream_info_response( &pli, info_url.as_str(), cluster, - ).await + ) + .await { return try_unwrap_body!(axum::response::Response::builder() .status(axum::http::StatusCode::OK) - .header( - axum::http::header::CONTENT_TYPE, - mime::APPLICATION_JSON.to_string() - ) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) .body(axum::body::Body::from(content))); } } @@ -845,14 +830,10 @@ pub async fn xtream_get_stream_info_response( return match cluster { XtreamCluster::Video => { - let content = - create_vod_info_from_item(target, user, &pli); + let content = create_vod_info_from_item(target, user, &pli); try_unwrap_body!(axum::response::Response::builder() .status(axum::http::StatusCode::OK) - .header( - axum::http::header::CONTENT_TYPE, - mime::APPLICATION_JSON.to_string() - ) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) .body(axum::body::Body::from(content))) } XtreamCluster::Live | XtreamCluster::Series => { @@ -877,31 +858,31 @@ async fn xtream_get_short_epg( Err(_) => return get_empty_epg_response().into_response(), }; - if let Ok(pli) = xtream_get_item_for_stream_id( - virtual_id, - app_state, - target, - None, - ).await { + if let Ok(pli) = xtream_get_item_for_stream_id(virtual_id, app_state, target, None).await { let config = &app_state.app_config.config.load(); if let (Some(epg_path), Some(channel_id)) = (get_epg_path_for_target(config, target), &pli.epg_channel_id) { if file_exists_async(&epg_path).await { - return serve_short_epg(app_state, epg_path.as_path(), user, target, channel_id, stream_id.intern(), limit).await; + return serve_short_epg( + app_state, + epg_path.as_path(), + user, + target, + channel_id, + stream_id.intern(), + limit, + ) + .await; } } if pli.provider_id > 0 { let input_name = &pli.input_name; if let Some(input) = app_state.app_config.get_input_by_name(input_name) { - if let Some(action_url) = xtream::get_xtream_player_api_action_url( - &input, - crate::model::XC_ACTION_GET_SHORT_EPG, - ) { - let mut info_url = format!( - "{action_url}&{}={}", - crate::model::XC_TAG_STREAM_ID, - pli.provider_id - ); + if let Some(action_url) = + xtream::get_xtream_player_api_action_url(&input, crate::model::XC_ACTION_GET_SHORT_EPG) + { + let mut info_url = + format!("{action_url}&{}={}", crate::model::XC_TAG_STREAM_ID, pli.provider_id); if limit > 0 { info_url = format!("{info_url}&limit={limit}"); } @@ -918,14 +899,11 @@ async fn xtream_get_short_epg( None, false, ) - .await + .await { Ok((content, _)) => ( axum::http::StatusCode::OK, - [( - axum::http::header::CONTENT_TYPE.to_string(), - mime::APPLICATION_JSON.to_string(), - )], + [(axum::http::header::CONTENT_TYPE.to_string(), mime::APPLICATION_JSON.to_string())], content, ) .into_response(), @@ -960,13 +938,8 @@ async fn xtream_player_api_handle_content_action( if let Ok(file_path) = xtream_get_collection_path(config, target_name, collection) { match tokio::fs::read_to_string(&file_path).await { Ok(content) => { - let filter = user_get_bouquet_filter( - config, - &user.username, - category_id, - TargetType::Xtream, - cluster, - ).await; + let filter = + user_get_bouquet_filter(config, &user.username, category_id, TargetType::Xtream, cluster).await; match serde_json::from_str::>(&content) { Ok(mut categories) => { @@ -1000,16 +973,11 @@ async fn xtream_get_catchup_response( return axum::Json(json!(ShortEpgResultDto::default())).into_response(); }; - let pli = try_result_bad_request!(xtream_get_item_for_stream_id( - req_virtual_id, - app_state, - target, - Some(XtreamCluster::Live) - ).await); + let pli = try_result_bad_request!( + xtream_get_item_for_stream_id(req_virtual_id, app_state, target, Some(XtreamCluster::Live)).await + ); - let input = try_option_bad_request!(app_state - .app_config - .get_input_by_name(&pli.input_name)); + let input = try_option_bad_request!(app_state.app_config.get_input_by_name(&pli.input_name)); let mut info_url = try_option_bad_request!(xtream::get_xtream_player_api_action_url( &input, @@ -1038,26 +1006,18 @@ async fn xtream_get_catchup_response( ); let mut doc: Map = try_result_bad_request!(serde_json::from_str(&content)); - let epg_listings = try_option_bad_request!(doc - .get_mut(crate::model::XC_TAG_EPG_LISTINGS) - .and_then(Value::as_array_mut)); + let epg_listings = + try_option_bad_request!(doc.get_mut(crate::model::XC_TAG_EPG_LISTINGS).and_then(Value::as_array_mut)); // Collect data and generate UUIDs without holding the lock. let mut tasks = Vec::new(); let pli_uuid_str = pli.get_uuid().to_string(); for (idx, epg_list_item) in epg_listings.iter().enumerate() { - if let Some(cp_id) = epg_list_item - .get(crate::model::XC_TAG_ID) - .and_then(Value::as_str) - .and_then(|id| id.parse::().ok()) + if let Some(cp_id) = + epg_list_item.get(crate::model::XC_TAG_ID).and_then(Value::as_str).and_then(|id| id.parse::().ok()) { - let uuid = generate_playlist_uuid( - &pli_uuid_str, - &cp_id.to_string(), - pli.item_type, - &pli.input_name, - ); + let uuid = generate_playlist_uuid(&pli_uuid_str, &cp_id.to_string(), pli.item_type, &pli.input_name); tasks.push((idx, uuid, cp_id)); } } @@ -1070,11 +1030,9 @@ async fn xtream_get_catchup_response( if !tasks.is_empty() { { - let Ok((mut target_id_mapping, file_lock)) = get_target_id_mapping( - &app_state.app_config, - &target_path, - target.use_memory_cache, - ).await else { + let Ok((mut target_id_mapping, file_lock)) = + get_target_id_mapping(&app_state.app_config, &target_path, target.use_memory_cache).await + else { return internal_server_error!(); }; @@ -1089,15 +1047,13 @@ async fn xtream_get_catchup_response( mapping_results.push((idx, virtual_id)); if target.use_memory_cache { - in_memory_updates.push( - VirtualIdRecord::new( - cp_id, - virtual_id, - PlaylistItemType::Catchup, - pli.provider_id, - uuid, - ), - ); + in_memory_updates.push(VirtualIdRecord::new( + cp_id, + virtual_id, + PlaylistItemType::Catchup, + pli.provider_id, + uuid, + )); } } @@ -1114,10 +1070,7 @@ async fn xtream_get_catchup_response( // Apply the new virtual IDs back to the JSON document for (idx, v_id) in mapping_results { if let Some(item) = epg_listings.get_mut(idx).and_then(Value::as_object_mut) { - item.insert( - crate::model::XC_TAG_ID.to_string(), - Value::String(v_id.to_string()), - ); + item.insert(crate::model::XC_TAG_ID.to_string(), Value::String(v_id.to_string())); } } @@ -1130,10 +1083,7 @@ async fn xtream_get_catchup_response( |result| { try_unwrap_body!(axum::response::Response::builder() .status(axum::http::StatusCode::OK) - .header( - axum::http::header::CONTENT_TYPE, - mime::APPLICATION_JSON.to_string() - ) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) .body(result)) }, ) @@ -1159,10 +1109,7 @@ macro_rules! skip_flag_optional { } #[allow(clippy::too_many_lines)] -async fn xtream_player_api( - api_req: UserApiRequest, - app_state: &Arc, -) -> impl IntoResponse + Send { +async fn xtream_player_api(api_req: UserApiRequest, app_state: &Arc) -> impl IntoResponse + Send { let user_target = get_user_target(&api_req, app_state); if let Some((user, target)) = user_target { if !target.has_output(TargetType::Xtream) { @@ -1224,13 +1171,7 @@ async fn xtream_player_api( ); } crate::model::XC_ACTION_GET_EPG | crate::model::XC_ACTION_GET_SHORT_EPG => { - return xtream_get_short_epg( - app_state, - &user, - &target, - api_req.stream_id.trim(), - api_req.get_limit(), - ) + return xtream_get_short_epg(app_state, &user, &target, api_req.stream_id.trim(), api_req.get_limit()) .await .into_response(); } @@ -1259,43 +1200,27 @@ async fn xtream_player_api( action, category_id, &user, - ).await { + ) + .await + { return response.into_response(); } let result = match action { crate::model::XC_ACTION_GET_LIVE_STREAMS => skip_flag_optional!( skip_live, - xtream_load_rewrite_playlist( - XtreamCluster::Live, - &app_state.app_config, - &target, - category_id, - &user - ) - .await + xtream_load_rewrite_playlist(XtreamCluster::Live, &app_state.app_config, &target, category_id, &user) + .await ), crate::model::XC_ACTION_GET_VOD_STREAMS => skip_flag_optional!( skip_vod, - xtream_load_rewrite_playlist( - XtreamCluster::Video, - &app_state.app_config, - &target, - category_id, - &user - ) - .await + xtream_load_rewrite_playlist(XtreamCluster::Video, &app_state.app_config, &target, category_id, &user) + .await ), crate::model::XC_ACTION_GET_SERIES => skip_flag_optional!( skip_series, - xtream_load_rewrite_playlist( - XtreamCluster::Series, - &app_state.app_config, - &target, - category_id, - &user - ) - .await + xtream_load_rewrite_playlist(XtreamCluster::Series, &app_state.app_config, &target, category_id, &user) + .await ), _ => Some(info_err_res!("Unknown api call: {action} for target: {}", &target.name)), }; @@ -1308,17 +1233,11 @@ async fn xtream_player_api( let content_stream = xtream_create_content_stream(xtream_iter); try_unwrap_body!(axum::response::Response::builder() .status(axum::http::StatusCode::OK) - .header( - axum::http::header::CONTENT_TYPE, - mime::APPLICATION_JSON.to_string() - ) + .header(axum::http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string()) .body(axum::body::Body::from_stream(content_stream))) } Err(err) => { - error!( - "Failed response for xtream target: {} action: {} error: {}", - &target.name, action, err - ); + error!("Failed response for xtream target: {} action: {} error: {}", &target.name, action, err); // Some players fail on NoContent, so we return an empty array api_utils::empty_json_list_response().into_response() } @@ -1339,11 +1258,9 @@ async fn xtream_player_api( } } -fn xtream_create_content_stream( - xtream_iter: S, -) -> impl Stream> +fn xtream_create_content_stream(xtream_iter: S) -> impl Stream> where - S: Stream + Send + Unpin + 'static, + S: Stream + Send + Unpin + 'static, { let mapped = xtream_iter.map(move |(mut line, has_next)| { if has_next { @@ -1414,12 +1331,7 @@ macro_rules! register_xtream_api_timeshift { async fn xtream_player_token_stream( fingerprint: Fingerprint, - axum::extract::Path((token, target_id, cluster, stream_id)): axum::extract::Path<( - String, - u16, - String, - String, - )>, + axum::extract::Path((token, target_id, cluster, stream_id)): axum::extract::Path<(String, u16, String, String)>, axum::extract::State(app_state): axum::extract::State>, req_headers: HeaderMap, ) -> impl IntoResponse + Send { @@ -1431,17 +1343,15 @@ async fn xtream_player_token_stream( target_id, ApiStreamRequest::from_access_token(ctxt, &token, &stream_id, ""), ) - .await - .into_response() + .await + .into_response() } pub fn xtream_api_register() -> axum::Router> { let router = axum::Router::new(); let mut router = register_xtream_api!(router, ["/player_api.php", "/panel_api.php", "/xtream"]); - router = router.route( - "/token/{token}/{target_id}/{cluster}/{stream_id}", - axum::routing::get(xtream_player_token_stream), - ); + router = router + .route("/token/{token}/{target_id}/{cluster}/{stream_id}", axum::routing::get(xtream_player_token_stream)); router = register_xtream_api_stream!( router, [ diff --git a/backend/src/api/hdhomerun_proprietary.rs b/backend/src/api/hdhomerun_proprietary.rs index e79d93738..cf30b4824 100644 --- a/backend/src/api/hdhomerun_proprietary.rs +++ b/backend/src/api/hdhomerun_proprietary.rs @@ -1,16 +1,22 @@ -use crate::api::model::AppState; -use crate::model::{AppConfig, HdHomeRunDeviceConfig, HdHomeRunFlags}; +use crate::{ + api::model::AppState, + model::{AppConfig, HdHomeRunDeviceConfig, HdHomeRunFlags}, +}; use bytes::{Buf, BufMut, BytesMut}; use log::{error, info, trace}; -use std::collections::HashMap; -use std::io::Cursor; -use std::net::{Ipv4Addr, SocketAddr}; -use std::sync::Arc; -use std::time::Duration; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use tokio::net::{TcpListener, UdpSocket}; -use tokio_util::sync::CancellationToken; use shared::utils::Internable; +use std::{ + collections::HashMap, + io::Cursor, + net::{Ipv4Addr, SocketAddr}, + sync::Arc, + time::Duration, +}; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::{TcpListener, UdpSocket}, +}; +use tokio_util::sync::CancellationToken; const HDHR_PROPRIETARY_PORT: u16 = 65001; @@ -81,36 +87,18 @@ fn build_discover_response(device: &HdHomeRunDeviceConfig, server_host: &str) -> let device_id = u32::from_str_radix(&device.device_id, 16).unwrap_or(0); - write_tlv_u32( - &mut payload, - packet::HDHOMERUN_TAG_DEVICE_TYPE, - packet::HDHOMERUN_DEVICE_TYPE_TUNER, - ); - write_tlv_u32( - &mut payload, - packet::HDHOMERUN_TAG_DEVICE_ID, - device_id, - ); - write_tlv_str( - &mut payload, - packet::HDHOMERUN_TAG_BASE_URL, - &base_url, - ); - write_tlv_u8( - &mut payload, - packet::HDHOMERUN_TAG_TUNER_COUNT, - device.tuner_count, - ); - write_tlv_str( - &mut payload, - packet::HDHOMERUN_TAG_LINEUP_URL, - &lineup_url, - ); + write_tlv_u32(&mut payload, packet::HDHOMERUN_TAG_DEVICE_TYPE, packet::HDHOMERUN_DEVICE_TYPE_TUNER); + write_tlv_u32(&mut payload, packet::HDHOMERUN_TAG_DEVICE_ID, device_id); + write_tlv_str(&mut payload, packet::HDHOMERUN_TAG_BASE_URL, &base_url); + write_tlv_u8(&mut payload, packet::HDHOMERUN_TAG_TUNER_COUNT, device.tuner_count); + write_tlv_str(&mut payload, packet::HDHOMERUN_TAG_LINEUP_URL, &lineup_url); let mut response = BytesMut::new(); response.put_u16(packet::HDHOMERUN_TYPE_DISCOVER_RSP); //response.put_u16(u16::try_from(payload.len()).unwrap_or(0)); - if let Ok(n) = u16::try_from(payload.len()) { response.put_u16(n) } else { + if let Ok(n) = u16::try_from(payload.len()) { + response.put_u16(n); + } else { error!("HDHR response payload too large ({} bytes)", payload.len()); return Vec::new(); } @@ -149,24 +137,24 @@ fn parse_tlv(cursor: &mut Cursor<&[u8]>) -> HashMap> { // } // let len = u16::from_be_bytes(len_buf) as usize; - // Variable-length TLV decoding - let mut first_len_byte = [0u8; 1]; - if std::io::Read::read_exact(cursor, &mut first_len_byte).is_err() { - break; - } - let len = if first_len_byte[0] & 0x80 == 0 { - // Single-byte length (≤127) - first_len_byte[0] as usize - } else { - // Two-byte length (≥128) - let mut second_len_byte = [0u8; 1]; - if std::io::Read::read_exact(cursor, &mut second_len_byte).is_err() { - break; - } - let low_bits = (first_len_byte[0] & 0x7F) as usize; - let high_bits = (second_len_byte[0] as usize) << 7; - low_bits | high_bits - }; + // Variable-length TLV decoding + let mut first_len_byte = [0u8; 1]; + if std::io::Read::read_exact(cursor, &mut first_len_byte).is_err() { + break; + } + let len = if first_len_byte[0] & 0x80 == 0 { + // Single-byte length (≤127) + first_len_byte[0] as usize + } else { + // Two-byte length (≥128) + let mut second_len_byte = [0u8; 1]; + if std::io::Read::read_exact(cursor, &mut second_len_byte).is_err() { + break; + } + let low_bits = (first_len_byte[0] & 0x7F) as usize; + let high_bits = (second_len_byte[0] as usize) << 7; + low_bits | high_bits + }; // Check for incomplete TLV let remaining = cursor.get_ref().len() as u64 - cursor.position(); @@ -185,11 +173,7 @@ fn parse_tlv(cursor: &mut Cursor<&[u8]>) -> HashMap> { tags } -async fn proprietary_discover_loop( - socket: UdpSocket, - app_config: Arc, - server_host: String, -) { +async fn proprietary_discover_loop(socket: UdpSocket, app_config: Arc, server_host: String) { let mut buf = [0; 1024]; loop { let (len, remote_addr) = match socket.recv_from(&mut buf).await { @@ -248,13 +232,16 @@ async fn proprietary_discover_loop( if should_reply { let response = build_discover_response(device, &server_host); if response.is_empty() { - error!("Failed to build discovery response for device '{}'", device.name); - continue; + error!("Failed to build discovery response for device '{}'", device.name); + continue; } if let Err(e) = socket.send_to(&response, remote_addr).await { error!("Failed to send proprietary discovery response to {remote_addr}: {e}"); } else { - trace!("Sent proprietary discovery response for device '{}' to {remote_addr}", device.name); + trace!( + "Sent proprietary discovery response for device '{}' to {remote_addr}", + device.name + ); } } } @@ -267,11 +254,7 @@ async fn proprietary_discover_loop( // --- TCP Get/Set Logic --- -async fn handle_tcp_connection( - mut stream: tokio::net::TcpStream, - _addr: SocketAddr, - app_state: Arc, -) { +async fn handle_tcp_connection(mut stream: tokio::net::TcpStream, _addr: SocketAddr, app_state: Arc) { let mut buf = [0; 1024]; loop { match stream.read(&mut buf).await { @@ -332,19 +315,11 @@ async fn process_getset_request(request: &[u8], app_state: &Arc) -> Ve trace!("Received GET/SET for: {name}"); let name_str = name.trim_end_matches('\0'); - write_tlv_str( - &mut response_payload, - packet::HDHOMERUN_TAG_GETSET_NAME, - name_str, - ); + write_tlv_str(&mut response_payload, packet::HDHOMERUN_TAG_GETSET_NAME, name_str); match name_str { "/sys/model" => { - write_tlv_str( - &mut response_payload, - packet::HDHOMERUN_TAG_GETSET_VALUE, - "hdhomerun4_atsc", - ); + write_tlv_str(&mut response_payload, packet::HDHOMERUN_TAG_GETSET_VALUE, "hdhomerun4_atsc"); } s if s.starts_with("/tuner") && s.ends_with("/status") => { let rest = &s[6..]; @@ -356,11 +331,7 @@ async fn process_getset_request(request: &[u8], app_state: &Arc) -> Ve } else { "ch=none lock=none ss=0 snq=0 seq=0 bps=0 pps=0".to_string() }; - write_tlv_str( - &mut response_payload, - packet::HDHOMERUN_TAG_GETSET_VALUE, - &status_str, - ); + write_tlv_str(&mut response_payload, packet::HDHOMERUN_TAG_GETSET_VALUE, &status_str); } } s if s.starts_with("/tuner") && s.ends_with("/vchannel") => { @@ -373,11 +344,7 @@ async fn process_getset_request(request: &[u8], app_state: &Arc) -> Ve } else { "none".intern() }; - write_tlv_str( - &mut response_payload, - packet::HDHOMERUN_TAG_GETSET_VALUE, - &vchannel, - ); + write_tlv_str(&mut response_payload, packet::HDHOMERUN_TAG_GETSET_VALUE, &vchannel); } } s if s.starts_with("/tuner") && s.ends_with("/lockkey") => { @@ -385,7 +352,7 @@ async fn process_getset_request(request: &[u8], app_state: &Arc) -> Ve write_tlv_str(&mut response_payload, packet::HDHOMERUN_TAG_ERROR_MESSAGE, err_msg); } _ => { - write_tlv_str(&mut response_payload,packet::HDHOMERUN_TAG_GETSET_VALUE,""); + write_tlv_str(&mut response_payload, packet::HDHOMERUN_TAG_GETSET_VALUE, ""); } } } @@ -406,10 +373,7 @@ async fn process_getset_request(request: &[u8], app_state: &Arc) -> Ve response.to_vec() } -async fn proprietary_tcp_listener_loop( - app_state: Arc, - cancel_token: CancellationToken, -) { +async fn proprietary_tcp_listener_loop(app_state: Arc, cancel_token: CancellationToken) { let addr = SocketAddr::from((Ipv4Addr::UNSPECIFIED, HDHR_PROPRIETARY_PORT)); let listener = match TcpListener::bind(addr).await { Ok(l) => l, @@ -439,11 +403,7 @@ async fn proprietary_tcp_listener_loop( } } -pub fn spawn_proprietary_tasks( - app_state: Arc, - server_host: String, - cancel_token: CancellationToken, -) { +pub fn spawn_proprietary_tasks(app_state: Arc, server_host: String, cancel_token: CancellationToken) { let app_config = Arc::clone(&app_state.app_config); let cancel_token_udp = cancel_token.clone(); diff --git a/backend/src/api/hdhomerun_ssdp.rs b/backend/src/api/hdhomerun_ssdp.rs index 823b1fd5f..542741abf 100644 --- a/backend/src/api/hdhomerun_ssdp.rs +++ b/backend/src/api/hdhomerun_ssdp.rs @@ -1,9 +1,11 @@ use crate::model::{AppConfig, HdHomeRunDeviceConfig, HdHomeRunFlags}; -use socket2::{Domain, Protocol, Socket, Type}; -use std::net::{Ipv4Addr, SocketAddr, UdpSocket as StdUdpSocket}; -use std::sync::Arc; -use std::time::Duration; use log::{error, info, trace}; +use socket2::{Domain, Protocol, Socket, Type}; +use std::{ + net::{Ipv4Addr, SocketAddr, UdpSocket as StdUdpSocket}, + sync::Arc, + time::Duration, +}; use tokio::net::UdpSocket; use tokio_util::sync::CancellationToken; @@ -36,27 +38,29 @@ async fn ssdp_task_loop(socket: UdpSocket, app_config: Arc, server_ho }; let request = String::from_utf8_lossy(&buf[..len]); - if !request.starts_with("M-SEARCH") { continue; } + if !request.starts_with("M-SEARCH") { + continue; + } let req = request.to_ascii_lowercase(); - if !req.contains(r#"man: "ssdp:discover""#) { continue; } + if !req.contains(r#"man: "ssdp:discover""#) { + continue; + } // Extract ST and MX (defaults) - let st = req.lines() + let st = req + .lines() .find_map(|l| l.strip_prefix("st:").map(|v| v.trim().to_string())) .unwrap_or_else(|| "ssdp:all".to_string()); - let mx: u64 = req.lines() - .find_map(|l| l.strip_prefix("mx:").and_then(|v| v.trim().parse().ok())) - .unwrap_or(1); + let mx: u64 = req.lines().find_map(|l| l.strip_prefix("mx:").and_then(|v| v.trim().parse().ok())).unwrap_or(1); // Normalize to the set we support - let supported = [ - "urn:schemas-upnp-org:device:mediaserver:1", - "upnp:rootdevice", - "ssdp:all", - ]; - if !supported.contains(&st.as_str()) { continue; } + let supported = ["urn:schemas-upnp-org:device:mediaserver:1", "upnp:rootdevice", "ssdp:all"]; + if !supported.contains(&st.as_str()) { + continue; + } // Randomized delay per MX - let delay_ms = (fastrand::u64(0..=mx*1000)).min(2000); - if delay_ms > 0 { tokio::time::sleep(Duration::from_millis(delay_ms)).await; } - + let delay_ms = (fastrand::u64(0..=mx * 1000)).min(2000); + if delay_ms > 0 { + tokio::time::sleep(Duration::from_millis(delay_ms)).await; + } trace!("Received HDHomeRun M-SEARCH from {remote_addr}"); let hdhomerun_guard = app_config.hdhomerun.load(); @@ -87,9 +91,13 @@ pub fn spawn_ssdp_discover_task(app_config: Arc, server_host: String, return; } }; - if let Err(e) = std_socket.set_reuse_address(true) { error!("Failed to set reuse_address on SSDP socket: {e}"); } + if let Err(e) = std_socket.set_reuse_address(true) { + error!("Failed to set reuse_address on SSDP socket: {e}"); + } #[cfg(not(windows))] - if let Err(e) = std_socket.set_reuse_port(true) { error!("Failed to set reuse_port on SSDP socket: {e}"); } + if let Err(e) = std_socket.set_reuse_port(true) { + error!("Failed to set reuse_port on SSDP socket: {e}"); + } if let Err(e) = std_socket.bind(&addr.into()) { error!("Failed to bind SSDP socket to {addr}: {e}"); return; diff --git a/backend/src/api/library_scan.rs b/backend/src/api/library_scan.rs index 3df94a041..33e84ad18 100644 --- a/backend/src/api/library_scan.rs +++ b/backend/src/api/library_scan.rs @@ -1,6 +1,8 @@ -use crate::api::model::{EventManager, EventMessage, UpdateGuardPermit}; -use crate::library::LibraryProcessor; -use crate::model::{LibraryConfig, MetadataUpdateConfig}; +use crate::{ + api::model::{EventManager, EventMessage, UpdateGuardPermit}, + library::LibraryProcessor, + model::{LibraryConfig, MetadataUpdateConfig}, +}; use log::{error, info}; use shared::model::{LibraryScanSummary, LibraryScanSummaryStatus}; use std::sync::Arc; @@ -25,10 +27,7 @@ pub(crate) fn spawn_library_scan( status: LibraryScanSummaryStatus::Success, message: format!( "{prefix}Scan completed: {} files scanned, {} added, {} updated, {} removed", - result.files_scanned, - result.files_added, - result.files_updated, - result.files_removed + result.files_scanned, result.files_added, result.files_updated, result.files_removed ), result: Some(result), }; diff --git a/backend/src/api/main_api.rs b/backend/src/api/main_api.rs index 8699c84c8..ce19f8967 100644 --- a/backend/src/api/main_api.rs +++ b/backend/src/api/main_api.rs @@ -1,44 +1,54 @@ -use crate::api::api_utils::{get_build_time, get_server_time}; -use crate::api::config_watch::exec_config_watch; -use crate::api::endpoints::custom_video_stream_api::cvs_api_register; -use crate::api::endpoints::hdhomerun_api::hdhr_api_register; -use crate::api::endpoints::hls_api::hls_api_register; -use crate::api::endpoints::m3u_api::m3u_api_register; -use crate::api::endpoints::v1_api::v1_api_register; -use crate::api::endpoints::web_index::{index_register_with_path, index_register_without_path}; -use crate::api::endpoints::websocket_api::ws_api_register; -use crate::api::endpoints::xmltv_api::xmltv_api_register; -use crate::api::endpoints::xtream_api::xtream_api_register; -use crate::api::hdhomerun_proprietary::spawn_proprietary_tasks; -use crate::api::hdhomerun_ssdp::spawn_ssdp_discover_task; -use crate::api::model::{ - create_cache, create_http_client, create_http_client_no_redirect, ActiveProviderManager, ActiveUserManager, - AppState, CancelTokens, ConnectionManager, DownloadQueue, EventManager, EventMessage, HdHomerunAppState, - MetadataUpdateManager, PlaylistStorageState, SharedStreamManager, UpdateGuard, +use crate::{ + api::{ + api_utils::{get_build_time, get_server_time}, + config_watch::exec_config_watch, + endpoints::{ + custom_video_stream_api::cvs_api_register, + hdhomerun_api::hdhr_api_register, + hls_api::hls_api_register, + m3u_api::m3u_api_register, + v1_api::v1_api_register, + web_index::{index_register_with_path, index_register_without_path}, + websocket_api::ws_api_register, + xmltv_api::xmltv_api_register, + xtream_api::xtream_api_register, + }, + hdhomerun_proprietary::spawn_proprietary_tasks, + hdhomerun_ssdp::spawn_ssdp_discover_task, + model::{ + create_cache, create_http_client, create_http_client_no_redirect, ActiveProviderManager, ActiveUserManager, + AppState, CancelTokens, ConnectionManager, DownloadQueue, EventManager, EventMessage, HdHomerunAppState, + MetadataUpdateManager, PlaylistStorageState, SharedStreamManager, UpdateGuard, + }, + panel_api::sync_panel_api_exp_dates_on_boot, + scheduler::{exec_interner_prune, exec_scheduler}, + serve::serve, + sys_usage::exec_system_usage, + }, + model::{AppConfig, Config, HdHomeRunFlags, Healthcheck, ProcessTargets, RateLimitConfig}, + processing::processor::exec_processing, + repository::{get_geoip_path, load_playlists_into_memory_cache}, + utils::{exec_file_lock_prune, get_default_web_root_path, GeoIp}, + VERSION, }; -use crate::api::panel_api::sync_panel_api_exp_dates_on_boot; -use crate::api::scheduler::{exec_interner_prune, exec_scheduler}; -use crate::api::serve::serve; -use crate::api::sys_usage::exec_system_usage; -use crate::model::{AppConfig, Config, HdHomeRunFlags, Healthcheck, ProcessTargets, RateLimitConfig}; -use crate::processing::processor::exec_processing; -use crate::repository::get_geoip_path; -use crate::repository::load_playlists_into_memory_cache; -use crate::utils::{exec_file_lock_prune, get_default_web_root_path, GeoIp}; -use crate::VERSION; use arc_swap::{ArcSwap, ArcSwapOption}; -use axum::extract::connect_info::ConnectInfo; -use axum::Router; -use axum::{extract::Request, middleware::Next}; +use axum::{ + extract::{connect_info::ConnectInfo, Request}, + middleware::Next, + Router, +}; use log::{debug, error, info, warn}; -use shared::error::TuliproxError; -use shared::utils::{concat_path_leading_slash, sanitize_sensitive_info}; -use shared::{info_err, info_err_res}; -use std::collections::{HashMap, HashSet}; -use std::net::SocketAddr; -use std::path::PathBuf; -use std::sync::atomic::AtomicI8; -use std::sync::Arc; +use shared::{ + error::TuliproxError, + info_err, info_err_res, + utils::{concat_path_leading_slash, sanitize_sensitive_info}, +}; +use std::{ + collections::{HashMap, HashSet}, + net::SocketAddr, + path::PathBuf, + sync::{atomic::AtomicI8, Arc}, +}; use tokio_util::sync::CancellationToken; use tower_governor::key_extractor::SmartIpKeyExtractor; use tower_http::services::ServeDir; @@ -62,9 +72,7 @@ fn create_healthcheck() -> Healthcheck { } } -async fn healthcheck() -> impl axum::response::IntoResponse { - axum::Json(create_healthcheck()) -} +async fn healthcheck() -> impl axum::response::IntoResponse { axum::Json(create_healthcheck()) } async fn create_shared_data( app_config: &Arc, @@ -124,7 +132,7 @@ async fn create_shared_data( }) } -fn exec_update_on_boot(client: &reqwest::Client, app_state: &Arc, targets: &Arc) { +fn exec_update_on_boot(client: &reqwest::Client, app_state: &Arc, targets: &Arc) -> bool { let cfg = &app_state.app_config; let update_on_boot = { let config = cfg.config.load(); @@ -140,6 +148,7 @@ fn exec_update_on_boot(client: &reqwest::Client, app_state: &Arc, targ let provider_manager = Arc::clone(&app_state.active_provider); let metadata_manager = Arc::clone(&app_state.metadata_manager); let event_manager = Some(Arc::clone(&app_state.event_manager)); + let app_state_clone = Arc::clone(app_state); tokio::spawn(async move { exec_processing( @@ -147,6 +156,7 @@ fn exec_update_on_boot(client: &reqwest::Client, app_state: &Arc, targ app_config_clone, targets_clone, event_manager, + Some(app_state_clone), Some(playlist_state), update_guard, disabled_headers, @@ -157,7 +167,9 @@ fn exec_update_on_boot(client: &reqwest::Client, app_state: &Arc, targ ) .await; }); + return true; } + false } fn is_web_auth_enabled(cfg: &Arc, web_ui_enabled: bool) -> bool { @@ -288,12 +300,12 @@ pub async fn start_server(app_config: Arc, targets: Arc, targets: &Arc, ) -> Self { - Self { - client_id, - allocation_id, - allocation, - cancel_token, - } + Self { client_id, allocation_id, allocation, cancel_token } } } @@ -106,9 +109,7 @@ impl ActiveProviderManager { cfg.sources.load().inputs.iter().filter(|i| i.enabled).map(Arc::clone).collect() } - fn get_grace_options(cfg: &AppConfig) -> GracePeriodOptions { - cfg.config.load().get_grace_options() - } + fn get_grace_options(cfg: &AppConfig) -> GracePeriodOptions { cfg.config.load().get_grace_options() } pub async fn update_config(&self, cfg: &AppConfig) { let grace_period_options = Self::get_grace_options(cfg); @@ -153,27 +154,14 @@ impl ActiveProviderManager { ) -> Option { // 1. Try to acquire directly let (allow_grace, allocation) = if force { - ( - true, - self.providers - .force_exact_acquire_connection(provider_or_input_name) - .await, - ) + (true, self.providers.force_exact_acquire_connection(provider_or_input_name).await) } else { match allow_grace_override { Some(allow_grace) => ( allow_grace, - self.providers - .acquire_connection_with_grace_override( - provider_or_input_name, - allow_grace, - ) - .await, - ), - None => ( - true, - self.providers.acquire_connection(provider_or_input_name).await, + self.providers.acquire_connection_with_grace_override(provider_or_input_name, allow_grace).await, ), + None => (true, self.providers.acquire_connection(provider_or_input_name).await), } }; @@ -207,19 +195,24 @@ impl ActiveProviderManager { let mut connections = self.connections.write().await; let per_addr = connections.single.entry(*addr).or_default(); - per_addr.insert(allocation_id, ActiveConnectionInfo { - allocation: allocation.clone(), - priority, - is_probe, - cancel_token: cancel_token.clone(), - created_at: Instant::now(), - }); + per_addr.insert( + allocation_id, + ActiveConnectionInfo { + allocation: allocation.clone(), + priority, + is_probe, + cancel_token: cancel_token.clone(), + created_at: Instant::now(), + }, + ); - connections.by_provider.entry(provider_name.clone()) - .or_default() - .push((*addr, allocation_id)); + connections.by_provider.entry(provider_name.clone()).or_default().push((*addr, allocation_id)); - debug_if_enabled!("Added provider connection {provider_name:?} for {} (prio={})", sanitize_sensitive_info(&addr.to_string()), priority); + debug_if_enabled!( + "Added provider connection {provider_name:?} for {} (prio={})", + sanitize_sensitive_info(&addr.to_string()), + priority + ); Some(ProviderHandle::new(*addr, allocation_id, allocation, Some(cancel_token))) } @@ -246,7 +239,8 @@ impl ActiveProviderManager { let is_better = match victim { None => true, Some((_, _, v_prio, v_created, _)) => { - info.priority > v_prio || (info.priority == v_prio && info.created_at < v_created) + info.priority > v_prio + || (info.priority == v_prio && info.created_at < v_created) } }; if is_better { @@ -267,13 +261,11 @@ impl ActiveProviderManager { None => true, Some((_, _, v_prio, v_created, _)) => { shared.priority > v_prio - || (shared.priority == v_prio - && shared.created_at < v_created) + || (shared.priority == v_prio && shared.created_at < v_created) } }; if is_better { - victim = - Some((*addr, *alloc_id, shared.priority, shared.created_at, true)); + victim = Some((*addr, *alloc_id, shared.priority, shared.created_at, true)); } } } @@ -284,9 +276,13 @@ impl ActiveProviderManager { } if let Some((addr, alloc_id, v_prio, victim_created_at, is_shared)) = victim { - debug_if_enabled!("Preempting {} connection from {} (prio={}) for higher priority request (prio={})", + debug_if_enabled!( + "Preempting {} connection from {} (prio={}) for higher priority request (prio={})", if is_shared { "shared" } else { "single" }, - sanitize_sensitive_info(&addr.to_string()), v_prio, new_priority); + sanitize_sensitive_info(&addr.to_string()), + v_prio, + new_priority + ); if is_shared { let released_shared_allocation = { @@ -304,10 +300,7 @@ impl ActiveProviderManager { if !still_single { None } else if let Some(shared) = connections.shared.by_key.remove(&key) { - connections - .shared - .shared_by_allocation_id - .remove(&shared.allocation_id); + connections.shared.shared_by_allocation_id.remove(&shared.allocation_id); for shared_addr in &shared.connections { connections.shared.key_by_addr.remove(shared_addr); } @@ -386,10 +379,7 @@ impl ActiveProviderManager { } // Now try acquire again preserving the original grace policy. - let allocation = self - .providers - .acquire_connection_with_grace_override(input_name, allow_grace) - .await; + let allocation = self.providers.acquire_connection_with_grace_override(input_name, allow_grace).await; if !matches!(allocation, ProviderAllocation::Exhausted) { return Some(allocation); } @@ -404,10 +394,7 @@ impl ActiveProviderManager { addr: &SocketAddr, allow_grace: bool, ) -> Option { - let allocation = self - .providers - .acquire_exact_connection_with_grace_override(provider_name, allow_grace) - .await; + let allocation = self.providers.acquire_exact_connection_with_grace_override(provider_name, allow_grace).await; if matches!(allocation, ProviderAllocation::Exhausted) { return None; } @@ -435,18 +422,13 @@ impl ActiveProviderManager { addr: &SocketAddr, allow_grace: bool, ) -> Option { - self.acquire_connection_inner(input_name, addr, false, Some(allow_grace), DEFAULT_USER_PRIORITY, false) - .await + self.acquire_connection_inner(input_name, addr, false, Some(allow_grace), DEFAULT_USER_PRIORITY, false).await } /// Acquire a provider connection while optionally disabling provider grace allocations. - pub async fn acquire_connection_for_probe( - &self, - input_name: &Arc, - ) -> Option { + pub async fn acquire_connection_for_probe(&self, input_name: &Arc) -> Option { // Probe is strictly low-priority and must never consume grace capacity. - self.acquire_connection_inner(input_name, &DUMMY_ADDR, false, Some(false), DEFAULT_PROBE_PRIORITY, true) - .await + self.acquire_connection_inner(input_name, &DUMMY_ADDR, false, Some(false), DEFAULT_PROBE_PRIORITY, true).await } // This method is used for redirects to cycle through the provider @@ -522,10 +504,7 @@ impl ActiveProviderManager { if shared.connections.is_empty() { // If this was the last user of the shared allocation: connections.shared.by_key.remove(&key); - connections - .shared - .shared_by_allocation_id - .remove(&shared.allocation_id); + connections.shared.shared_by_allocation_id.remove(&shared.allocation_id); if let Some(name) = shared.allocation.get_provider_name() { if let Some(list) = connections.by_provider.get_mut(&name) { list.retain(|(_, i)| *i != shared.allocation_id); @@ -543,10 +522,10 @@ impl ActiveProviderManager { if let Some(allocation) = shared_allocation { allocation.release().await; debug_if_enabled!( - "Released last shared connection for provider {}, releasing allocation {}", - allocation.get_provider_name().unwrap_or_default(), - sanitize_sensitive_info(&addr.to_string()) - ); + "Released last shared connection for provider {}, releasing allocation {}", + allocation.get_provider_name().unwrap_or_default(), + sanitize_sensitive_info(&addr.to_string()) + ); } } @@ -566,7 +545,9 @@ impl ActiveProviderManager { // Remove from by_provider index if let Some(name) = released.as_ref().and_then(ProviderAllocation::get_provider_name) { if let Some(list) = connections.by_provider.get_mut(&name) { - if let Some(idx) = list.iter().position(|(a, i)| *a == handle.client_id && *i == handle.allocation_id) { + if let Some(idx) = + list.iter().position(|(a, i)| *a == handle.client_id && *i == handle.allocation_id) + { list.remove(idx); } } @@ -576,11 +557,7 @@ impl ActiveProviderManager { if released.is_none() { // Try removing from Shared - if let Some(key) = connections - .shared - .shared_by_allocation_id - .remove(&handle.allocation_id) - { + if let Some(key) = connections.shared.shared_by_allocation_id.remove(&handle.allocation_id) { if let Some(shared) = connections.shared.by_key.remove(&key) { released = Some(shared.allocation); for addr in shared.connections { @@ -615,7 +592,6 @@ impl ActiveProviderManager { } else { let mut iter = m.drain(); if let Some((id, info)) = iter.next() { - // Collect others as extras to release for (_, extra_info) in iter { extras.push(extra_info.allocation); @@ -639,8 +615,9 @@ impl ActiveProviderManager { None } } - } else { None }; - + } else { + None + }; if let Some(handle) = &handle { let provider_name = handle.0.allocation.get_provider_name().unwrap_or_default(); @@ -661,10 +638,7 @@ impl ActiveProviderManager { }, ); connections.shared.key_by_addr.insert(*addr, key.to_string()); - connections - .shared - .shared_by_allocation_id - .insert(handle.0.allocation_id, key.to_string()); + connections.shared.shared_by_allocation_id.insert(handle.0.allocation_id, key.to_string()); } extras }; @@ -686,31 +660,27 @@ impl ActiveProviderManager { shared_allocation.connections.insert(*addr); connections.shared.key_by_addr.insert(*addr, key.to_string()); } else { - error!( - "Failed to add shared connection for {addr}: url {} not found", - sanitize_sensitive_info(key) - ); + error!("Failed to add shared connection for {addr}: url {} not found", sanitize_sensitive_info(key)); } } - pub async fn get_provider_connections_count(&self) -> usize { - self.providers.active_connection_count().await - } + pub async fn get_provider_connections_count(&self) -> usize { self.providers.active_connection_count().await } } #[cfg(test)] mod tests { use super::ActiveProviderManager; - use crate::api::model::EventManager; - use crate::model::{AppConfig, Config, ConfigInput, ConfigInputAlias, SourcesConfig}; - use crate::utils::FileLockManager; + use crate::{ + api::model::EventManager, + model::{AppConfig, Config, ConfigInput, ConfigInputAlias, SourcesConfig}, + utils::FileLockManager, + }; use arc_swap::{ArcSwap, ArcSwapOption}; - use shared::model::{ConfigPaths, InputFetchMethod, InputType}; - use shared::utils::Internable; - use std::collections::HashMap; - use std::net::SocketAddr; - use std::sync::Arc; - use std::time::Duration; + use shared::{ + model::{ConfigPaths, InputFetchMethod, InputType}, + utils::Internable, + }; + use std::{collections::HashMap, net::SocketAddr, sync::Arc, time::Duration}; fn build_test_app_config(aliases: Option>, max_connections: u16) -> AppConfig { let input = Arc::new(ConfigInput { @@ -729,10 +699,7 @@ mod tests { ..ConfigInput::default() }); - let sources = SourcesConfig { - inputs: vec![input], - ..SourcesConfig::default() - }; + let sources = SourcesConfig { inputs: vec![input], ..SourcesConfig::default() }; AppConfig { config: Arc::new(ArcSwap::from_pointee(Config::default())), @@ -775,9 +742,7 @@ mod tests { ) } - fn create_test_app_config_single_provider_pool() -> AppConfig { - build_test_app_config(None, 1) - } + fn create_test_app_config_single_provider_pool() -> AppConfig { build_test_app_config(None, 1) } #[tokio::test] async fn test_force_exact_acquire_does_not_overallocate_busy_provider() { @@ -789,20 +754,13 @@ mod tests { let client_1_addr: SocketAddr = "127.0.0.1:40001".parse().unwrap(); let client_2_addr: SocketAddr = "127.0.0.1:40002".parse().unwrap(); - let first_alloc = manager - .acquire_connection(&input_name, &client_1_addr) - .await - .expect("client1 initial allocation"); - let pinned_provider = first_alloc - .allocation - .get_provider_name() - .expect("provider name expected"); + let first_alloc = + manager.acquire_connection(&input_name, &client_1_addr).await.expect("client1 initial allocation"); + let pinned_provider = first_alloc.allocation.get_provider_name().expect("provider name expected"); assert_eq!(pinned_provider.as_ref(), "provider_1"); // provider_1 has max_connections=1 and is already in use by client1 - let forced = manager - .force_exact_acquire_connection(&pinned_provider, &client_2_addr) - .await; + let forced = manager.force_exact_acquire_connection(&pinned_provider, &client_2_addr).await; assert!(forced.is_none(), "forced exact acquire must not over-allocate busy provider"); manager.release_connection(&client_1_addr).await; @@ -820,17 +778,16 @@ mod tests { let client_2_addr: SocketAddr = "127.0.0.1:41002".parse().unwrap(); // Step 1: Client1 starts movie -> provider_1 - let first_alloc = manager.acquire_connection(&input_name, &client_1_addr).await.expect("client1 initial allocation"); - assert_eq!( - first_alloc.allocation.get_provider_name().as_deref(), - Some(input_name.as_ref()) - ); + let first_alloc = + manager.acquire_connection(&input_name, &client_1_addr).await.expect("client1 initial allocation"); + assert_eq!(first_alloc.allocation.get_provider_name().as_deref(), Some(input_name.as_ref())); // Step 2: Client1 stops -> release provider_1 manager.release_connection(&client_1_addr).await; // Step 3: Client2 starts live -> provider_1 - let live_alloc = manager.acquire_connection(&input_name, &client_2_addr).await.expect("client2 live allocation"); + let live_alloc = + manager.acquire_connection(&input_name, &client_2_addr).await.expect("client2 live allocation"); let busy_provider = live_alloc.allocation.get_provider_name().expect("provider name expected"); assert_eq!(busy_provider.as_ref(), input_name.as_ref()); assert!(manager.is_exhausted(&busy_provider).await); @@ -841,10 +798,7 @@ mod tests { .acquire_connection_with_grace(&input_name, &client_1_addr, false) .await .expect("client1 fallback allocation without grace"); - let fallback_provider = fallback_alloc - .allocation - .get_provider_name() - .expect("fallback provider expected"); + let fallback_provider = fallback_alloc.allocation.get_provider_name().expect("fallback provider expected"); assert_ne!(fallback_provider.as_ref(), busy_provider.as_ref()); assert_eq!(fallback_provider.as_ref(), "provider_2"); @@ -864,25 +818,14 @@ mod tests { let client_2_addr: SocketAddr = "127.0.0.1:42002".parse().unwrap(); // Initial playback for client1. - let first_alloc = manager - .acquire_connection(&input_name, &client_1_addr) - .await - .expect("client1 initial allocation"); - let pinned_provider = first_alloc - .allocation - .get_provider_name() - .expect("provider name expected"); + let first_alloc = + manager.acquire_connection(&input_name, &client_1_addr).await.expect("client1 initial allocation"); + let pinned_provider = first_alloc.allocation.get_provider_name().expect("provider name expected"); assert_eq!(pinned_provider.as_ref(), "provider_1"); // Another client occupies the alternate account while client1 keeps seeking. - let second_alloc = manager - .acquire_connection(&input_name, &client_2_addr) - .await - .expect("client2 allocation"); - let second_provider = second_alloc - .allocation - .get_provider_name() - .expect("provider name expected"); + let second_alloc = manager.acquire_connection(&input_name, &client_2_addr).await.expect("client2 allocation"); + let second_provider = second_alloc.allocation.get_provider_name().expect("provider name expected"); assert_eq!(second_provider.as_ref(), "provider_2"); // Simulate repeated seek/range reconnects for client1: @@ -893,10 +836,7 @@ mod tests { .force_exact_acquire_connection(&pinned_provider, &client_1_addr) .await .expect("seek reacquire should stay on pinned provider"); - let seek_provider = seek_alloc - .allocation - .get_provider_name() - .expect("provider name expected"); + let seek_provider = seek_alloc.allocation.get_provider_name().expect("provider name expected"); assert_eq!(seek_provider.as_ref(), pinned_provider.as_ref()); } @@ -914,24 +854,16 @@ mod tests { let input_name = "provider_1".intern(); let user_addr: SocketAddr = "127.0.0.1:43001".parse().unwrap(); - let probe_handle = manager - .acquire_connection_for_probe(&input_name) - .await - .expect("probe allocation should succeed"); - let probe_token = probe_handle - .cancel_token - .clone() - .expect("probe handle must carry cancel token"); + let probe_handle = + manager.acquire_connection_for_probe(&input_name).await.expect("probe allocation should succeed"); + let probe_token = probe_handle.cancel_token.clone().expect("probe handle must carry cancel token"); // User request should preempt probe and immediately acquire released capacity. let user_alloc = manager .acquire_connection_with_grace(&input_name, &user_addr, false) .await .expect("user allocation should preempt probe"); - assert_eq!( - user_alloc.allocation.get_provider_name().as_deref(), - Some(input_name.as_ref()) - ); + assert_eq!(user_alloc.allocation.get_provider_name().as_deref(), Some(input_name.as_ref())); // Let the detached cancellation task start and arm its sleep timer on the paused clock. tokio::task::yield_now().await; @@ -942,16 +874,12 @@ mod tests { let cancel_wait_timeout = super::PREEMPTED_PROBE_CANCEL_GRACE + Duration::from_millis(500); let wait_token = probe_token.clone(); - let cancel_wait = tokio::spawn(async move { - tokio::time::timeout(cancel_wait_timeout, wait_token.cancelled()).await - }); + let cancel_wait = + tokio::spawn(async move { tokio::time::timeout(cancel_wait_timeout, wait_token.cancelled()).await }); tokio::task::yield_now().await; tokio::time::advance(cancel_wait_timeout).await; assert!( - cancel_wait - .await - .expect("cancel wait task should join") - .is_ok(), + cancel_wait.await.expect("cancel wait task should join").is_ok(), "probe token should be cancelled before timeout after grace" ); diff --git a/backend/src/api/model/active_user_manager.rs b/backend/src/api/model/active_user_manager.rs index e59897e4a..2d4da81fb 100644 --- a/backend/src/api/model/active_user_manager.rs +++ b/backend/src/api/model/active_user_manager.rs @@ -1,30 +1,40 @@ -use crate::api::model::{CustomVideoStreamType, EventManager, EventMessage}; -use crate::auth::Fingerprint; -use crate::model::Config; -use crate::model::ProxyUserCredentials; -use crate::utils::GeoIp; +use crate::{ + api::model::{CustomVideoStreamType, EventManager, EventMessage}, + auth::Fingerprint, + model::{Config, ProxyUserCredentials}, + utils::{debug_if_enabled, GeoIp}, +}; use arc_swap::ArcSwapOption; use jsonwebtoken::get_current_timestamp; use log::{debug, info}; -use shared::model::{ActiveUserConnectionChange, StreamChannel, StreamInfo, UserConnectionPermission, VirtualId}; -use shared::utils::{current_time_secs, default_grace_period_millis, default_grace_period_timeout_secs, - sanitize_sensitive_info, strip_port, Internable}; -use std::borrow::Cow; -use std::collections::{HashMap, HashSet}; -use std::net::SocketAddr; -use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}; -use std::sync::Arc; +use shared::{ + model::{ActiveUserConnectionChange, StreamChannel, StreamInfo, UserConnectionPermission, VirtualId}, + utils::{ + current_time_secs, default_grace_period_millis, default_grace_period_timeout_secs, sanitize_sensitive_info, + strip_port, Internable, + }, +}; +use std::{ + borrow::Cow, + collections::{HashMap, HashSet}, + net::SocketAddr, + sync::{ + atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, + Arc, + }, +}; use tokio::sync::RwLock; -use crate::utils::debug_if_enabled; -const USER_GC_TTL: u64 = 900; // 15 Min -const USER_CON_TTL: u64 = 10_800; // 3 hours +const USER_GC_TTL: u64 = 900; // 15 Min +const USER_CON_TTL: u64 = 10_800; // 3 hours const USER_SESSION_LIMIT: usize = 50; fn get_grace_options(config: &Config) -> (u64, u64) { - 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)); + 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), + ); (grace_period_millis, grace_period_timeout_secs) } @@ -95,9 +105,7 @@ pub struct ActiveUserManager { } impl ActiveUserManager { - pub fn new(config: &Config, - geoip: &Arc>, - event_manager: &Arc, ) -> Self { + pub fn new(config: &Config, geoip: &Arc>, event_manager: &Arc) -> Self { let log_active_user: bool = config.log.as_ref().is_some_and(|l| l.log_active_user); let (grace_period_millis, grace_period_timeout_secs) = get_grace_options(config); @@ -116,10 +124,11 @@ impl ActiveUserManager { async fn log_active_user(&self) { let is_log_user_enabled = self.is_log_user_enabled(); - let (user_count, user_connection_count) = { - self.active_users_and_connections().await - }; - self.event_manager.send_event(EventMessage::ActiveUser(ActiveUserConnectionChange::Connections(user_count, user_connection_count))); + let (user_count, user_connection_count) = { self.active_users_and_connections().await }; + self.event_manager.send_event(EventMessage::ActiveUser(ActiveUserConnectionChange::Connections( + user_count, + user_connection_count, + ))); if is_log_user_enabled { let last_user_count = self.last_logged_user_count.load(Ordering::Relaxed); let last_connection_count = self.last_logged_user_connection_count.load(Ordering::Relaxed); @@ -155,7 +164,10 @@ impl ActiveUserManager { if let Some(username) = disconnected_user { if !username.is_empty() { - debug_if_enabled!("Released connection for user {username} at {}", sanitize_sensitive_info(&addr.to_string())); + debug_if_enabled!( + "Released connection for user {username} at {}", + sanitize_sensitive_info(&addr.to_string()) + ); } } @@ -179,7 +191,11 @@ impl ActiveUserManager { 0 } - fn check_connection_permission(&self, username: &str, connection_data: &mut UserConnectionData) -> UserConnectionPermission { + fn check_connection_permission( + &self, + username: &str, + connection_data: &mut UserConnectionData, + ) -> UserConnectionPermission { let current_connections = connection_data.connections; if current_connections < connection_data.max_connections { @@ -192,7 +208,9 @@ impl ActiveUserManager { let now = get_current_timestamp(); // Check if user already used a grace period if connection_data.granted_grace { - if current_connections > connection_data.max_connections && now - connection_data.grace_ts <= self.grace_period_timeout_secs.load(Ordering::Relaxed) { + if current_connections > connection_data.max_connections + && now - connection_data.grace_ts <= self.grace_period_timeout_secs.load(Ordering::Relaxed) + { // Grace timeout, still active, deny connection debug!("User access denied, grace exhausted, too many connections: {username}"); return UserConnectionPermission::Exhausted; @@ -202,7 +220,9 @@ impl ActiveUserManager { connection_data.grace_ts = 0; } - if self.grace_period_millis.load(Ordering::Relaxed) > 0 && current_connections == connection_data.max_connections { + if self.grace_period_millis.load(Ordering::Relaxed) > 0 + && current_connections == connection_data.max_connections + { // Allow a grace period once connection_data.granted_grace = true; connection_data.grace_ts = now; @@ -215,11 +235,7 @@ impl ActiveUserManager { UserConnectionPermission::Exhausted } - pub async fn connection_permission( - &self, - username: &str, - max_connections: u32, - ) -> UserConnectionPermission { + pub async fn connection_permission(&self, username: &str, max_connections: u32) -> UserConnectionPermission { if max_connections > 0 { if let Some(connection_data) = self.connections.write().await.by_key.get_mut(username) { return self.check_connection_permission(username, connection_data); @@ -234,12 +250,14 @@ impl ActiveUserManager { .by_key .values() .filter(|c| c.connections > 0) - .fold((0usize, 0usize), |(user_count, conn_count), c| { - (user_count + 1, conn_count + c.connections as usize) - }) + .fold((0usize, 0usize), |(user_count, conn_count), c| (user_count + 1, conn_count + c.connections as usize)) } - pub async fn update_stream_detail(&self, addr: &SocketAddr, video_type: CustomVideoStreamType) -> Option { + pub async fn update_stream_detail( + &self, + addr: &SocketAddr, + video_type: CustomVideoStreamType, + ) -> Option { let mut user_connections = self.connections.write().await; let username = { match user_connections.key_by_addr.get(addr) { @@ -269,8 +287,16 @@ impl ActiveUserManager { } #[allow(clippy::too_many_arguments)] - pub async fn update_connection(&self, username: &str, max_connections: u32, fingerprint: &Fingerprint, - provider: &str, stream_channel: StreamChannel, user_agent: Cow<'_, str>, session_token: Option<&str>) -> Option { + pub async fn update_connection( + &self, + username: &str, + max_connections: u32, + fingerprint: &Fingerprint, + provider: &str, + stream_channel: StreamChannel, + user_agent: Cow<'_, str>, + session_token: Option<&str>, + ) -> Option { let stream_info = { let mut user_connections = self.connections.write().await; @@ -281,18 +307,16 @@ impl ActiveUserManager { user_connections.key_by_addr.insert(fingerprint.addr, username.to_string()); - let connection_data = user_connections.by_key + let connection_data = user_connections + .by_key .entry(username.to_string()) .or_insert_with(|| UserConnectionData::new(0, max_connections)); connection_data.max_connections = max_connections; let user_agent_string = user_agent.to_string(); - let existing_stream_info = connection_data - .streams - .iter_mut() - .find(|s| s.addr == fingerprint.addr) - .map(|stream_info| { + let existing_stream_info = + connection_data.streams.iter_mut().find(|s| s.addr == fingerprint.addr).map(|stream_info| { stream_info.channel = stream_channel.clone(); stream_info.provider = provider.to_string(); stream_info.user_agent.clone_from(&user_agent_string); @@ -304,7 +328,9 @@ impl ActiveUserManager { stream_info.clone() }); - if let Some(stream_info) = existing_stream_info { stream_info } else { + if let Some(stream_info) = existing_stream_info { + stream_info + } else { let country = { let geoip = self.geo_ip.load(); if let Some(geoip_db) = (*geoip).as_ref() { @@ -342,12 +368,16 @@ impl ActiveUserManager { Some(stream_info) } - fn is_log_user_enabled(&self) -> bool { - self.log_active_user.load(Ordering::Relaxed) - } + fn is_log_user_enabled(&self) -> bool { self.log_active_user.load(Ordering::Relaxed) } - fn new_user_session(session_token: &str, virtual_id: u32, provider: &str, stream_url: &str, addr: &SocketAddr, - connection_permission: UserConnectionPermission) -> UserSession { + fn new_user_session( + session_token: &str, + virtual_id: u32, + provider: &str, + stream_url: &str, + addr: &SocketAddr, + connection_permission: UserConnectionPermission, + ) -> UserSession { UserSession { token: session_token.to_string(), virtual_id, @@ -360,9 +390,16 @@ impl ActiveUserManager { } #[allow(clippy::too_many_arguments)] - pub async fn create_user_session(&self, user: &ProxyUserCredentials, session_token: &str, virtual_id: u32, - provider: &str, stream_url: &str, addr: &SocketAddr, - connection_permission: UserConnectionPermission) -> String { + pub async fn create_user_session( + &self, + user: &ProxyUserCredentials, + session_token: &str, + virtual_id: u32, + provider: &str, + stream_url: &str, + addr: &SocketAddr, + connection_permission: UserConnectionPermission, + ) -> String { self.gc(); let username = user.username.clone(); @@ -370,7 +407,8 @@ impl ActiveUserManager { let connection_data = user_connections.by_key.entry(username.clone()).or_insert_with(|| { debug_if_enabled!("Creating first session for user {username} {}", sanitize_sensitive_info(stream_url)); let mut data = UserConnectionData::new(0, user.max_connections); - let session = Self::new_user_session(session_token, virtual_id, provider, stream_url, addr, connection_permission); + let session = + Self::new_user_session(session_token, virtual_id, provider, stream_url, addr, connection_permission); data.add_session(session); data }); @@ -386,15 +424,23 @@ impl ActiveUserManager { session.provider = provider.intern(); } session.permission = connection_permission; - debug_if_enabled!("Using session for user {} with url: {}", user.username, sanitize_sensitive_info(stream_url)); + debug_if_enabled!( + "Using session for user {} with url: {}", + user.username, + sanitize_sensitive_info(stream_url) + ); return session.token.clone(); } } // If no session exists, create one - debug_if_enabled!("Creating session for user {} with url: {}", - user.username, sanitize_sensitive_info(stream_url)); - let session = Self::new_user_session(session_token, virtual_id, provider, stream_url, addr, connection_permission); + debug_if_enabled!( + "Creating session for user {} with url: {}", + user.username, + sanitize_sensitive_info(stream_url) + ); + let session = + Self::new_user_session(session_token, virtual_id, provider, stream_url, addr, connection_permission); let token = session.token.clone(); connection_data.add_session(session); token @@ -412,9 +458,11 @@ impl ActiveUserManager { stream.addr = *addr; } } - debug_if_enabled!("Updated session {token} for {username} address {} -> {}", + debug_if_enabled!( + "Updated session {token} for {username} address {} -> {}", sanitize_sensitive_info(&previous_addr.to_string()), - sanitize_sensitive_info(&addr.to_string())); + sanitize_sensitive_info(&addr.to_string()) + ); } } } @@ -473,17 +521,9 @@ impl ActiveUserManager { .map(|stream| stream.addr.to_string()) .collect::>() .join(", "); - let recent_sockets = if recent_sockets.is_empty() { - String::from("n/a") - } else { - recent_sockets - }; - let unique_clients = connection_data - .streams - .iter() - .map(|stream| &stream.client_ip) - .collect::>() - .len(); + let recent_sockets = if recent_sockets.is_empty() { String::from("n/a") } else { recent_sockets }; + let unique_clients = + connection_data.streams.iter().map(|stream| &stream.client_ip).collect::>().len(); debug!( "User {username} exceeded configured max connections ({}/{}). Unique clients: {}, recent sockets [{}]", active_for_user, @@ -529,9 +569,7 @@ impl ActiveUserManager { let now = current_time_secs(); if now.saturating_sub(ts) > USER_GC_TTL - && gc_ts - .compare_exchange(ts, now, Ordering::AcqRel, Ordering::Relaxed) - .is_ok() + && gc_ts.compare_exchange(ts, now, Ordering::AcqRel, Ordering::Relaxed).is_ok() { if let Ok(mut user_connections) = self.connections.try_write() { user_connections.kicked.retain(|_, (expires_at, _)| *expires_at > now); @@ -572,7 +610,6 @@ impl ActiveUserManager { // .collect(); // - // for handle in handles { // handle.join().unwrap(); // } diff --git a/backend/src/api/model/app_state.rs b/backend/src/api/model/app_state.rs index 4a31810e3..440d11677 100644 --- a/backend/src/api/model/app_state.rs +++ b/backend/src/api/model/app_state.rs @@ -1,30 +1,40 @@ -use crate::api::config_watch::exec_config_watch; -use crate::api::model::UpdateGuard; -use crate::api::model::{ActiveProviderManager, ConnectionManager, EventManager, PlaylistStorage, PlaylistStorageState, SharedStreamManager}; -use crate::api::model::{ActiveUserManager, DownloadQueue}; -use crate::api::scheduler::exec_scheduler; -use crate::model::{AppConfig, Config, ConfigTarget, GracePeriodOptions, HdHomeRunConfig, HdHomeRunDeviceConfig, ProcessTargets, ReverseProxyDisabledHeaderConfig, ScheduleConfig, SourcesConfig}; -use crate::repository::get_geoip_path; -use crate::repository::load_target_into_memory_cache; -use crate::tools::lru_cache::LRUResourceCache; -use crate::utils::request::{create_client, create_client_with_redirect}; -use crate::utils::GeoIp; +use crate::{ + api::{ + config_watch::exec_config_watch, + model::{ + metadata_update_manager::MetadataUpdateManager, ActiveProviderManager, ActiveUserManager, + ConnectionManager, DownloadQueue, EventManager, PlaylistStorage, PlaylistStorageState, SharedStreamManager, + UpdateGuard, + }, + scheduler::exec_scheduler, + }, + model::{ + AppConfig, Config, ConfigTarget, GracePeriodOptions, HdHomeRunConfig, HdHomeRunDeviceConfig, ProcessTargets, + ReverseProxyDisabledHeaderConfig, ScheduleConfig, SourcesConfig, + }, + repository::{get_geoip_path, load_target_into_memory_cache}, + tools::lru_cache::LRUResourceCache, + utils::{ + request::{create_client, create_client_with_redirect}, + GeoIp, + }, +}; use arc_swap::{ArcSwap, ArcSwapOption}; use log::{error, info}; use reqwest::Client; -use shared::create_bitset; -use shared::error::TuliproxError; -use shared::info_err_res; -use shared::model::UserConnectionPermission; -use shared::utils::small_vecs_equal_unordered; -use std::collections::HashMap; -use std::sync::atomic::AtomicI8; -use std::sync::Arc; -use std::time::Duration; -use tokio::sync::Mutex; -use tokio::task; +use shared::{ + create_bitset, error::TuliproxError, info_err_res, model::UserConnectionPermission, + utils::small_vecs_equal_unordered, +}; +use std::{ + collections::HashMap, + ffi::OsStr, + sync::{atomic::AtomicI8, Arc}, + time::Duration, +}; +use tokio::{sync::Mutex, task}; use tokio_util::sync::CancellationToken; -use crate::api::model::metadata_update_manager::MetadataUpdateManager; +use url::Url; macro_rules! cancel_service { ($field: ident, $flag:expr, $changes:expr, $cancel_tokens:expr) => { @@ -67,9 +77,7 @@ pub(in crate::api) struct UpdateChanges { } impl UpdateChanges { - pub(in crate::api) fn modified(&self) -> bool { - !self.flags.is_empty() - } + pub(in crate::api) fn modified(&self) -> bool { !self.flags.is_empty() } fn set_flag_if(&mut self, condition: bool, flag: UpdateChangesFlags) { if condition { @@ -78,10 +86,7 @@ impl UpdateChanges { } } -async fn update_target_caches( - app_state: &Arc, - target_changes: Option<&HashMap>, -) { +async fn update_target_caches(app_state: &Arc, target_changes: Option<&HashMap>) { if let Some(target_changes) = target_changes { let mut to_remove = Vec::new(); for target in target_changes.values() { @@ -112,19 +117,13 @@ async fn update_target_caches( } } -pub async fn update_app_state_config( - app_state: &Arc, - config: Config, -) -> Result<(), TuliproxError> { +pub async fn update_app_state_config(app_state: &Arc, config: Config) -> Result<(), TuliproxError> { let updates = app_state.set_config(config).await?; restart_services(app_state, &updates); Ok(()) } -pub async fn update_app_state_sources( - app_state: &Arc, - sources: SourcesConfig, -) -> Result<(), TuliproxError> { +pub async fn update_app_state_sources(app_state: &Arc, sources: SourcesConfig) -> Result<(), TuliproxError> { let targets = sources.validate_targets(Some(&app_state.forced_targets.load().target_names))?; app_state.forced_targets.store(Arc::new(targets)); let updates = app_state.set_sources(sources).await?; @@ -151,11 +150,7 @@ fn cancel_services(app_state: &Arc, changes: &UpdateChanges) { let hdhomerun = cancel_service!(hdhomerun, UpdateChangesFlags::Hdhomerun, changes, cancel_tokens); let file_watch = cancel_service!(file_watch, UpdateChangesFlags::FileWatch, changes, cancel_tokens); - let tokens = CancelTokens { - scheduler, - hdhomerun, - file_watch, - }; + let tokens = CancelTokens { scheduler, hdhomerun, file_watch }; app_state.cancel_tokens.store(Arc::new(tokens)); } @@ -209,9 +204,7 @@ pub fn create_http_client(app_config: &AppConfig) -> Result Result { +pub fn create_http_client_no_redirect(app_config: &AppConfig) -> Result { let builder = create_client_with_redirect(app_config, reqwest::redirect::Policy::none()).http1_only(); let config = app_config.config.load(); build_http_client_with_fallback( @@ -243,9 +236,7 @@ fn build_http_client_with_fallback( let proxy_configured = config.proxy.is_some(); if config.connect_timeout_secs > 0 { - builder = builder.connect_timeout(Duration::from_secs( - u64::from(config.connect_timeout_secs), - )); + builder = builder.connect_timeout(Duration::from_secs(u64::from(config.connect_timeout_secs))); } if let Ok(client) = builder.build() { @@ -262,17 +253,13 @@ fn build_http_client_with_fallback( } pub fn create_cache(config: &Config) -> Option>> { - let lru_cache = config - .reverse_proxy - .as_ref() - .and_then(|r| r.cache.as_ref()) - .and_then(|c| { - if c.enabled { - Some(LRUResourceCache::new(c.size, c.dir.as_str())) - } else { - None - } - }); + let lru_cache = config.reverse_proxy.as_ref().and_then(|r| r.cache.as_ref()).and_then(|c| { + if c.enabled { + Some(LRUResourceCache::new(c.size, c.dir.as_str())) + } else { + None + } + }); let cache_enabled = lru_cache.is_some(); if cache_enabled { info!("Scanning cache"); @@ -396,7 +383,10 @@ impl AppState { Ok(()) } - pub(in crate::api::model) async fn set_sources(&self, sources: SourcesConfig) -> Result { + pub(in crate::api::model) async fn set_sources( + &self, + sources: SourcesConfig, + ) -> Result { let changes = self.detect_changes_for_sources(&sources); self.app_config.set_sources(sources)?; self.active_provider.update_config(&self.app_config).await; @@ -409,14 +399,8 @@ impl AppState { self.active_users.user_connections(username).await } - pub async fn get_connection_permission( - &self, - username: &str, - max_connections: u32, - ) -> UserConnectionPermission { - self.active_users - .connection_permission(username, max_connections) - .await + pub async fn get_connection_permission(&self, username: &str, max_connections: u32) -> UserConnectionPermission { + self.active_users.connection_permission(username, max_connections).await } fn detect_changes_for_config(&self, config: &Config) -> UpdateChanges { @@ -425,23 +409,14 @@ impl AppState { change_detect!(schedules_changed, old_config.schedules.as_ref(), config.schedules.as_ref()); let changed_hdhomerun = change_detect!(hdhomerun_changed, old_config.hdhomerun.as_ref(), config.hdhomerun.as_ref()); - let changed_file_watch = change_detect!( - string_changed, - old_config.mapping_path.as_ref(), - config.mapping_path.as_ref() - ) || change_detect!( - string_changed, - old_config.template_path.as_ref(), - config.template_path.as_ref() - ); + let changed_file_watch = + change_detect!(string_changed, old_config.mapping_path.as_ref(), config.mapping_path.as_ref()) + || change_detect!(string_changed, old_config.template_path.as_ref(), config.template_path.as_ref()); let geoip_enabled = config.is_geoip_enabled(); let geoip_enabled_old = old_config.is_geoip_enabled(); - let mut changes = UpdateChanges { - flags: UpdateChangesFlagsSet::new(), - targets: None, - }; + let mut changes = UpdateChanges { flags: UpdateChangesFlagsSet::new(), targets: None }; changes.set_flag_if(changed_schedules, UpdateChangesFlags::Scheduler); changes.set_flag_if(changed_hdhomerun, UpdateChangesFlags::Hdhomerun); changes.set_flag_if(changed_file_watch, UpdateChangesFlags::FileWatch); @@ -493,12 +468,8 @@ impl AppState { Some(changes) => { changes.status = TargetStatus::Keep; changes.cache_status = match (changes.cache_status, target.use_memory_cache) { - (TargetCacheState::UnchangedFalse, true) => { - TargetCacheState::ChangedToTrue - } - (TargetCacheState::UnchangedTrue, false) => { - TargetCacheState::ChangedToFalse - } + (TargetCacheState::UnchangedFalse, true) => TargetCacheState::ChangedToTrue, + (TargetCacheState::UnchangedTrue, false) => TargetCacheState::ChangedToFalse, (x, _) => x, }; } @@ -509,10 +480,7 @@ impl AppState { (file_watch_changed, target_changes) }; - let mut changes = UpdateChanges { - flags: UpdateChangesFlagsSet::new(), - targets: Some(target_changes), - }; + let mut changes = UpdateChanges { flags: UpdateChangesFlagsSet::new(), targets: Some(target_changes) }; changes.set_flag_if(file_watch_changed, UpdateChangesFlags::FileWatch); changes } @@ -525,26 +493,67 @@ impl AppState { self.app_config.get_disabled_headers() } - pub fn get_grace_options(&self) -> GracePeriodOptions { - self.app_config.get_grace_options() - } + pub fn get_grace_options(&self) -> GracePeriodOptions { self.app_config.get_grace_options() } pub fn should_use_manual_redirects(&self) -> bool { let config = self.app_config.config.load(); - config.proxy.is_some() || proxy_env_present() + config.proxy.as_ref().is_some_and(|proxy| should_use_manual_redirect_for_proxy(proxy.url.as_str())) + || proxy_env_present() } } -fn proxy_env_present() -> bool { - const ENV_KEYS: [&str; 3] = [ - "HTTP_PROXY", - "HTTPS_PROXY", - "ALL_PROXY", - ]; +fn proxy_env_present() -> bool { should_use_manual_redirects_for_env_vars(std::env::vars_os()) } - std::env::vars().any(|(key, value)| { - ENV_KEYS.iter().any(|k| k.eq_ignore_ascii_case(&key)) - && !value.trim().is_empty() +fn parse_proxy_url_with_http_fallback(proxy_url: &str) -> Option { + let trimmed = proxy_url.trim(); + if trimmed.is_empty() { + return None; + } + + if let Ok(url) = Url::parse(trimmed) { + if matches!(url.scheme().to_ascii_lowercase().as_str(), "http" | "https") { + return Some(url); + } + if trimmed.contains("://") { + return None; + } + } + + if trimmed.contains("://") { + return None; + } + if trimmed.starts_with('/') || trimmed.starts_with('\\') { + return None; + } + + Url::parse(format!("http://{trimmed}").as_str()).ok() +} + +fn should_use_manual_redirect_for_proxy(proxy_url: &str) -> bool { + parse_proxy_url_with_http_fallback(proxy_url).is_some_and(|url| { + matches!(url.scheme().to_ascii_lowercase().as_str(), "http" | "https") && url.host_str().is_some() + }) +} + +fn should_use_manual_redirects_for_env_vars(vars: I) -> bool +where + I: IntoIterator, + K: AsRef, + V: AsRef, +{ + const ENV_KEYS: [&str; 3] = ["HTTP_PROXY", "HTTPS_PROXY", "ALL_PROXY"]; + + vars.into_iter().any(|(key, value)| { + let Some(key) = key.as_ref().to_str() else { + return false; + }; + let Some(value) = value.as_ref().to_str() else { + return false; + }; + let value = value.trim(); + ENV_KEYS.iter().any(|candidate| candidate.eq_ignore_ascii_case(key)) + && !value.is_empty() + && should_use_manual_redirect_for_proxy(value) }) } @@ -579,9 +588,7 @@ fn hdhomerun_changed(a: &HdHomeRunConfig, b: &HdHomeRunConfig) -> bool { false } -fn string_changed(a: &str, b: &str) -> bool { - a != b -} +fn string_changed(a: &str, b: &str) -> bool { a != b } #[derive(Clone)] pub struct HdHomerunAppState { @@ -589,3 +596,44 @@ pub struct HdHomerunAppState { pub device: Arc, pub hd_scan_state: Arc, } + +#[cfg(test)] +mod tests { + use super::{should_use_manual_redirect_for_proxy, should_use_manual_redirects_for_env_vars}; + + #[test] + fn should_use_manual_redirect_for_proxy_only_http_or_https() { + assert!(should_use_manual_redirect_for_proxy("http://proxy.local:8080")); + assert!(should_use_manual_redirect_for_proxy("https://proxy.local:8443")); + assert!(should_use_manual_redirect_for_proxy("proxy.local:8080")); + assert!(should_use_manual_redirect_for_proxy("127.0.0.1:8888")); + assert!(!should_use_manual_redirect_for_proxy("socks5://proxy.local:1080")); + assert!(!should_use_manual_redirect_for_proxy("socks5h://proxy.local:1080")); + assert!(!should_use_manual_redirect_for_proxy("://invalid")); + assert!(!should_use_manual_redirect_for_proxy("/tmp/proxy.socket")); + } + + #[test] + fn should_use_manual_redirects_for_env_vars_only_when_http_proxy_is_present() { + assert!(should_use_manual_redirects_for_env_vars(vec![( + "HTTP_PROXY".to_string(), + "http://proxy.local:8080".to_string(), + )])); + assert!(should_use_manual_redirects_for_env_vars(vec![( + "all_proxy".to_string(), + "https://proxy.local:8443".to_string(), + )])); + assert!(should_use_manual_redirects_for_env_vars(vec![( + "HTTP_PROXY".to_string(), + "127.0.0.1:8888".to_string(), + )])); + assert!(!should_use_manual_redirects_for_env_vars(vec![( + "ALL_PROXY".to_string(), + "socks5://proxy.local:1080".to_string(), + )])); + assert!(!should_use_manual_redirects_for_env_vars(vec![( + "NO_PROXY".to_string(), + "http://localhost".to_string(), + )])); + } +} diff --git a/backend/src/api/model/batch_result_collector.rs b/backend/src/api/model/batch_result_collector.rs index 61725157b..3b034d53c 100644 --- a/backend/src/api/model/batch_result_collector.rs +++ b/backend/src/api/model/batch_result_collector.rs @@ -1,6 +1,6 @@ -use std::mem; -use shared::model::{VideoStreamProperties, SeriesStreamProperties, LiveStreamProperties}; use crate::api::model::ProviderIdType; +use shared::model::{LiveStreamProperties, SeriesStreamProperties, VideoStreamProperties}; +use std::mem; const BATCH_THRESHOLD: usize = 200; @@ -20,22 +20,14 @@ impl BatchResultCollector { } } - pub fn add_vod(&mut self, id: ProviderIdType, props: VideoStreamProperties) { - self.vod.push((id, props)); - } + pub fn add_vod(&mut self, id: ProviderIdType, props: VideoStreamProperties) { self.vod.push((id, props)); } - pub fn add_series(&mut self, id: ProviderIdType, props: SeriesStreamProperties) { - self.series.push((id, props)); - } + pub fn add_series(&mut self, id: ProviderIdType, props: SeriesStreamProperties) { self.series.push((id, props)); } - pub fn add_live(&mut self, id: ProviderIdType, props: LiveStreamProperties) { - self.live.push((id, props)); - } + pub fn add_live(&mut self, id: ProviderIdType, props: LiveStreamProperties) { self.live.push((id, props)); } pub fn should_flush(&self) -> bool { - self.vod.len() >= BATCH_THRESHOLD || - self.series.len() >= BATCH_THRESHOLD || - self.live.len() >= BATCH_THRESHOLD + self.vod.len() >= BATCH_THRESHOLD || self.series.len() >= BATCH_THRESHOLD || self.live.len() >= BATCH_THRESHOLD } pub fn take_vod_updates(&mut self) -> Vec<(ProviderIdType, VideoStreamProperties)> { @@ -53,7 +45,7 @@ impl BatchResultCollector { mem::take(&mut self.series) } } - + pub fn take_live_updates(&mut self) -> Vec<(ProviderIdType, LiveStreamProperties)> { if self.live.is_empty() { Vec::new() @@ -61,8 +53,6 @@ impl BatchResultCollector { mem::take(&mut self.live) } } - - pub fn is_empty(&self) -> bool { - self.vod.is_empty() && self.series.is_empty() && self.live.is_empty() - } + + pub fn is_empty(&self) -> bool { self.vod.is_empty() && self.series.is_empty() && self.live.is_empty() } } diff --git a/backend/src/api/model/connection_manager.rs b/backend/src/api/model/connection_manager.rs index a6303435b..64fd89d08 100644 --- a/backend/src/api/model/connection_manager.rs +++ b/backend/src/api/model/connection_manager.rs @@ -1,12 +1,17 @@ -use crate::api::model::{ActiveProviderManager, ActiveUserManager, CustomVideoStreamType, EventManager, EventMessage, ProviderHandle, SharedStreamManager}; -use crate::auth::Fingerprint; -use crate::utils::debug_if_enabled; -use log::{warn}; -use shared::model::{ActiveUserConnectionChange, StreamChannel, VirtualId}; -use shared::utils::sanitize_sensitive_info; -use std::borrow::Cow; -use std::net::SocketAddr; -use std::sync::Arc; +use crate::{ + api::model::{ + ActiveProviderManager, ActiveUserManager, CustomVideoStreamType, EventManager, EventMessage, ProviderHandle, + SharedStreamManager, + }, + auth::Fingerprint, + utils::debug_if_enabled, +}; +use log::warn; +use shared::{ + model::{ActiveUserConnectionChange, StreamChannel, VirtualId}, + utils::sanitize_sensitive_info, +}; +use std::{borrow::Cow, net::SocketAddr, sync::Arc}; pub struct ConnectionManager { pub user_manager: Arc, @@ -39,13 +44,19 @@ impl ConnectionManager { } pub async fn kick_connection(&self, addr: &SocketAddr, virtual_id: VirtualId, block_secs: u64) -> bool { - debug_if_enabled!("User {} kicked for stream with virtual_id {virtual_id} for {block_secs} seconds with addr {}.", - self.user_manager.get_username_for_addr(addr).await.unwrap_or_default(), sanitize_sensitive_info(&addr.to_string())); + debug_if_enabled!( + "User {} kicked for stream with virtual_id {virtual_id} for {block_secs} seconds with addr {}.", + self.user_manager.get_username_for_addr(addr).await.unwrap_or_default(), + sanitize_sensitive_info(&addr.to_string()) + ); if block_secs > 0 { self.user_manager.block_user_for_stream(addr, virtual_id, block_secs).await; } if let Err(e) = self.close_socket_signal_tx.send(*addr) { - debug_if_enabled!("No active receivers for close signal ({}): {e:?}", sanitize_sensitive_info(&addr.to_string())); + debug_if_enabled!( + "No active receivers for close signal ({}): {e:?}", + sanitize_sensitive_info(&addr.to_string()) + ); return false; } true @@ -70,14 +81,32 @@ impl ConnectionManager { } #[allow(clippy::too_many_arguments)] - pub async fn add_connection(&self, addr: &SocketAddr) { - self.user_manager.add_connection(addr).await; - } + pub async fn add_connection(&self, addr: &SocketAddr) { self.user_manager.add_connection(addr).await; } #[allow(clippy::too_many_arguments)] - pub async fn update_connection(&self, username: &str, max_connections: u32, fingerprint: &Fingerprint, - provider: &str, stream_channel: StreamChannel, user_agent: Cow<'_, str>, session_token: Option<&str>) { - if let Some(stream_info) = self.user_manager.update_connection(username, max_connections, fingerprint, provider, stream_channel, user_agent, session_token).await { + pub async fn update_connection( + &self, + username: &str, + max_connections: u32, + fingerprint: &Fingerprint, + provider: &str, + stream_channel: StreamChannel, + user_agent: Cow<'_, str>, + session_token: Option<&str>, + ) { + if let Some(stream_info) = self + .user_manager + .update_connection( + username, + max_connections, + fingerprint, + provider, + stream_channel, + user_agent, + session_token, + ) + .await + { self.event_manager.send_event(EventMessage::ActiveUser(ActiveUserConnectionChange::Updated(stream_info))); } else { warn!("Failed to register connection for user {username} at {}; disconnecting client", fingerprint.addr); diff --git a/backend/src/api/model/download.rs b/backend/src/api/model/download.rs index ccafd5086..3d1c48948 100644 --- a/backend/src/api/model/download.rs +++ b/backend/src/api/model/download.rs @@ -1,12 +1,13 @@ -use std::collections::VecDeque; -use std::ffi::OsStr; -use std::path::{Path, PathBuf}; -use tokio::sync::{RwLock, Mutex}; -use std::sync::Arc; +use crate::model::VideoDownloadConfig; use serde::{Deserialize, Serialize}; - -use crate::model::{VideoDownloadConfig}; -use shared::utils::{CONSTANTS, FILENAME_TRIM_PATTERNS, hash_string_as_hex, deunicode_string}; +use shared::utils::{deunicode_string, hash_string_as_hex, CONSTANTS, FILENAME_TRIM_PATTERNS}; +use std::{ + collections::VecDeque, + ffi::OsStr, + path::{Path, PathBuf}, + sync::Arc, +}; +use tokio::sync::{Mutex, RwLock}; /// File-Download information. #[derive(Clone)] @@ -59,7 +60,6 @@ fn get_download_directory(download_cfg: &VideoDownloadConfig, filestem: &str) -> } impl FileDownload { - // TODO read header size info and restart support // "content-type" => ".../..." // "content-length" => "1975828544" @@ -69,12 +69,17 @@ impl FileDownload { pub fn new(req_url: &str, req_filename: &str, download_cfg: &VideoDownloadConfig) -> Option { match reqwest::Url::parse(req_url) { Ok(url) => { - let tmp_filename = CONSTANTS.re_filename.replace_all(&deunicode_string(req_filename) - .replace(' ', "_"), "") + let tmp_filename = CONSTANTS + .re_filename + .replace_all(&deunicode_string(req_filename).replace(' ', "_"), "") .replace("__", "_") .replace("_-_", "-"); let filename_path = Path::new(&tmp_filename); - let file_stem = filename_path.file_stem().and_then(OsStr::to_str).unwrap_or("").trim_matches(FILENAME_TRIM_PATTERNS); + let file_stem = filename_path + .file_stem() + .and_then(OsStr::to_str) + .unwrap_or("") + .trim_matches(FILENAME_TRIM_PATTERNS); let file_ext = filename_path.extension().and_then(OsStr::to_str).unwrap_or(""); let mut filename = format!("{file_stem}.{file_ext}"); @@ -102,7 +107,7 @@ impl FileDownload { error: None, }) } - Err(_) => None + Err(_) => None, } } } @@ -114,9 +119,7 @@ pub struct DownloadQueue { } impl Default for DownloadQueue { - fn default() -> Self { - Self::new() - } + fn default() -> Self { Self::new() } } impl DownloadQueue { @@ -129,7 +132,6 @@ impl DownloadQueue { } } - #[derive(Deserialize, Serialize, Debug, Clone)] pub struct FileDownloadRequest { pub url: String, diff --git a/backend/src/api/model/event_manager.rs b/backend/src/api/model/event_manager.rs index 843e9a250..295318715 100644 --- a/backend/src/api/model/event_manager.rs +++ b/backend/src/api/model/event_manager.rs @@ -1,13 +1,13 @@ -use std::sync::Arc; -use log::{trace}; +use log::trace; use shared::model::{ActiveUserConnectionChange, ConfigType, LibraryScanSummary, PlaylistUpdateState, SystemInfo}; +use std::sync::Arc; #[allow(clippy::large_enum_variant)] #[derive(Clone, PartialEq)] pub enum EventMessage { ServerError(String), ActiveUser(ActiveUserConnectionChange), // user_count, connection count - ActiveProvider(Arc, usize), // provider name, connections + ActiveProvider(Arc, usize), // provider name, connections ConfigChange(ConfigType), PlaylistUpdate(PlaylistUpdateState), PlaylistUpdateProgress(String, String), @@ -33,9 +33,7 @@ impl EventManager { } } - pub fn get_event_channel(&self) -> tokio::sync::broadcast::Receiver { - self.channel_tx.subscribe() - } + pub fn get_event_channel(&self) -> tokio::sync::broadcast::Receiver { self.channel_tx.subscribe() } pub fn send_event(&self, event: EventMessage) -> bool { if let Err(err) = self.channel_tx.send(event) { @@ -60,7 +58,5 @@ impl EventManager { } impl Default for EventManager { - fn default() -> Self { - Self::new() - } -} \ No newline at end of file + fn default() -> Self { Self::new() } +} diff --git a/backend/src/api/model/mod.rs b/backend/src/api/model/mod.rs index b249f7c4c..550cd2551 100644 --- a/backend/src/api/model/mod.rs +++ b/backend/src/api/model/mod.rs @@ -1,37 +1,28 @@ +mod active_provider_manager; +mod active_user_manager; mod app_state; -mod request; +mod connection_manager; mod download; -mod xtream; +mod event_manager; +mod metadata_update_manager; mod model_utils; +mod playlist_mem_cache; +mod provider_config; +mod provider_lineup_manager; +mod request; +mod stream; mod stream_error; mod streams; -mod active_user_manager; -mod active_provider_manager; -mod stream; -mod provider_config; -mod event_manager; -mod playlist_mem_cache; -mod provider_lineup_manager; -mod connection_manager; mod update_guard; -mod metadata_update_manager; +mod xtream; -pub use self::active_provider_manager::*; -pub(in crate::api) use self::active_user_manager::*; -pub use self::app_state::*; -pub use self::connection_manager::*; -pub(in crate::api) use self::download::*; -pub use self::event_manager::*; -pub(in crate::api) use self::model_utils::*; -pub use self::playlist_mem_cache::*; -pub(in crate::api) use self::provider_config::*; -pub use self::provider_lineup_manager::*; -pub(in crate::api) use self::request::*; -pub use self::stream::*; -pub(in crate::api) use self::stream_error::*; pub(crate) use self::streams::*; -pub(in crate::api) use self::xtream::*; -pub use self::update_guard::*; -pub use self::metadata_update_manager::*; +pub use self::{ + active_provider_manager::*, app_state::*, connection_manager::*, event_manager::*, metadata_update_manager::*, + playlist_mem_cache::*, provider_lineup_manager::*, stream::*, update_guard::*, +}; +pub(in crate::api) use self::{ + active_user_manager::*, download::*, model_utils::*, provider_config::*, request::*, stream_error::*, xtream::*, +}; mod batch_result_collector; pub use self::batch_result_collector::*; diff --git a/backend/src/api/model/model_utils.rs b/backend/src/api/model/model_utils.rs index f3e406f5b..599b718a6 100644 --- a/backend/src/api/model/model_utils.rs +++ b/backend/src/api/model/model_utils.rs @@ -1,19 +1,18 @@ -use reqwest::{StatusCode}; -use std::collections::{HashSet}; -use std::str::FromStr; -use reqwest::header::HeaderMap; -use shared::utils::{filter_response_header}; +use reqwest::{header::HeaderMap, StatusCode}; +use shared::utils::filter_response_header; +use std::{collections::HashSet, str::FromStr}; pub fn get_response_headers(headers: &HeaderMap) -> Vec<(String, String)> { - headers.iter() + headers + .iter() .filter(|(key, _)| filter_response_header(key.as_str())) - .filter_map(|(key, value)| { - value.to_str().ok().map(|v| (key.to_string(), v.to_string())) - }) + .filter_map(|(key, value)| value.to_str().ok().map(|v| (key.to_string(), v.to_string()))) .collect() } -pub fn get_stream_response_with_headers(custom: Option<(Vec<(String, String)>, StatusCode)>) -> (axum::http::StatusCode, axum::http::HeaderMap) { +pub fn get_stream_response_with_headers( + custom: Option<(Vec<(String, String)>, StatusCode)>, +) -> (axum::http::StatusCode, axum::http::HeaderMap) { let mut headers = HeaderMap::new(); let mut added_headers: HashSet = HashSet::new(); let mut status = StatusCode::OK; @@ -21,20 +20,22 @@ pub fn get_stream_response_with_headers(custom: Option<(Vec<(String, String)>, S if let Some((custom_headers, status_code)) = custom { status = status_code; for (key, value) in custom_headers { - if let (Ok(name), Ok(val)) = (axum::http::HeaderName::from_str(&key), axum::http::HeaderValue::from_str(&value)) { + if let (Ok(name), Ok(val)) = + (axum::http::HeaderName::from_str(&key), axum::http::HeaderValue::from_str(&value)) + { headers.insert(name.clone(), val); added_headers.insert(key); } } } - let default_headers = vec![ - ("content-type", "application/octet-stream"), - ]; + let default_headers = vec![("content-type", "application/octet-stream")]; for (key, value) in default_headers { if !added_headers.contains(key) { - if let (Ok(name), Ok(val)) = (axum::http::HeaderName::from_str(key), axum::http::HeaderValue::from_str(value)) { + if let (Ok(name), Ok(val)) = + (axum::http::HeaderName::from_str(key), axum::http::HeaderValue::from_str(value)) + { headers.insert(name, val); } } @@ -45,4 +46,4 @@ pub fn get_stream_response_with_headers(custom: Option<(Vec<(String, String)>, S } (status, headers) -} \ No newline at end of file +} diff --git a/backend/src/api/model/playlist_mem_cache.rs b/backend/src/api/model/playlist_mem_cache.rs index 857c2b78d..afb5f7e63 100644 --- a/backend/src/api/model/playlist_mem_cache.rs +++ b/backend/src/api/model/playlist_mem_cache.rs @@ -1,8 +1,10 @@ +use crate::{ + model::ConfigTarget, + repository::{BPlusTree, VirtualIdRecord}, +}; +use shared::model::{M3uPlaylistItem, PlaylistItem, VirtualId, XtreamCluster, XtreamPlaylistItem}; use std::collections::HashMap; use tokio::sync::RwLock; -use shared::model::{M3uPlaylistItem, PlaylistItem, VirtualId, XtreamCluster, XtreamPlaylistItem}; -use crate::model::ConfigTarget; -use crate::repository::{BPlusTree, VirtualIdRecord}; pub struct PlaylistXtreamStorage { pub live: BPlusTree, @@ -30,12 +32,7 @@ pub struct PlaylistStorageState { } impl PlaylistStorageState { - - pub(crate) fn new() -> Self { - Self { - data: RwLock::new(HashMap::new()), - } - } + pub(crate) fn new() -> Self { Self { data: RwLock::new(HashMap::new()) } } pub async fn update_target_id_mapping(&self, target: &ConfigTarget, mapping: Vec) { if target.use_memory_cache { @@ -57,8 +54,9 @@ impl PlaylistStorageState { match pli.xtream_cluster { XtreamCluster::Live => &mut xtream.live, XtreamCluster::Video => &mut xtream.vod, - XtreamCluster::Series => &mut xtream.series, - }.insert(pli.virtual_id, pli.clone()); + XtreamCluster::Series => &mut xtream.series, + } + .insert(pli.virtual_id, pli.clone()); } } } @@ -73,8 +71,9 @@ impl PlaylistStorageState { match pli.header.xtream_cluster { XtreamCluster::Live => &mut xtream.live, XtreamCluster::Video => &mut xtream.vod, - XtreamCluster::Series => &mut xtream.series, - }.insert(pli.header.virtual_id, XtreamPlaylistItem::from(&pli)); + XtreamCluster::Series => &mut xtream.series, + } + .insert(pli.header.virtual_id, XtreamPlaylistItem::from(&pli)); } } } @@ -88,11 +87,7 @@ impl PlaylistStorageState { storage.id_mapping = Some(id_mapping); } std::collections::hash_map::Entry::Vacant(entry) => { - entry.insert(TargetPlaylistStorage { - xtream: None, - m3u: None, - id_mapping: Some(id_mapping), - }); + entry.insert(TargetPlaylistStorage { xtream: None, m3u: None, id_mapping: Some(id_mapping) }); } } } diff --git a/backend/src/api/model/provider_config.rs b/backend/src/api/model/provider_config.rs index 23c268315..640f3216c 100644 --- a/backend/src/api/model/provider_config.rs +++ b/backend/src/api/model/provider_config.rs @@ -1,16 +1,13 @@ -use std::fmt; -use crate::model::{is_input_expired, ConfigInput, ConfigInputAlias, InputUserInfo}; +use crate::{ + api::model::ProviderAllocation, + model::{is_input_expired, ConfigInput, ConfigInputAlias, InputUserInfo}, + utils::debug_if_enabled, +}; use jsonwebtoken::get_current_timestamp; -use log::{debug}; -use std::ops::Deref; -use std::sync::{Arc}; +use log::debug; +use shared::{model::InputType, utils::sanitize_sensitive_info, write_if_some}; +use std::{fmt, ops::Deref, sync::Arc}; use tokio::sync::RwLock; -use shared::model::InputType; -use shared::utils::sanitize_sensitive_info; -use shared::write_if_some; -use crate::api::model::ProviderAllocation; -use crate::utils::debug_if_enabled; - pub type ProviderConnectionChangeCallback = Arc, usize) + Send + Sync>; @@ -69,9 +66,7 @@ impl fmt::Display for ProviderConfig { } impl fmt::Debug for ProviderConfig { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - write!(f, "{self}") - } + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{self}") } } impl PartialEq for ProviderConfig { @@ -85,7 +80,7 @@ impl PartialEq for ProviderConfig { && self.max_connections == other.max_connections && self.priority == other.priority && self.exp_date == other.exp_date - // Note: self.connection is skipped + // Note: self.connection is skipped } } @@ -101,7 +96,11 @@ macro_rules! modify_connections { } impl ProviderConfig { - pub fn new(cfg: &ConfigInput, connection: Arc>, on_connection_change: ProviderConnectionChangeCallback) -> Self { + pub fn new( + cfg: &ConfigInput, + connection: Arc>, + on_connection_change: ProviderConnectionChangeCallback, + ) -> Self { let panel_api_enabled = cfg.panel_api.as_ref().is_some_and(|panel_api| panel_api.enabled); // Logic change: panel api accounts are not considering unlimited provider access! let effective_max_connections = if panel_api_enabled && cfg.max_connections == 0 { @@ -124,7 +123,7 @@ impl ProviderConfig { priority: cfg.priority, exp_date: cfg.exp_date, connection, - on_connection_change + on_connection_change, } } @@ -160,14 +159,10 @@ impl ProviderConfig { } #[inline] - pub fn max_connections(&self) -> usize { - self.max_connections - } + pub fn max_connections(&self) -> usize { self.max_connections } #[inline] - pub(crate) fn exp_date(&self) -> Option { - self.exp_date - } + pub(crate) fn exp_date(&self) -> Option { self.exp_date } pub fn get_user_info(&self) -> Option { InputUserInfo::new(self.input_type, self.username.as_deref(), self.password.as_deref(), &self.url) @@ -244,8 +239,9 @@ impl ProviderConfig { } let now = get_current_timestamp(); - if guard.granted_grace && now - guard.grace_ts <= grace_period_timeout_secs { - if guard.current_connections > self.max_connections && now - guard.grace_ts <= grace_period_timeout_secs { + if guard.granted_grace && now - guard.grace_ts <= grace_period_timeout_secs { + if guard.current_connections > self.max_connections && 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; @@ -290,7 +286,10 @@ impl ProviderConfig { if guard.granted_grace { if connections > self.max_connections && now - guard.grace_ts <= grace_period_timeout_secs { // Grace timeout still active, deny connection - debug!("Provider access denied, grace exhausted, too many connections, no connection available: {}", self.name); + debug!( + "Provider access denied, grace exhausted, too many connections, no connection available: {}", + self.name + ); return false; } // Grace timeout expired, reset grace counters @@ -314,14 +313,10 @@ impl ProviderConfig { } #[inline] - pub(crate) async fn get_current_connections(&self) -> usize { - self.connection.read().await.current_connections - } + 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 - } + pub(crate) fn get_priority(&self) -> i16 { self.priority } } #[derive(Clone, Debug)] @@ -330,17 +325,11 @@ pub(in crate::api::model) struct ProviderConfigWrapper { } impl fmt::Display for ProviderConfigWrapper { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - write!(f, "{}", self.inner) - } + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{}", self.inner) } } impl ProviderConfigWrapper { - pub fn new(cfg: ProviderConfig) -> Self { - Self { - inner: Arc::new(cfg) - } - } + pub fn new(cfg: ProviderConfig) -> Self { Self { inner: Arc::new(cfg) } } pub async fn force_allocate(&self) -> ProviderAllocation { if self.inner.force_allocate().await { @@ -368,7 +357,5 @@ impl ProviderConfigWrapper { impl Deref for ProviderConfigWrapper { type Target = ProviderConfig; - fn deref(&self) -> &Self::Target { - &self.inner - } + fn deref(&self) -> &Self::Target { &self.inner } } diff --git a/backend/src/api/model/provider_lineup_manager.rs b/backend/src/api/model/provider_lineup_manager.rs index cf9904ec8..aac3419a3 100644 --- a/backend/src/api/model/provider_lineup_manager.rs +++ b/backend/src/api/model/provider_lineup_manager.rs @@ -1,20 +1,31 @@ -use crate::api::model::provider_config::ProviderConfigWrapper; -use crate::api::model::{EventManager, ProviderConfig, ProviderConfigConnection, ProviderConnectionChangeCallback}; -use crate::model::{is_input_expired, ConfigInput, GracePeriodOptions}; -use crate::utils::debug_if_enabled; +use crate::{ + api::model::{ + provider_config::ProviderConfigWrapper, EventManager, ProviderConfig, ProviderConfigConnection, + ProviderConnectionChangeCallback, + }, + model::{is_input_expired, ConfigInput, GracePeriodOptions}, + utils::debug_if_enabled, +}; use arc_swap::ArcSwap; use dashmap::DashMap; use log::{debug, log_enabled}; use shared::utils::{display_vec, sanitize_sensitive_info}; -use std::collections::HashMap; -use std::fmt; -use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; -use std::sync::Arc; +use std::{ + collections::HashMap, + fmt, + sync::{ + atomic::{AtomicU64, AtomicUsize, Ordering}, + Arc, + }, +}; use tokio::sync::RwLock; macro_rules! gen_provider_search { ($fn_name:ident, $field: ident, $crit_type:ty) => { - fn $fn_name<'a>(criteria: $crit_type, providers: &'a Vec) -> Option<(&'a ProviderLineup, &'a ProviderConfigWrapper)> { + fn $fn_name<'a>( + criteria: $crit_type, + providers: &'a Vec, + ) -> Option<(&'a ProviderLineup, &'a ProviderConfigWrapper)> { for lineup in providers { match lineup { ProviderLineup::Single(single) => { @@ -44,7 +55,7 @@ macro_rules! gen_provider_search { } None } - } + }; } fn get_or_create_provider_connection( @@ -73,49 +84,35 @@ impl ProviderAllocation { } } - pub fn new_available(config: Arc) -> Self { - ProviderAllocation::Available(config) - } + pub fn new_available(config: Arc) -> Self { ProviderAllocation::Available(config) } - pub fn new_grace_period(config: Arc) -> Self { - ProviderAllocation::GracePeriod(config) - } + pub fn new_grace_period(config: Arc) -> Self { ProviderAllocation::GracePeriod(config) } pub fn get_provider_name(&self) -> Option> { match self { ProviderAllocation::Exhausted => None, - ProviderAllocation::Available(ref cfg) | - ProviderAllocation::GracePeriod(ref cfg) => { - Some(cfg.name.clone()) - } + ProviderAllocation::Available(ref cfg) | ProviderAllocation::GracePeriod(ref cfg) => Some(cfg.name.clone()), } } pub fn get_provider_id(&self) -> Option { match self { ProviderAllocation::Exhausted => None, - ProviderAllocation::Available(ref cfg) | - ProviderAllocation::GracePeriod(ref cfg) => { - Some(cfg.id) - } + ProviderAllocation::Available(ref cfg) | ProviderAllocation::GracePeriod(ref cfg) => Some(cfg.id), } } pub fn get_provider_config(&self) -> Option> { match self { ProviderAllocation::Exhausted => None, - ProviderAllocation::Available(ref cfg) | - ProviderAllocation::GracePeriod(ref cfg) => { - Some(Arc::clone(cfg)) - } + ProviderAllocation::Available(ref cfg) | ProviderAllocation::GracePeriod(ref cfg) => Some(Arc::clone(cfg)), } } pub async fn release(&self) { match &self { ProviderAllocation::Exhausted => {} - ProviderAllocation::Available(config) | - ProviderAllocation::GracePeriod(config) => { + ProviderAllocation::Available(config) | ProviderAllocation::GracePeriod(config) => { config.release().await; } } @@ -185,7 +182,11 @@ struct SingleProviderLineup { } impl SingleProviderLineup { - fn new(cfg: &ConfigInput, connection: Arc>, connection_change: &ProviderConnectionChangeCallback) -> Self { + fn new( + cfg: &ConfigInput, + connection: Arc>, + connection_change: &ProviderConnectionChangeCallback, + ) -> Self { Self { provider: ProviderConfigWrapper::new(ProviderConfig::new(cfg, connection, Arc::clone(connection_change))), } @@ -207,7 +208,6 @@ impl SingleProviderLineup { } } - /// Manages provider groups based on priority: /// /// `SingleProviderGroup(ProviderConfig)`: A single provider. @@ -246,7 +246,11 @@ impl MultiProviderLineup { connection_change: &ProviderConnectionChangeCallback, ) -> Self { let input_connection = get_or_create_provider_connection(provider_connections, &cfg_input.name); - let mut inputs = vec![ProviderConfigWrapper::new(ProviderConfig::new(cfg_input, input_connection, Arc::clone(connection_change)))]; + let mut inputs = vec![ProviderConfigWrapper::new(ProviderConfig::new( + cfg_input, + input_connection, + Arc::clone(connection_change), + ))]; if let Some(aliases) = cfg_input.get_enabled_aliases() { for alias in aliases { let alias_connection = get_or_create_provider_connection(provider_connections, &alias.name); @@ -261,24 +265,22 @@ impl MultiProviderLineup { let mut providers = HashMap::new(); for provider in inputs { let priority = provider.get_priority(); - providers.entry(priority) - .or_insert_with(Vec::new) - .push(provider); + providers.entry(priority).or_insert_with(Vec::new).push(provider); } let mut values: Vec<(i16, Vec)> = providers.into_iter().collect(); values.sort_by_key(|(p1, _)| *p1); - let providers: Vec = values.into_iter().map(|(_, mut group)| { - if group.len() > 1 { - ProviderPriorityGroup::MultiProviderGroup(AtomicUsize::new(0), group) - } else { - ProviderPriorityGroup::SingleProviderGroup(group.remove(0)) - } - }).collect(); + let providers: Vec = values + .into_iter() + .map(|(_, mut group)| { + if group.len() > 1 { + ProviderPriorityGroup::MultiProviderGroup(AtomicUsize::new(0), group) + } else { + ProviderPriorityGroup::SingleProviderGroup(group.remove(0)) + } + }) + .collect(); - Self { - name: cfg_input.name.clone(), - providers, - } + Self { name: cfg_input.name.clone(), providers } } /// Attempts to acquire the next available provider from a specific priority group. @@ -308,7 +310,11 @@ impl MultiProviderLineup { /// } /// } /// ``` - async fn acquire_next_provider_from_group(priority_group: &ProviderPriorityGroup, grace: bool, grace_period_timeout_secs: u64) -> 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, grace_period_timeout_secs).await; @@ -354,7 +360,11 @@ impl MultiProviderLineup { } // Used for redirect to cyclce through provider - async fn get_next_provider_from_group(priority_group: &ProviderPriorityGroup, grace: bool, grace_period_timeout_secs: u64) -> 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, grace_period_timeout_secs).await; @@ -462,7 +472,6 @@ impl MultiProviderLineup { None } - #[cfg(test)] async fn release(&self, provider_name: &Arc) { for g in &self.providers { @@ -519,19 +528,17 @@ struct LineupSnapshot { } impl ProviderLineupManager { - pub fn new(inputs: Vec>, grace_period_options: GracePeriodOptions, event_manager: &Arc) -> Self { + pub fn new( + inputs: Vec>, + grace_period_options: GracePeriodOptions, + event_manager: &Arc, + ) -> Self { let provider_connections: DashMap, Arc>> = DashMap::new(); - let lineups = inputs - .iter() - .map(|i| Self::create_lineup(i, &provider_connections, event_manager)) - .collect(); + let lineups = inputs.iter().map(|i| Self::create_lineup(i, &provider_connections, event_manager)).collect(); Self { grace_period_millis: AtomicU64::new(grace_period_options.period_millis), grace_period_timeout_secs: AtomicU64::new(grace_period_options.timeout_secs), - snapshot: Arc::new(ArcSwap::from_pointee(LineupSnapshot { - inputs, - providers: lineups, - })), + snapshot: Arc::new(ArcSwap::from_pointee(LineupSnapshot { inputs, providers: lineups })), provider_connections, event_manager: Arc::clone(event_manager), } @@ -543,9 +550,10 @@ impl ProviderLineupManager { event_manager: &Arc, ) -> ProviderLineup { let event_manager = Arc::clone(event_manager); - let on_connection_change: ProviderConnectionChangeCallback = Arc::new(move |name: &Arc, connections: usize| { - event_manager.send_provider_event(name, connections); - }); + let on_connection_change: ProviderConnectionChangeCallback = + Arc::new(move |name: &Arc, connections: usize| { + event_manager.send_provider_event(name, connections); + }); if cfg_input.has_enabled_aliases() { ProviderLineup::Multi(MultiProviderLineup::new(cfg_input, provider_connections, &on_connection_change)) @@ -632,19 +640,15 @@ impl ProviderLineupManager { debug_if_enabled!("inputs {}", sanitize_sensitive_info(&display_vec(&new_inputs))); debug_if_enabled!("lineup {}", sanitize_sensitive_info(&display_vec(&new_lineups))); - self.snapshot.store(Arc::new(LineupSnapshot { - inputs: new_inputs, - providers: new_lineups, - })); + self.snapshot.store(Arc::new(LineupSnapshot { inputs: new_inputs, providers: new_lineups })); } pub async fn reconcile_connections(&self, mut counts: HashMap, usize>) { // 1. Synchronize known providers from actual counts. // We take a snapshot of the keys and locks to avoid holding the DashMap's internal // shard locks while awaiting the RwLock of each provider. Holding both can lead to deadlocks. - let snapshot: Vec<_> = self.provider_connections.iter() - .map(|e| (e.key().clone(), Arc::clone(e.value()))) - .collect(); + let snapshot: Vec<_> = + self.provider_connections.iter().map(|e| (e.key().clone(), Arc::clone(e.value()))).collect(); for (name, conn_lock) in snapshot { let count = counts.remove(&name).unwrap_or(0); @@ -660,7 +664,8 @@ impl ProviderLineupManager { // Same lock-ordering rule as Phase 1: clone the Arc> out of DashMap first, // then await on the provider RwLock without holding any DashMap shard lock. for (name, count) in counts { - let conn_lock = self.provider_connections + let conn_lock = self + .provider_connections .entry(name) .or_insert_with(|| Arc::new(RwLock::new(ProviderConfigConnection::default()))) .clone(); @@ -671,9 +676,8 @@ impl ProviderLineupManager { // 3. Broadcast status updates to the UI/Event system. // We drop the snapshot before GC to ensure Arc counts are accurate. { - let snapshot: Vec<_> = self.provider_connections.iter() - .map(|e| (e.key().clone(), Arc::clone(e.value()))) - .collect(); + let snapshot: Vec<_> = + self.provider_connections.iter().map(|e| (e.key().clone(), Arc::clone(e.value()))).collect(); for (name, conn_lock) in snapshot { let count = conn_lock.read().await.current_connections; @@ -690,14 +694,20 @@ impl ProviderLineupManager { let snapshot = self.snapshot.load(); for lineup in &snapshot.providers { match lineup { - ProviderLineup::Single(s) => { names.insert(s.provider.name.clone()); } + ProviderLineup::Single(s) => { + names.insert(s.provider.name.clone()); + } ProviderLineup::Multi(m) => { names.insert(m.name.clone()); for g in &m.providers { match g { - ProviderPriorityGroup::SingleProviderGroup(p) => { names.insert(p.name.clone()); } + ProviderPriorityGroup::SingleProviderGroup(p) => { + names.insert(p.name.clone()); + } ProviderPriorityGroup::MultiProviderGroup(_, group) => { - for p in group { names.insert(p.name.clone()); } + for p in group { + names.insert(p.name.clone()); + } } } } @@ -722,7 +732,6 @@ impl ProviderLineupManager { gen_provider_search!(get_provider_config_by_name, name, &Arc); - fn log_allocation(allocation: &ProviderAllocation) { match allocation { ProviderAllocation::Exhausted => {} @@ -731,8 +740,7 @@ impl ProviderLineupManager { "Using provider {} (pool user: {}, max_connections: {})", sanitize_sensitive_info(&cfg.name), cfg.get_user_info() - .map_or_else(|| "?".to_string(), |ui| sanitize_sensitive_info(&ui.username).to_string()) - , + .map_or_else(|| "?".to_string(), |ui| sanitize_sensitive_info(&ui.username).to_string()), cfg.max_connections() ); } @@ -762,9 +770,7 @@ impl ProviderLineupManager { let allocation = match Self::get_provider_config_by_name(provider_name, &snapshot.providers) { None => ProviderAllocation::Exhausted, Some((_lineup, config)) => { - config - .try_allocate(with_grace, self.grace_period_timeout_secs.load(Ordering::Acquire)) - .await + config.try_allocate(with_grace, self.grace_period_timeout_secs.load(Ordering::Acquire)).await } }; Self::log_allocation(&allocation); @@ -791,9 +797,7 @@ impl ProviderLineupManager { let allocation = match lineup_opt { None => ProviderAllocation::Exhausted, // No Name matched, we don't have this provider Some((lineup, _config)) => { - lineup - .acquire(with_grace, self.grace_period_timeout_secs.load(Ordering::Acquire)) - .await + lineup.acquire(with_grace, self.grace_period_timeout_secs.load(Ordering::Acquire)).await } }; if matches!(allocation, ProviderAllocation::Exhausted) { @@ -926,29 +930,26 @@ impl ProviderLineupManager { pub fn is_provider_for_input(&self, provider_name: &str, input_name: &str) -> bool { let snapshot = self.snapshot.load(); - if let Some((lineup, _)) = - Self::get_provider_config_by_name(&provider_name.into(), &snapshot.providers) - { - match lineup { - ProviderLineup::Single(_) => return input_name == provider_name, - ProviderLineup::Multi(m) => return m.name.as_ref() == input_name, - } + if let Some((lineup, _)) = Self::get_provider_config_by_name(&provider_name.into(), &snapshot.providers) { + match lineup { + ProviderLineup::Single(_) => return input_name == provider_name, + ProviderLineup::Multi(m) => return m.name.as_ref() == input_name, + } } false } } - #[cfg(test)] mod tests { use super::*; - use crate::model::ConfigInputAlias; - use crate::Arc; - use shared::model::{InputFetchMethod, InputType}; - use std::sync::atomic::AtomicU16; - use std::thread; - use shared::concat_string; - use shared::utils::Internable; + use crate::{model::ConfigInputAlias, Arc}; + use shared::{ + concat_string, + model::{InputFetchMethod, InputType}, + utils::Internable, + }; + use std::{sync::atomic::AtomicU16, thread}; macro_rules! should_available { ($lineup:expr, $provider_id:expr, $grace_period_timeout_secs: expr) => { @@ -956,7 +957,9 @@ mod tests { match $lineup.acquire(true, $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), + ProviderAllocation::GracePeriod(provider) => { + assert!(false, "Should available and not grace period: {}", provider.id) + } } }; } @@ -965,7 +968,9 @@ mod tests { thread::sleep(std::time::Duration::from_millis(200)); match $lineup.acquire(true, $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::Available(provider) => { + assert!(false, "Should grace period and not available: {}", provider.id) + } ProviderAllocation::GracePeriod(provider) => assert_eq!(provider.id, $provider_id), } }; @@ -975,9 +980,13 @@ mod tests { ($lineup:expr, $grace_period_timeout_secs: expr) => { thread::sleep(std::time::Duration::from_millis(200)); match $lineup.acquire(true, $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), + 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) + } } }; } @@ -1026,7 +1035,6 @@ mod tests { fn dummy_callback(_: &Arc, _: usize) {} - // Test acquiring with an alias #[test] fn test_provider_with_alias() { @@ -1104,7 +1112,6 @@ mod tests { }); } - // Test acquiring when all aliases are exhausted #[test] fn test_provider_with_exhausted_aliases() { @@ -1159,7 +1166,6 @@ mod tests { }); } - // Test releasing a connection #[test] fn test_release_connection() { diff --git a/backend/src/api/model/request.rs b/backend/src/api/model/request.rs index 7a9a7e866..d9eb7ee52 100644 --- a/backend/src/api/model/request.rs +++ b/backend/src/api/model/request.rs @@ -38,4 +38,4 @@ impl UserApiRequest { self.limit.parse::().unwrap_or(0) } } -} \ No newline at end of file +} diff --git a/backend/src/api/model/stream.rs b/backend/src/api/model/stream.rs index 693b13903..40df403b7 100644 --- a/backend/src/api/model/stream.rs +++ b/backend/src/api/model/stream.rs @@ -1,12 +1,13 @@ -use std::collections::HashMap; -use std::sync::Arc; -use crate::api::model::{CustomVideoStreamType, ProviderHandle, StreamError}; +use crate::{ + api::model::{CustomVideoStreamType, ProviderHandle, StreamError}, + model::GracePeriodOptions, + tools::atomic_once_flag::AtomicOnceFlag, +}; use axum::http::StatusCode; use bytes::Bytes; use futures::stream::BoxStream; +use std::{collections::HashMap, sync::Arc}; use url::Url; -use crate::model::GracePeriodOptions; -use crate::tools::atomic_once_flag::AtomicOnceFlag; pub type BoxedProviderStream = BoxStream<'static, Result>; pub type ProviderStreamHeader = Vec<(String, String)>; @@ -50,14 +51,10 @@ impl StreamDetails { } } #[inline] - pub fn has_stream(&self) -> bool { - self.stream.is_some() - } + pub fn has_stream(&self) -> bool { self.stream.is_some() } #[inline] - pub fn has_grace_period(&self) -> bool { - self.grace_period.period_millis > 0 - } + pub fn has_grace_period(&self) -> bool { self.grace_period.period_millis > 0 } } pub struct StreamingStrategy { diff --git a/backend/src/api/model/stream_error.rs b/backend/src/api/model/stream_error.rs index 5a04d00ba..feb8d34b8 100644 --- a/backend/src/api/model/stream_error.rs +++ b/backend/src/api/model/stream_error.rs @@ -7,16 +7,14 @@ pub enum StreamError { // ReceiverClosed, ReceiverError(BroadcastStreamRecvError), LockError(String), - Stream(String) + Stream(String), } impl StreamError { -// pub(crate) fn std_io(msg: String) -> Self { -// StreamError::StdIo(std::io::Error::new(std::io::ErrorKind::Other, msg)) -// } - pub fn reqwest(err: &reqwest::Error) -> Self { - Self::Reqwest(err.to_string()) - } + // pub(crate) fn std_io(msg: String) -> Self { + // StreamError::StdIo(std::io::Error::new(std::io::ErrorKind::Other, msg)) + // } + pub fn reqwest(err: &reqwest::Error) -> Self { Self::Reqwest(err.to_string()) } } impl std::error::Error for StreamError {} @@ -27,8 +25,8 @@ impl std::fmt::Display for StreamError { StreamError::Reqwest(e) => write!(f, "Reqwest error: {e}"), StreamError::StdIo(e) => write!(f, "IO error: {e}"), // StreamError::ReceiverClosed => write!(f, "Receiver closed"), - StreamError::ReceiverError(e) => write!(f, "Receiver error {e}"), - StreamError::Stream(e) | StreamError::LockError(e) => write!(f, "{e}") + StreamError::ReceiverError(e) => write!(f, "Receiver error {e}"), + StreamError::Stream(e) | StreamError::LockError(e) => write!(f, "{e}"), } } -} \ No newline at end of file +} diff --git a/backend/src/api/model/streams/active_client_stream.rs b/backend/src/api/model/streams/active_client_stream.rs index 208ec615d..aae9bee92 100644 --- a/backend/src/api/model/streams/active_client_stream.rs +++ b/backend/src/api/model/streams/active_client_stream.rs @@ -1,26 +1,32 @@ -use crate::api::model::BoxedProviderStream; -use crate::api::model::StreamError; -use crate::api::model::TimedClientStream; -use crate::api::model::TransportStreamBuffer; -use crate::api::model::{AppState, ConnectionManager, CustomVideoStreamType, ProviderHandle, StreamDetails}; -use crate::api::panel_api::{can_provision_on_exhausted, find_input_by_provider_name, run_panel_api_provisioning_probe}; -use crate::auth::Fingerprint; -use crate::model::{ConfigInput, ProxyUserCredentials}; -use crate::tools::atomic_once_flag::AtomicOnceFlag; -use crate::utils::debug_if_enabled; -use axum::http::header::USER_AGENT; -use axum::http::HeaderMap; +use crate::{ + api::{ + model::{ + AppState, BoxedProviderStream, ConnectionManager, CustomVideoStreamType, ProviderHandle, StreamDetails, + StreamError, TimedClientStream, TransportStreamBuffer, + }, + panel_api::{can_provision_on_exhausted, find_input_by_provider_name, run_panel_api_provisioning_probe}, + }, + auth::Fingerprint, + model::{ConfigInput, ProxyUserCredentials}, + tools::atomic_once_flag::AtomicOnceFlag, + utils::debug_if_enabled, +}; +use axum::http::{header::USER_AGENT, HeaderMap}; use bytes::Bytes; -use futures::task::AtomicWaker; -use futures::Stream; -use futures::StreamExt; +use futures::{task::AtomicWaker, Stream, StreamExt}; use log::{debug, error, info}; -use shared::model::{StreamChannel, UserConnectionPermission, VirtualId}; -use shared::utils::sanitize_sensitive_info; -use std::pin::Pin; -use std::sync::atomic::{AtomicU8, Ordering}; -use std::sync::Arc; -use std::task::{Context, Poll}; +use shared::{ + model::{StreamChannel, UserConnectionPermission, VirtualId}, + utils::sanitize_sensitive_info, +}; +use std::{ + pin::Pin, + sync::{ + atomic::{AtomicU8, Ordering}, + Arc, + }, + task::{Context, Poll}, +}; const INNER_STREAM: u8 = 0_u8; const USER_EXHAUSTED_STREAM: u8 = 1_u8; @@ -253,15 +259,9 @@ pub(crate) async fn create_active_client_stream( } let grant_user_grace_period = connection_permission == UserConnectionPermission::GracePeriod; let username = user.username.as_str(); - let provider_name = stream_details - .provider_name - .as_ref() - .map_or_else(String::new, ToString::to_string); + let provider_name = stream_details.provider_name.as_ref().map_or_else(String::new, ToString::to_string); - let user_agent = req_headers - .get(USER_AGENT) - .map(|h| String::from_utf8_lossy(h.as_bytes())) - .unwrap_or_default(); + let user_agent = req_headers.get(USER_AGENT).map(|h| String::from_utf8_lossy(h.as_bytes())).unwrap_or_default(); let virtual_id = stream_channel.virtual_id; app_state @@ -277,36 +277,17 @@ pub(crate) async fn create_active_client_stream( ) .await; if let Some((_, _, _m_, Some(cvt))) = stream_details.stream_info.as_ref() { - app_state - .connection_manager - .update_stream_detail(&fingerprint.addr, *cvt) - .await; + app_state.connection_manager.update_stream_detail(&fingerprint.addr, *cvt).await; } let provisioning_info = resolve_grace_period_provisioning(app_state, &stream_details); let has_provisioning = provisioning_info.is_some(); let hold_stream = stream_details.grace_period.hold_stream; - let (grace_stop_flag, waker) = if grant_user_grace_period - || (stream_details.has_grace_period() && stream_details.provider_name.is_some()) - { - let waker = Arc::new(AtomicWaker::new()); - let flag = stream_grace_period( - app_state, - &stream_details, - grant_user_grace_period, - user, - fingerprint, - virtual_id, - provisioning_info, - Some(Arc::clone(&waker)), - hold_stream, - ); - let maybe_waker = flag.as_ref().map(|_| waker); - (flag, maybe_waker) - } else { - ( - stream_grace_period( + let (grace_stop_flag, waker) = + if grant_user_grace_period || (stream_details.has_grace_period() && stream_details.provider_name.is_some()) { + let waker = Arc::new(AtomicWaker::new()); + let flag = stream_grace_period( app_state, &stream_details, grant_user_grace_period, @@ -314,12 +295,27 @@ pub(crate) async fn create_active_client_stream( fingerprint, virtual_id, provisioning_info, - None, + Some(Arc::clone(&waker)), hold_stream, - ), - None, - ) - }; + ); + let maybe_waker = flag.as_ref().map(|_| waker); + (flag, maybe_waker) + } else { + ( + stream_grace_period( + app_state, + &stream_details, + grant_user_grace_period, + user, + fingerprint, + virtual_id, + provisioning_info, + None, + hold_stream, + ), + None, + ) + }; let cfg = &app_state.app_config; let custom_response = cfg.custom_stream_response.load(); @@ -335,10 +331,7 @@ pub(crate) async fn create_active_client_stream( let stream = match stream_details.stream.take() { None => { let provider_handle = stream_details.provider_handle.take(); - app_state - .connection_manager - .release_provider_handle(provider_handle) - .await; + app_state.connection_manager.release_provider_handle(provider_handle).await; futures::stream::empty::>().boxed() } Some(stream) => { @@ -346,11 +339,9 @@ pub(crate) async fn create_active_client_stream( match config.sleep_timer_mins { None => stream, Some(mins) => { - let secs = - u32::try_from((u64::from(mins) * 60).min(u64::from(u32::MAX))).unwrap_or(0); + let secs = u32::try_from((u64::from(mins) * 60).min(u64::from(u32::MAX))).unwrap_or(0); if secs > 0 { - TimedClientStream::new(app_state, stream, secs, fingerprint.addr, virtual_id) - .boxed() + TimedClientStream::new(app_state, stream, secs, fingerprint.addr, virtual_id).boxed() } else { stream } @@ -385,8 +376,7 @@ fn resolve_grace_period_provisioning( return None; } let provider_name = stream_details.provider_name.as_deref(); - let input = - provider_name.and_then(|name| find_input_by_provider_name(app_state.as_ref(), name))?; + let input = provider_name.and_then(|name| find_input_by_provider_name(app_state.as_ref(), name))?; if !can_provision_on_exhausted(app_state, &input) { return None; } @@ -431,7 +421,7 @@ fn stream_grace_period( debug!("hold stream {hold_stream}"); if provider_grace_check.is_some() || user_grace_check.is_some() { - let stream_strategy_flag = Arc::new(AtomicU8::new(if hold_stream {GRACE_PENDING} else {INNER_STREAM})); + let stream_strategy_flag = Arc::new(AtomicU8::new(if hold_stream { GRACE_PENDING } else { INNER_STREAM })); let stream_strategy_flag_copy = Arc::clone(&stream_strategy_flag); let grace_period_millis = stream_details.grace_period.period_millis; @@ -450,10 +440,7 @@ fn stream_grace_period( if active_connections > max_connections { stream_strategy_flag_copy.store(USER_EXHAUSTED_STREAM, Ordering::Release); connection_manager - .update_stream_detail( - &fingerprint.addr, - CustomVideoStreamType::UserConnectionsExhausted, - ) + .update_stream_detail(&fingerprint.addr, CustomVideoStreamType::UserConnectionsExhausted) .await; // Release the shared stream subscription to stop the subscriber loop connection_manager.shared_stream_manager.release_connection(&fingerprint.addr, true).await; @@ -468,10 +455,7 @@ fn stream_grace_period( if let Some(provisioning_info) = provisioning_info { stream_strategy_flag_copy.store(PROVISIONING_STREAM, Ordering::Release); connection_manager - .update_stream_detail( - &fingerprint.addr, - CustomVideoStreamType::Provisioning, - ) + .update_stream_detail(&fingerprint.addr, CustomVideoStreamType::Provisioning) .await; debug_if_enabled!( "Provider grace period exhausted; provisioning for active clients: {provider_name}" @@ -481,20 +465,15 @@ fn stream_grace_period( let stop_signal = provisioning_info.stop_signal; let addr = fingerprint.addr; tokio::spawn(async move { - if let Err(err) = run_panel_api_provisioning_probe( - app_state, - input, - stop_signal, - addr, - virtual_id, - ) - .await { + if let Err(err) = + run_panel_api_provisioning_probe(app_state, input, stop_signal, addr, virtual_id) + .await + { error!("Error running Probe: {err:?}"); } }); } else { - stream_strategy_flag_copy - .store(PROVIDER_EXHAUSTED_STREAM, Ordering::Release); + stream_strategy_flag_copy.store(PROVIDER_EXHAUSTED_STREAM, Ordering::Release); connection_manager .update_stream_detail( &fingerprint.addr, diff --git a/backend/src/api/model/streams/buffered_stream.rs b/backend/src/api/model/streams/buffered_stream.rs index f0055b973..e82731120 100644 --- a/backend/src/api/model/streams/buffered_stream.rs +++ b/backend/src/api/model/streams/buffered_stream.rs @@ -1,16 +1,20 @@ -use crate::api::model::{BoxedProviderStream, StreamError, STREAM_IDLE_TIMEOUT}; -use crate::tools::atomic_once_flag::AtomicOnceFlag; -use futures::{stream::Stream, task::{Context, Poll}, StreamExt}; -use log::{debug}; -use std::{ - cmp::max, - pin::Pin, - sync::Arc, +use crate::{ + api::model::{BoxedProviderStream, StreamError, STREAM_IDLE_TIMEOUT}, + tools::atomic_once_flag::AtomicOnceFlag, +}; +use futures::{ + stream::Stream, + task::{Context, Poll}, + StreamExt, +}; +use log::debug; +use std::{cmp::max, pin::Pin, sync::Arc}; +use tokio::{ + select, + sync::mpsc::{channel, Sender}, + time::{sleep, Duration, Instant}, }; -use tokio::select; -use tokio::sync::mpsc::{channel, Sender}; use tokio_stream::wrappers::ReceiverStream; -use tokio::time::{sleep, Duration, Instant}; pub const CHANNEL_SIZE: usize = 1024; @@ -20,14 +24,16 @@ pub(in crate::api::model) struct BufferedStream { } impl BufferedStream { - pub fn new(stream: BoxedProviderStream, buffer_size: usize, client_close_signal: Arc, _url: &str) -> Self { + pub fn new( + stream: BoxedProviderStream, + buffer_size: usize, + client_close_signal: Arc, + _url: &str, + ) -> Self { // TODO make channel_size based on bytes not entries let (tx, rx) = channel(max(buffer_size, CHANNEL_SIZE)); tokio::spawn(Self::buffer_stream(tx, stream, Arc::clone(&client_close_signal))); - Self { - stream: ReceiverStream::new(rx), - close_signal: client_close_signal, - } + Self { stream: ReceiverStream::new(rx), close_signal: client_close_signal } } async fn buffer_stream( diff --git a/backend/src/api/model/streams/client_stream.rs b/backend/src/api/model/streams/client_stream.rs index 64f0c5fd7..32223ba1a 100644 --- a/backend/src/api/model/streams/client_stream.rs +++ b/backend/src/api/model/streams/client_stream.rs @@ -1,15 +1,20 @@ -use crate::api::model::BoxedProviderStream; -use crate::api::model::StreamError; -use crate::tools::atomic_once_flag::AtomicOnceFlag; -use crate::utils::trace_if_enabled; +use crate::{ + api::model::{BoxedProviderStream, StreamError}, + tools::atomic_once_flag::AtomicOnceFlag, + utils::trace_if_enabled, +}; use bytes::Bytes; use futures::Stream; use log::trace; use shared::utils::sanitize_sensitive_info; -use std::pin::Pin; -use std::sync::atomic::{AtomicUsize, Ordering}; -use std::sync::Arc; -use std::task::Poll; +use std::{ + pin::Pin, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }, + task::Poll, +}; /// This stream counts the send bytes for reconnecting to the actual position and /// sets the `close_signal` if the client drops the connection. @@ -22,17 +27,19 @@ pub(in crate::api::model) struct ClientStream { } impl ClientStream { - pub(crate) fn new(inner: BoxedProviderStream, close_signal: Arc, total_bytes: Arc>, url: &str) -> Self { + pub(crate) fn new( + inner: BoxedProviderStream, + close_signal: Arc, + total_bytes: Arc>, + url: &str, + ) -> Self { Self { inner, close_signal, total_bytes, url: url.to_string() } } } impl Stream for ClientStream { type Item = Result; - fn poll_next( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> Poll> { + fn poll_next(mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>) -> Poll> { if self.close_signal.is_active() { match Pin::as_mut(&mut self.inner).poll_next(cx) { Poll::Ready(Some(Ok(bytes))) => { @@ -41,7 +48,7 @@ impl Stream for ClientStream { // Empty payload signals upstream closure; notify and let consumer see final chunk self.close_signal.notify(); } else if let Some(counter) = self.total_bytes.as_ref() { - counter.fetch_add(bytes.len(), Ordering::AcqRel); + counter.fetch_add(bytes.len(), Ordering::AcqRel); } Poll::Ready(Some(Ok(bytes))) @@ -63,10 +70,9 @@ impl Stream for ClientStream { } } - impl Drop for ClientStream { fn drop(&mut self) { trace_if_enabled!("Client disconnected {}", sanitize_sensitive_info(&self.url)); self.close_signal.notify(); } -} \ No newline at end of file +} diff --git a/backend/src/api/model/streams/custom_video_stream.rs b/backend/src/api/model/streams/custom_video_stream.rs index 609d272f5..c84090b90 100644 --- a/backend/src/api/model/streams/custom_video_stream.rs +++ b/backend/src/api/model/streams/custom_video_stream.rs @@ -1,21 +1,17 @@ -use crate::api::model::StreamError; -use crate::api::model::TransportStreamBuffer; +use crate::api::model::{StreamError, TransportStreamBuffer}; use bytes::Bytes; use futures::Stream; -use std::pin::Pin; -use std::task::{Context, Poll}; - +use std::{ + pin::Pin, + task::{Context, Poll}, +}; pub struct CustomVideoStream { buffer: TransportStreamBuffer, } impl CustomVideoStream { - pub fn new(buffer: TransportStreamBuffer) -> Self { - Self { - buffer - } - } + pub fn new(buffer: TransportStreamBuffer) -> Self { Self { buffer } } } impl Stream for CustomVideoStream { diff --git a/backend/src/api/model/streams/mod.rs b/backend/src/api/model/streams/mod.rs index 6ace8d62d..0a3a62028 100644 --- a/backend/src/api/model/streams/mod.rs +++ b/backend/src/api/model/streams/mod.rs @@ -1,26 +1,22 @@ -mod timed_client_stream; mod buffered_stream; mod client_stream; mod custom_video_stream; mod provisioning_stream; +mod timed_client_stream; mod transport_stream_buffer; // mod chunked_buffer; +mod active_client_stream; +pub mod persist_pipe_stream; mod provider_stream; mod provider_stream_factory; mod shared_stream_manager; -mod active_client_stream; mod throttled_stream; -pub mod persist_pipe_stream; -pub(in crate) use self::transport_stream_buffer::*; -pub(in crate::api) use self::provider_stream::*; -pub(in crate::api) use self::provider_stream_factory::*; -pub(in crate::api) use self::shared_stream_manager::*; -pub(in crate::api) use self::active_client_stream::*; -pub(in crate::api) use self::throttled_stream::*; -pub(in crate::api) use self::timed_client_stream::*; -pub(in crate::api) use self::custom_video_stream::*; -pub(in crate::api) use self::provisioning_stream::*; pub use self::persist_pipe_stream::*; +pub(crate) use self::transport_stream_buffer::*; +pub(in crate::api) use self::{ + active_client_stream::*, custom_video_stream::*, provider_stream::*, provider_stream_factory::*, + provisioning_stream::*, shared_stream_manager::*, throttled_stream::*, timed_client_stream::*, +}; -pub const STREAM_IDLE_TIMEOUT: u64 = 60; \ No newline at end of file +pub const STREAM_IDLE_TIMEOUT: u64 = 60; diff --git a/backend/src/api/model/streams/persist_pipe_stream.rs b/backend/src/api/model/streams/persist_pipe_stream.rs index d9d352f15..44e0a48eb 100644 --- a/backend/src/api/model/streams/persist_pipe_stream.rs +++ b/backend/src/api/model/streams/persist_pipe_stream.rs @@ -1,16 +1,16 @@ -use crate::api::model::{StreamError, STREAM_IDLE_TIMEOUT}; -use crate::utils::request::DynReader; -use crate::utils::{async_file_writer, IO_BUFFER_SIZE}; +use crate::{ + api::model::{StreamError, STREAM_IDLE_TIMEOUT}, + utils::{async_file_writer, debug_if_enabled, request::DynReader, IO_BUFFER_SIZE}, +}; use bytes::Bytes; use log::{debug, error}; -use std::path::Path; -use std::sync::Arc; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use tokio::select; -use tokio_stream::wrappers::ReceiverStream; -use tokio_stream::StreamExt; -use tokio::time::{sleep, Duration, Instant}; -use crate::utils::debug_if_enabled; +use std::{path::Path, sync::Arc}; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + select, + time::{sleep, Duration, Instant}, +}; +use tokio_stream::{wrappers::ReceiverStream, StreamExt}; pub fn tee_stream( mut stream: S, @@ -19,7 +19,7 @@ pub fn tee_stream( callback: Arc, ) -> ReceiverStream> where - S: tokio_stream::Stream> + Send + Unpin + 'static, + S: tokio_stream::Stream> + Send + Unpin + 'static, W: tokio::io::AsyncWrite + Send + Unpin + 'static, { let (tx, rx) = tokio::sync::mpsc::channel::>(32); @@ -150,4 +150,4 @@ pub async fn tee_dyn_reader( }); Box::pin(rx) as DynReader -} \ No newline at end of file +} diff --git a/backend/src/api/model/streams/provider_stream.rs b/backend/src/api/model/streams/provider_stream.rs index 1304fabb4..a386fa08f 100644 --- a/backend/src/api/model/streams/provider_stream.rs +++ b/backend/src/api/model/streams/provider_stream.rs @@ -1,19 +1,20 @@ -use crate::api::api_utils::{HeaderFilter}; -use crate::api::model::{AppState, CustomVideoStream, ProvisioningStream, ThrottledStream}; -use crate::model::{AppConfig}; -use shared::model::PlaylistItemType; -use log::{trace}; -use reqwest::StatusCode; +use crate::{ + api::{ + api_utils::{try_unwrap_body, HeaderFilter}, + model::{ + stream::ProviderStreamResponse, AppState, CustomVideoStream, ProvisioningStream, ThrottledStream, + TransportStreamBuffer, + }, + }, + model::AppConfig, + tools::atomic_once_flag::AtomicOnceFlag, +}; use axum::response::IntoResponse; -use crate::api::model::stream::ProviderStreamResponse; -use crate::api::model::TransportStreamBuffer; -use crate::api::api_utils::try_unwrap_body; -use crate::tools::atomic_once_flag::AtomicOnceFlag; -use std::str::FromStr; -use std::fmt; -use std::net::SocketAddr; -use std::sync::Arc; -use serde::{Serialize, Deserialize, Serializer, Deserializer}; +use log::trace; +use reqwest::StatusCode; +use serde::{Deserialize, Deserializer, Serialize, Serializer}; +use shared::model::PlaylistItemType; +use std::{fmt, net::SocketAddr, str::FromStr, sync::Arc}; #[derive(Debug, Copy, Clone)] pub enum CustomVideoStreamType { @@ -70,48 +71,86 @@ impl<'de> Deserialize<'de> for CustomVideoStreamType { } } -fn create_video_stream(stream_type: CustomVideoStreamType, video_buffer: Option<&TransportStreamBuffer>, headers: &[(String, String)], log_message: &str) -> ProviderStreamResponse { +fn create_video_stream( + stream_type: CustomVideoStreamType, + video_buffer: Option<&TransportStreamBuffer>, + headers: &[(String, String)], + log_message: &str, +) -> ProviderStreamResponse { if let Some(video) = video_buffer { trace!("{log_message}"); - let mut response_headers: Vec<(String, String)> = headers.iter() + let mut response_headers: Vec<(String, String)> = headers + .iter() .filter(|(key, _)| !(key.eq("content-type") || key.eq("content-length") || key.contains("range"))) - .map(|(key, value)| (key.clone(), value.clone())).collect(); + .map(|(key, value)| (key.clone(), value.clone())) + .collect(); response_headers.push(("content-type".to_string(), "video/mp2t".to_string())); - (Some(Box::pin(ThrottledStream::new(CustomVideoStream::new(video.clone()), 8000))), Some((response_headers, StatusCode::OK, None, Some(stream_type)))) + ( + Some(Box::pin(ThrottledStream::new(CustomVideoStream::new(video.clone()), 8000))), + Some((response_headers, StatusCode::OK, None, Some(stream_type))), + ) } else { (None, None) } } -pub fn create_channel_unavailable_stream(cfg: &AppConfig, headers: &[(String, String)], status: StatusCode) -> ProviderStreamResponse { +pub fn create_channel_unavailable_stream( + cfg: &AppConfig, + headers: &[(String, String)], + status: StatusCode, +) -> ProviderStreamResponse { let custom_stream_response = cfg.custom_stream_response.load(); let video = custom_stream_response.as_ref().and_then(|c| c.channel_unavailable.as_ref()); - create_video_stream(CustomVideoStreamType::ChannelUnavailable, video, headers, &format!("Streaming response channel unavailable for status {status}")) + create_video_stream( + CustomVideoStreamType::ChannelUnavailable, + video, + headers, + &format!("Streaming response channel unavailable for status {status}"), + ) } -pub fn create_user_connections_exhausted_stream(cfg: &AppConfig, headers: &[(String, String)]) -> ProviderStreamResponse { +pub fn create_user_connections_exhausted_stream( + cfg: &AppConfig, + headers: &[(String, String)], +) -> ProviderStreamResponse { let custom_stream_response = cfg.custom_stream_response.load(); let video = custom_stream_response.as_ref().and_then(|c| c.user_connections_exhausted.as_ref()); - create_video_stream(CustomVideoStreamType::UserConnectionsExhausted, video, headers, "Streaming response user connections exhausted") + create_video_stream( + CustomVideoStreamType::UserConnectionsExhausted, + video, + headers, + "Streaming response user connections exhausted", + ) } -pub fn create_provider_connections_exhausted_stream(cfg: &AppConfig, headers: &[(String, String)]) -> ProviderStreamResponse { +pub fn create_provider_connections_exhausted_stream( + cfg: &AppConfig, + headers: &[(String, String)], +) -> ProviderStreamResponse { let custom_stream_response = cfg.custom_stream_response.load(); let video = custom_stream_response.as_ref().and_then(|c| c.provider_connections_exhausted.as_ref()); - create_video_stream(CustomVideoStreamType::ProviderConnectionsExhausted, video, headers, "Streaming response provider connections exhausted") + create_video_stream( + CustomVideoStreamType::ProviderConnectionsExhausted, + video, + headers, + "Streaming response provider connections exhausted", + ) } pub fn create_user_account_expired_stream(cfg: &AppConfig, headers: &[(String, String)]) -> ProviderStreamResponse { let custom_stream_response = cfg.custom_stream_response.load(); let video = custom_stream_response.as_ref().and_then(|c| c.user_account_expired.as_ref()); - create_video_stream(CustomVideoStreamType::UserAccountExpired, video, headers, "Streaming response user account expired") + create_video_stream( + CustomVideoStreamType::UserAccountExpired, + video, + headers, + "Streaming response user account expired", + ) } pub fn create_panel_api_provisioning_stream(cfg: &AppConfig, headers: &[(String, String)]) -> ProviderStreamResponse { let custom_stream_response = cfg.custom_stream_response.load(); - let video = custom_stream_response - .as_ref() - .and_then(|c| c.panel_api_provisioning.as_ref()); + let video = custom_stream_response.as_ref().and_then(|c| c.panel_api_provisioning.as_ref()); create_video_stream( CustomVideoStreamType::Provisioning, video, @@ -126,9 +165,7 @@ pub fn create_panel_api_provisioning_stream_with_stop( stop_signal: Arc, ) -> ProviderStreamResponse { let custom_stream_response = cfg.custom_stream_response.load(); - let video = custom_stream_response - .as_ref() - .and_then(|c| c.panel_api_provisioning.as_ref()); + let video = custom_stream_response.as_ref().and_then(|c| c.panel_api_provisioning.as_ref()); if let Some(video) = video { trace!("Streaming response panel api provisioning"); let mut response_headers: Vec<(String, String)> = headers @@ -140,31 +177,33 @@ pub fn create_panel_api_provisioning_stream_with_stop( let stream = ProvisioningStream::new(video.clone(), stop_signal); ( Some(Box::pin(ThrottledStream::new(stream, 8000))), - Some(( - response_headers, - StatusCode::OK, - None, - Some(CustomVideoStreamType::Provisioning), - )), + Some((response_headers, StatusCode::OK, None, Some(CustomVideoStreamType::Provisioning))), ) } else { (None, None) } } -pub async fn create_custom_video_stream_response(app_state: &Arc, addr: &SocketAddr, video_response: CustomVideoStreamType) -> impl axum::response::IntoResponse + Send { +pub async fn create_custom_video_stream_response( + app_state: &Arc, + addr: &SocketAddr, + video_response: CustomVideoStreamType, +) -> impl axum::response::IntoResponse + Send { let config = &app_state.app_config; if let (Some(stream), Some((headers, status_code, _, _))) = match video_response { - CustomVideoStreamType::ChannelUnavailable => create_channel_unavailable_stream(config, &[], StatusCode::BAD_REQUEST), + CustomVideoStreamType::ChannelUnavailable => { + create_channel_unavailable_stream(config, &[], StatusCode::BAD_REQUEST) + } CustomVideoStreamType::UserConnectionsExhausted => create_user_connections_exhausted_stream(config, &[]), - CustomVideoStreamType::ProviderConnectionsExhausted => create_provider_connections_exhausted_stream(config, &[]), + CustomVideoStreamType::ProviderConnectionsExhausted => { + create_provider_connections_exhausted_stream(config, &[]) + } CustomVideoStreamType::UserAccountExpired => create_user_account_expired_stream(config, &[]), CustomVideoStreamType::Provisioning => create_panel_api_provisioning_stream(config, &[]), } { app_state.connection_manager.update_stream_detail(addr, video_response).await; app_state.connection_manager.release_provider_connection(addr).await; - let mut builder = axum::response::Response::builder() - .status(status_code); + let mut builder = axum::response::Response::builder().status(status_code); for (key, value) in headers { builder = builder.header(key, value); } diff --git a/backend/src/api/model/streams/provider_stream_factory.rs b/backend/src/api/model/streams/provider_stream_factory.rs index 64e1201ba..10904db78 100644 --- a/backend/src/api/model/streams/provider_stream_factory.rs +++ b/backend/src/api/model/streams/provider_stream_factory.rs @@ -1,28 +1,43 @@ -use crate::api::api_utils::{get_headers_from_request, StreamOptions}; -use crate::api::model::{get_response_headers, AppState, CustomVideoStreamType}; -use crate::api::model::StreamError; -use crate::api::model::{create_channel_unavailable_stream, get_header_filter_for_item_type}; -use crate::api::model::{BoxedProviderStream, ProviderStreamFactoryResponse}; -use crate::model::{ConfigProvider, ReverseProxyDisabledHeaderConfig}; -use crate::tools::atomic_once_flag::AtomicOnceFlag; -use crate::utils::debug_if_enabled; -use crate::utils::request::{classify_content_type, get_request_headers, MimeCategory, send_with_retry_and_provider}; -use futures::stream::{self}; -use futures::{StreamExt, TryStreamExt}; +use crate::{ + api::{ + api_utils::{get_headers_from_request, StreamOptions}, + model::{ + create_channel_unavailable_stream, get_header_filter_for_item_type, get_response_headers, + streams::{buffered_stream::BufferedStream, client_stream::ClientStream}, + AppState, BoxedProviderStream, CustomVideoStreamType, ProviderStreamFactoryResponse, StreamError, + }, + }, + model::{ConfigProvider, ReverseProxyDisabledHeaderConfig}, + tools::atomic_once_flag::AtomicOnceFlag, + utils::{ + debug_if_enabled, + request::{classify_content_type, get_request_headers, send_with_retry_and_provider, MimeCategory}, + }, +}; +use futures::{ + stream::{self}, + StreamExt, TryStreamExt, +}; use log::{debug, log_enabled, warn}; -use reqwest::header::{HeaderMap, RANGE}; -use reqwest::StatusCode; -use shared::create_bitset; -use shared::model::{PlaylistItemType, DEFAULT_USER_AGENT}; -use shared::utils::{filter_request_header, sanitize_sensitive_info}; -use std::collections::HashMap; -use std::net::SocketAddr; -use std::sync::atomic::{AtomicUsize, Ordering}; -use std::sync::Arc; -use std::time::{Duration, Instant}; +use reqwest::{ + header::{HeaderMap, RANGE}, + StatusCode, +}; +use shared::{ + create_bitset, + model::{PlaylistItemType, DEFAULT_USER_AGENT}, + utils::{filter_request_header, sanitize_sensitive_info}, +}; +use std::{ + collections::HashMap, + net::SocketAddr, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }, + time::{Duration, Instant}, +}; use url::Url; -use crate::api::model::streams::buffered_stream::BufferedStream; -use crate::api::model::streams::client_stream::ClientStream; const RETRY_SECONDS: u64 = 5; const ERR_MAX_RETRY_COUNT: u32 = 5; @@ -64,23 +79,14 @@ impl ProviderStreamFactoryOptions { disabled_headers: Option<&ReverseProxyDisabledHeaderConfig>, default_user_agent: Option<&str>, ) -> Self { - let buffer_size = if stream_options.buffer_enabled { - stream_options.buffer_size - } else { - 0 - }; + let buffer_size = if stream_options.buffer_enabled { stream_options.buffer_size } else { 0 }; let filter_header = get_header_filter_for_item_type(item_type); let mut req_headers = get_headers_from_request(req_headers, &filter_header); let requested_range = get_request_range_start_bytes(&req_headers); req_headers.remove("range"); // We merge configured input headers with the headers from the request. - let headers = get_request_headers( - input_headers, - Some(&req_headers), - disabled_headers, - default_user_agent, - ); + let headers = get_request_headers(input_headers, Some(&req_headers), disabled_headers, default_user_agent); let default_user_agent = default_user_agent .and_then(|ua| { @@ -126,71 +132,43 @@ impl ProviderStreamFactoryOptions { } } - pub fn set_provider(&mut self, provider: Option>) { - self.provider = provider; - } + pub fn set_provider(&mut self, provider: Option>) { self.provider = provider; } - pub fn get_provider(&self) -> Option<&Arc> { - self.provider.as_ref() - } + pub fn get_provider(&self) -> Option<&Arc> { self.provider.as_ref() } #[inline] - fn is_piped(&self) -> bool { - self.flags.contains(ProviderStreamFactoryFlags::PipeStream) - } + fn is_piped(&self) -> bool { self.flags.contains(ProviderStreamFactoryFlags::PipeStream) } #[inline] - fn is_buffer_enabled(&self) -> bool { - self.flags.contains(ProviderStreamFactoryFlags::BufferEnabled) - } + fn is_buffer_enabled(&self) -> bool { self.flags.contains(ProviderStreamFactoryFlags::BufferEnabled) } #[inline] - fn is_shared_stream(&self) -> bool { - self.flags.contains(ProviderStreamFactoryFlags::ShareStream) - } + fn is_shared_stream(&self) -> bool { self.flags.contains(ProviderStreamFactoryFlags::ShareStream) } #[inline] - pub(crate) fn get_buffer_size(&self) -> usize { - self.buffer_size - } + pub(crate) fn get_buffer_size(&self) -> usize { self.buffer_size } #[inline] - pub fn get_reconnect_flag_clone(&self) -> Arc { - Arc::clone(&self.reconnect_flag) - } + pub fn get_reconnect_flag_clone(&self) -> Arc { Arc::clone(&self.reconnect_flag) } #[inline] - pub fn cancel_reconnect(&self) { - self.reconnect_flag.notify(); - } + pub fn cancel_reconnect(&self) { self.reconnect_flag.notify(); } #[inline] - pub fn get_url(&self) -> &Url { - &self.url - } + pub fn get_url(&self) -> &Url { &self.url } #[inline] - pub fn get_url_as_str(&self) -> &str { - self.url.as_str() - } + pub fn get_url_as_str(&self) -> &str { self.url.as_str() } #[inline] - pub fn should_reconnect(&self) -> bool { - self.flags - .contains(ProviderStreamFactoryFlags::ReconnectEnabled) - } + pub fn should_reconnect(&self) -> bool { self.flags.contains(ProviderStreamFactoryFlags::ReconnectEnabled) } #[inline] - pub fn get_headers(&self) -> &HeaderMap { - &self.headers - } + pub fn get_headers(&self) -> &HeaderMap { &self.headers } #[inline] pub fn get_total_bytes_send(&self) -> Option { - self.range_bytes - .as_ref() - .as_ref() - .map(|atomic| atomic.load(Ordering::Acquire)) + self.range_bytes.as_ref().as_ref().map(|atomic| atomic.load(Ordering::Acquire)) } // pub fn get_range_bytes(&self) -> &Arc> { @@ -198,20 +176,13 @@ impl ProviderStreamFactoryOptions { // } #[inline] - pub fn get_range_bytes_clone(&self) -> Arc> { - Arc::clone(&self.range_bytes) - } + pub fn get_range_bytes_clone(&self) -> Arc> { Arc::clone(&self.range_bytes) } #[inline] - pub fn should_continue(&self) -> bool { - self.reconnect_flag.is_active() - } + pub fn should_continue(&self) -> bool { self.reconnect_flag.is_active() } #[inline] - pub fn was_range_requested(&self) -> bool { - self.flags.contains(ProviderStreamFactoryFlags::RangeRequested) - } - + pub fn was_range_requested(&self) -> bool { self.flags.contains(ProviderStreamFactoryFlags::RangeRequested) } } fn get_request_range_start_bytes(req_headers: &HashMap>) -> Option { @@ -264,15 +235,12 @@ fn prepare_client( remove_sensitive_headers_on_cross_origin(&mut headers, original_url, url_override); prepare_default_headers(&mut headers, stream_options); - let partial = prepare_partial_request_headers(&mut headers, stream_options, range_start); + let partial = prepare_partial_request_headers(&mut headers, stream_options, range_start); if log_enabled!(log::Level::Debug) { let message = format!( "Stream requested with headers: {:?}", - headers - .iter() - .map(|header| (header.0, String::from_utf8_lossy(header.1.as_ref()))) - .collect::>() + headers.iter().map(|header| (header.0, String::from_utf8_lossy(header.1.as_ref()))).collect::>() ); debug!("{}", sanitize_sensitive_info(&message)); } @@ -291,11 +259,9 @@ fn remove_sensitive_headers_on_cross_origin( return; }; - let cross_origin = - override_url.scheme() != original_url.scheme() - || override_url.host_str() != original_url.host_str() - || override_url.port_or_known_default() - != original_url.port_or_known_default(); + let cross_origin = override_url.scheme() != original_url.scheme() + || override_url.host_str() != original_url.host_str() + || override_url.port_or_known_default() != original_url.port_or_known_default(); if !cross_origin { return; @@ -308,10 +274,7 @@ fn remove_sensitive_headers_on_cross_origin( fn prepare_default_headers(headers: &mut axum::http::HeaderMap, stream_options: &ProviderStreamFactoryOptions) { // Force Connection: close so the provider releases its slot immediately when the stream ends. // This prevents 509 errors from providers counting idle pooled connections against limits. - headers.insert( - axum::http::header::CONNECTION, - axum::http::header::HeaderValue::from_static("close"), - ); + headers.insert(axum::http::header::CONNECTION, axum::http::header::HeaderValue::from_static("close")); if !headers.contains_key(axum::http::header::USER_AGENT) { headers.insert( @@ -324,7 +287,11 @@ fn prepare_default_headers(headers: &mut axum::http::HeaderMap, stream_options: } } -fn prepare_partial_request_headers(headers: &mut HeaderMap, stream_options: &ProviderStreamFactoryOptions, range_start: Option) -> bool { +fn prepare_partial_request_headers( + headers: &mut HeaderMap, + stream_options: &ProviderStreamFactoryOptions, + range_start: Option, +) -> bool { if let Some(range) = range_start { if range > 0 || stream_options.was_range_requested() { let range_header = format!("bytes={range}-"); @@ -341,25 +308,14 @@ fn prepare_partial_request_headers(headers: &mut HeaderMap, stream_options: &Pro } fn collect_debug_headers(headers: &HeaderMap) -> Vec<(String, String)> { - const HEADER_NAMES: [&str; 8] = [ - "proxy-authenticate", - "via", - "server", - "location", - "x-cache", - "x-cache-status", - "x-served-by", - "x-proxy-id", - ]; + const HEADER_NAMES: [&str; 8] = + ["proxy-authenticate", "via", "server", "location", "x-cache", "x-cache-status", "x-served-by", "x-proxy-id"]; HEADER_NAMES .iter() .filter_map(|name| { headers.get_all(*name).iter().next().map(|value| { - let value = value - .to_str() - .unwrap_or("") - .to_string(); + let value = value.to_str().unwrap_or("").to_string(); ((*name).to_string(), value) }) }) @@ -381,8 +337,9 @@ async fn send_with_manual_redirects( ¤t_url, provider.as_ref(), true, - |resolved_url| prepare_client(request_client, stream_options, Some(resolved_url)).0 - ).await; + |resolved_url| prepare_client(request_client, stream_options, Some(resolved_url)).0, + ) + .await; let response = match result { Ok(resp) => resp, @@ -408,9 +365,7 @@ async fn send_with_manual_redirects( let Ok(location_str) = location.to_str() else { return Ok(response); }; - let next_url = current_url - .join(location_str) - .or_else(|_| Url::parse(location_str)); + let next_url = current_url.join(location_str).or_else(|_| Url::parse(location_str)); let Ok(next_url) = next_url else { return Ok(response); }; @@ -436,16 +391,11 @@ async fn provider_stream_request( let url = stream_options.get_url(); let provider = stream_options.get_provider().cloned(); - send_with_retry_and_provider( - &app_state.app_config, - url, - provider.as_ref(), - false, - |resolved_url| { - let (client, _partial_content) = prepare_client(request_client, stream_options, Some(resolved_url)); - client - } - ).await + send_with_retry_and_provider(&app_state.app_config, url, provider.as_ref(), false, |resolved_url| { + let (client, _partial_content) = prepare_client(request_client, stream_options, Some(resolved_url)); + client + }) + .await }; match response_result { Ok(mut response) => { @@ -453,9 +403,8 @@ async fn provider_stream_request( let response_url = response.url().clone(); if log_enabled!(log::Level::Debug) && !status.is_success() { let debug_headers = collect_debug_headers(response.headers()); - let message = format!( - "Provider response error: status={status}, url={response_url}, headers={debug_headers:?}" - ); + let message = + format!("Provider response error: status={status}, url={response_url}, headers={debug_headers:?}"); debug!("{}", sanitize_sensitive_info(&message)); } if status.is_success() { @@ -471,16 +420,10 @@ async fn provider_stream_request( debug!("{}", sanitize_sensitive_info(&message)); } - let response_headers: Vec<(String, String)> = - get_response_headers(response.headers()); + let response_headers: Vec<(String, String)> = get_response_headers(response.headers()); //let url = stream_options.get_url(); // debug!("First headers {headers:?} {} {}", sanitize_sensitive_info(url.as_str())); - Some(( - response_headers, - response.status(), - Some(response.url().clone()), - None, - )) + Some((response_headers, response.status(), Some(response.url().clone()), None)) }; let provider_stream = response @@ -501,9 +444,7 @@ async fn provider_stream_request( | StatusCode::UNAUTHORIZED | StatusCode::PROXY_AUTHENTICATION_REQUIRED | StatusCode::METHOD_NOT_ALLOWED - | StatusCode::BAD_REQUEST => { - handle_channel_unavailable_stream(app_state, stream_options).await - } + | StatusCode::BAD_REQUEST => handle_channel_unavailable_stream(app_state, stream_options).await, _ => Err(status), }; } @@ -513,9 +454,7 @@ async fn provider_stream_request( StatusCode::INTERNAL_SERVER_ERROR | StatusCode::BAD_GATEWAY | StatusCode::SERVICE_UNAVAILABLE - | StatusCode::GATEWAY_TIMEOUT => { - handle_channel_unavailable_stream(app_state, stream_options).await - } + | StatusCode::GATEWAY_TIMEOUT => handle_channel_unavailable_stream(app_state, stream_options).await, _ => Err(status), }; } @@ -528,16 +467,21 @@ async fn provider_stream_request( } } -async fn handle_channel_unavailable_stream(app_state: &Arc, - stream_options: &ProviderStreamFactoryOptions +async fn handle_channel_unavailable_stream( + app_state: &Arc, + stream_options: &ProviderStreamFactoryOptions, ) -> Result, StatusCode> { - app_state.connection_manager.update_stream_detail(&stream_options.addr, CustomVideoStreamType::ChannelUnavailable).await; + app_state + .connection_manager + .update_stream_detail(&stream_options.addr, CustomVideoStreamType::ChannelUnavailable) + .await; app_state.connection_manager.release_provider_connection(&stream_options.addr).await; - if let (Some(boxed_provider_stream), response_info) = - create_channel_unavailable_stream(&app_state.app_config,&get_response_headers(stream_options.get_headers()), - StatusCode::SERVICE_UNAVAILABLE) - { + if let (Some(boxed_provider_stream), response_info) = create_channel_unavailable_stream( + &app_state.app_config, + &get_response_headers(stream_options.get_headers()), + StatusCode::SERVICE_UNAVAILABLE, + ) { Ok(Some((boxed_provider_stream, response_info))) } else { Err(StatusCode::SERVICE_UNAVAILABLE) @@ -561,18 +505,34 @@ async fn get_provider_stream( } Ok(None) => { if connect_err > ERR_MAX_RETRY_COUNT { - warn!("The stream could be unavailable. {}", sanitize_sensitive_info(stream_options.get_url().as_str())); + warn!( + "The stream could be unavailable. {}", + sanitize_sensitive_info(stream_options.get_url().as_str()) + ); break; } } Err(status) => { debug!("Provider stream response error status response : {status}"); - if matches!(status, StatusCode::FORBIDDEN | StatusCode::SERVICE_UNAVAILABLE | StatusCode::UNAUTHORIZED | StatusCode::PROXY_AUTHENTICATION_REQUIRED | StatusCode::RANGE_NOT_SATISFIABLE) { - warn!("The stream could be unavailable. ({status}) {}",sanitize_sensitive_info(stream_options.get_url().as_str())); + if matches!( + status, + StatusCode::FORBIDDEN + | StatusCode::SERVICE_UNAVAILABLE + | StatusCode::UNAUTHORIZED + | StatusCode::PROXY_AUTHENTICATION_REQUIRED + | StatusCode::RANGE_NOT_SATISFIABLE + ) { + warn!( + "The stream could be unavailable. ({status}) {}", + sanitize_sensitive_info(stream_options.get_url().as_str()) + ); break; } if connect_err > ERR_MAX_RETRY_COUNT { - warn!("The stream could be unavailable. ({status}) {}",sanitize_sensitive_info(stream_options.get_url().as_str())); + warn!( + "The stream could be unavailable. ({status}) {}", + sanitize_sensitive_info(stream_options.get_url().as_str()) + ); break; } } @@ -604,27 +564,19 @@ pub async fn create_provider_stream( stream_options: ProviderStreamFactoryOptions, ) -> Option { let client_stream_factory = |stream, reconnect_flag, range_cnt| { - let stream = if !stream_options.is_piped() - && stream_options.is_buffer_enabled() - && !stream_options.is_shared_stream() - { - BufferedStream::new( - stream, - stream_options.get_buffer_size(), - stream_options.get_reconnect_flag_clone(), - stream_options.get_url_as_str(), - ) - .boxed() - } else { - stream - }; - ClientStream::new( - stream, - reconnect_flag, - range_cnt, - stream_options.get_url_as_str(), - ) - .boxed() + let stream = + if !stream_options.is_piped() && stream_options.is_buffer_enabled() && !stream_options.is_shared_stream() { + BufferedStream::new( + stream, + stream_options.get_buffer_size(), + stream_options.get_reconnect_flag_clone(), + stream_options.get_url_as_str(), + ) + .boxed() + } else { + stream + }; + ClientStream::new(stream, reconnect_flag, range_cnt, stream_options.get_url_as_str()).boxed() }; match get_provider_stream(app_state, client, &stream_options).await { @@ -653,7 +605,10 @@ pub async fn create_provider_stream( match get_provider_stream(&app_state_clone, &client, &stream_opts).await { Ok(Some((stream, _info))) => Some((stream, ())), Ok(None) => { - app_state_clone.connection_manager.release_provider_connection(&stream_opts.addr).await; + app_state_clone + .connection_manager + .release_provider_connection(&stream_opts.addr) + .await; continue_streaming.notify(); if let (Some(boxed_provider_stream), _response_info) = create_channel_unavailable_stream( @@ -667,7 +622,10 @@ pub async fn create_provider_stream( None } Err(status) => { - app_state_clone.connection_manager.release_provider_connection(&stream_opts.addr).await; + app_state_clone + .connection_manager + .release_provider_connection(&stream_opts.addr) + .await; continue_streaming.notify(); if let (Some(boxed_provider_stream), _response_info) = create_channel_unavailable_stream( @@ -728,19 +686,15 @@ pub async fn create_provider_stream( #[cfg(test)] mod tests { use super::*; - use shared::model::PlaylistItemType; use axum::http::HeaderMap; + use shared::model::PlaylistItemType; #[test] fn test_provider_stream_factory_options_range_logic() { let addr = "127.0.0.1:8080".parse().unwrap(); let stream_url = Url::parse("http://example.com/stream").unwrap(); - let stream_options = StreamOptions { - stream_retry: true, - buffer_enabled: true, - buffer_size: 1024, - pipe_provider_stream: false, - }; + let stream_options = + StreamOptions { stream_retry: true, buffer_enabled: true, buffer_size: 1024, pipe_provider_stream: false }; let disabled_headers = None; // Case 1: VOD, no initial range requested @@ -806,6 +760,6 @@ mod tests { None, ); assert!(!options.was_range_requested()); // Stripped by filter - assert_eq!(options.get_total_bytes_send(), None); + assert_eq!(options.get_total_bytes_send(), None); } } diff --git a/backend/src/api/model/streams/provisioning_stream.rs b/backend/src/api/model/streams/provisioning_stream.rs index 521d66471..ca9397168 100644 --- a/backend/src/api/model/streams/provisioning_stream.rs +++ b/backend/src/api/model/streams/provisioning_stream.rs @@ -1,11 +1,14 @@ -use crate::api::model::StreamError; -use crate::api::model::TransportStreamBuffer; -use crate::tools::atomic_once_flag::AtomicOnceFlag; +use crate::{ + api::model::{StreamError, TransportStreamBuffer}, + tools::atomic_once_flag::AtomicOnceFlag, +}; use bytes::Bytes; use futures::Stream; -use std::pin::Pin; -use std::sync::Arc; -use std::task::{Context, Poll}; +use std::{ + pin::Pin, + sync::Arc, + task::{Context, Poll}, +}; pub struct ProvisioningStream { buffer: TransportStreamBuffer, @@ -13,9 +16,7 @@ pub struct ProvisioningStream { } impl ProvisioningStream { - pub fn new(buffer: TransportStreamBuffer, stop_signal: Arc) -> Self { - Self { buffer, stop_signal } - } + pub fn new(buffer: TransportStreamBuffer, stop_signal: Arc) -> Self { Self { buffer, stop_signal } } } impl Stream for ProvisioningStream { diff --git a/backend/src/api/model/streams/shared_stream_manager.rs b/backend/src/api/model/streams/shared_stream_manager.rs index 949f74af8..e488b8183 100644 --- a/backend/src/api/model/streams/shared_stream_manager.rs +++ b/backend/src/api/model/streams/shared_stream_manager.rs @@ -1,26 +1,28 @@ -use crate::api::model::{AppState, STREAM_IDLE_TIMEOUT}; -use crate::api::model::{ActiveProviderManager, ProviderHandle, StreamError}; -use crate::model::Config; -use crate::utils::debug_if_enabled; +use crate::{ + api::model::{ + streams::buffered_stream::CHANNEL_SIZE, ActiveProviderManager, AppState, BoxedProviderStream, ProviderHandle, + StreamError, STREAM_IDLE_TIMEOUT, + }, + model::Config, + utils::{debug_if_enabled, trace_if_enabled}, +}; use bytes::Bytes; -use futures::stream::BoxStream; -use futures::{Stream, StreamExt}; -use std::collections::{HashMap, VecDeque}; -use std::fmt; -use std::fmt::{Debug, Formatter}; -use std::net::SocketAddr; -use std::sync::Arc; - -use crate::api::model::streams::buffered_stream::CHANNEL_SIZE; -use crate::api::model::BoxedProviderStream; -use crate::utils::trace_if_enabled; +use futures::{stream::BoxStream, Stream, StreamExt}; use log::{debug, trace, warn}; use shared::utils::sanitize_sensitive_info; -use std::pin::Pin; -use std::task::{Context, Poll}; -use tokio::sync::mpsc::Sender; -use tokio::sync::{mpsc, RwLock}; -use tokio::time::{sleep, Duration, Instant}; +use std::{ + collections::{HashMap, VecDeque}, + fmt, + fmt::{Debug, Formatter}, + net::SocketAddr, + pin::Pin, + sync::Arc, + task::{Context, Poll}, +}; +use tokio::{ + sync::{mpsc, mpsc::Sender, RwLock}, + time::{sleep, Duration, Instant}, +}; use tokio_stream::wrappers::ReceiverStream; use tokio_util::sync::CancellationToken; @@ -37,7 +39,7 @@ struct ReceiverStreamWrapper { impl Stream for ReceiverStreamWrapper where - S: Stream + Unpin, + S: Stream + Unpin, { type Item = Result; @@ -66,7 +68,6 @@ fn convert_stream(stream: BoxStream) -> BoxStream>, buffer_size: usize, @@ -85,16 +86,10 @@ impl Debug for BurstBuffer { impl BurstBuffer { pub fn new(buf_size: usize) -> Self { - Self { - buffer: VecDeque::with_capacity(buf_size), - buffer_size: buf_size, - current_bytes: 0, - } + Self { buffer: VecDeque::with_capacity(buf_size), buffer_size: buf_size, current_bytes: 0 } } - pub fn snapshot(&self) -> VecDeque> { - self.buffer.iter().cloned().collect::>>() - } + pub fn snapshot(&self) -> VecDeque> { self.buffer.iter().cloned().collect::>>() } pub fn push(&mut self, packet: Arc) { while self.current_bytes + packet.len() > self.buffer_size { @@ -110,13 +105,15 @@ impl BurstBuffer { } } - async fn send_burst_buffer( start_buffer: &VecDeque>, client_tx: &Sender, - cancellation_token: &CancellationToken) { + cancellation_token: &CancellationToken, +) { for buf in start_buffer { - if cancellation_token.is_cancelled() { return; } + if cancellation_token.is_cancelled() { + return; + } if let Err(err) = client_tx.send(buf.as_ref().clone()).await { warn!("Error sending burst-buffer chunk to client: {err}"); return; // stop on send error @@ -140,8 +137,12 @@ pub struct SharedStreamState { } impl SharedStreamState { - fn new(headers: Vec<(String, String)>, buf_size: usize, - provider_guard: Option, min_burst_buffer_size: usize) -> Self { + fn new( + headers: Vec<(String, String)>, + buf_size: usize, + provider_guard: Option, + min_burst_buffer_size: usize, + ) -> Self { let (broadcaster, _) = tokio::sync::broadcast::channel(buf_size); // TODO channel size versus byte size, channels are chunk sized, burst_buffer byte sized let burst_buffer_size_in_bytes = min_burst_buffer_size.max(buf_size * 1024 * 12); @@ -157,7 +158,11 @@ impl SharedStreamState { } } - async fn subscribe(&self, addr: &SocketAddr, manager: Arc) -> (BoxedProviderStream, Option>) { + async fn subscribe( + &self, + addr: &SocketAddr, + manager: Arc, + ) -> (BoxedProviderStream, Option>) { let (client_tx, client_rx) = mpsc::channel(self.buf_size); let mut broadcast_rx = self.broadcaster.subscribe(); let cancel_token = CancellationToken::new(); @@ -170,8 +175,11 @@ impl SharedStreamState { { let mut subs = self.subscribers.write().await; subs.insert(*addr, cancel_token.clone()); - debug_if_enabled!("Shared stream subscriber added {}; total subscribers={}", - sanitize_sensitive_info(&addr.to_string()), subs.len()); + debug_if_enabled!( + "Shared stream subscriber added {}; total subscribers={}", + sanitize_sensitive_info(&addr.to_string()), + subs.len() + ); } let client_tx_clone = client_tx.clone(); @@ -196,62 +204,62 @@ impl SharedStreamState { let mut loop_cnt = 0; loop { tokio::select! { - biased; + biased; - // canceled - () = cancel_token.cancelled() => { - debug!("Client disconnected from shared stream: {address}"); - break; - } - - // timeout handling - () = sleep(Duration::from_secs(1)) => { - if last_active.elapsed() > timeout_duration { - debug!("Client timed out due to inactivity: {address}"); - cancel_token.cancel(); + // canceled + () = cancel_token.cancelled() => { + debug!("Client disconnected from shared stream: {address}"); break; } - } - // receive broadcast data - result = broadcast_rx.recv() => { - match result { - Ok(data) => { - // If the client press pause, skip - if client_tx_clone.is_closed() { - continue; - } - - if let Err(err) = client_tx.send(data).await { - debug!("Shared stream client send error: {address} {err}"); - break; - } - loop_cnt += 1; - last_active = Instant::now(); - - if loop_cnt >= yield_counter { - tokio::task::yield_now().await; - loop_cnt = 0; - } + // timeout handling + () = sleep(Duration::from_secs(1)) => { + if last_active.elapsed() > timeout_duration { + debug!("Client timed out due to inactivity: {address}"); + cancel_token.cancel(); + break; } - Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => { - if last_lag_log.elapsed() > Duration::from_secs(5) { - let buffered_bytes = { - let buffer = burst_buffer_for_log.read().await; - buffer.current_bytes - }; - warn!("Shared stream client lagged behind {address}. Skipped {skipped} messages (buffered {buffered_bytes} bytes, yield counter {yield_counter})"); - last_lag_log = Instant::now(); + } + + // receive broadcast data + result = broadcast_rx.recv() => { + match result { + Ok(data) => { + // If the client press pause, skip + if client_tx_clone.is_closed() { + continue; + } + + if let Err(err) = client_tx.send(data).await { + debug!("Shared stream client send error: {address} {err}"); + break; + } + loop_cnt += 1; + last_active = Instant::now(); + + if loop_cnt >= yield_counter { + tokio::task::yield_now().await; + loop_cnt = 0; + } } - // Sleep to prevent CPU spin when client is persistently lagging behind. - // This provides backpressure for slow clients that can't keep up. - sleep(Duration::from_millis(50)).await; + Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => { + if last_lag_log.elapsed() > Duration::from_secs(5) { + let buffered_bytes = { + let buffer = burst_buffer_for_log.read().await; + buffer.current_bytes + }; + warn!("Shared stream client lagged behind {address}. Skipped {skipped} messages (buffered {buffered_bytes} bytes, yield counter {yield_counter})"); + last_lag_log = Instant::now(); + } + // Sleep to prevent CPU spin when client is persistently lagging behind. + // This provides backpressure for slow clients that can't keep up. + sleep(Duration::from_millis(50)).await; + } + Err(_) => break, } - Err(_) => break, } } } - } manager.release_connection(&address, false).await; }); @@ -262,14 +270,9 @@ impl SharedStreamState { (convert_stream(ReceiverStream::new(client_rx).boxed()), provider) } - fn broadcast( - &self, - stream_url: &str, - bytes_stream: S, - shared_streams: Arc, - ) + fn broadcast(&self, stream_url: &str, bytes_stream: S, shared_streams: Arc) where - S: Stream> + Unpin + 'static + Send, + S: Stream> + Unpin + 'static + Send, E: std::fmt::Debug + Send, { let mut source_stream = Box::pin(bytes_stream); @@ -286,61 +289,64 @@ impl SharedStreamState { loop { tokio::select! { - biased; + biased; - () = stop_token.cancelled() => { - debug_if_enabled!("No shared stream subscribers left. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url)); + () = stop_token.cancelled() => { + debug_if_enabled!("No shared stream subscribers left. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url)); + break; + }, + + () = &mut idle => { + debug!("shared stream idle for too long, closing"); + stop_token.cancel(); break; - }, + } - () = &mut idle => { - debug!("shared stream idle for too long, closing"); - stop_token.cancel(); - break; - } + chunk = source_stream.next() => { + idle.as_mut().reset(Instant::now() + idle_timeout); + match chunk { + Some(Ok(data)) => { + let arc_data = Arc::new(data); + { + let mut buffer = burst_buffer.write().await; + buffer.push(arc_data.clone()); + } - chunk = source_stream.next() => { - idle.as_mut().reset(Instant::now() + idle_timeout); - match chunk { - Some(Ok(data)) => { - let arc_data = Arc::new(data); - { - let mut buffer = burst_buffer.write().await; - buffer.push(arc_data.clone()); - } + match sender.send(arc_data.as_ref().clone()) { + Ok(clients) => { + if clients == 0 { + debug_if_enabled!("No shared stream subscribers closing {}", sanitize_sensitive_info(&streaming_url)); + break; + } + counter += 1; + if counter >= YIELD_COUNTER { + tokio::task::yield_now().await; + counter = 0; + } + } + Err(_e) => { + debug_if_enabled!("Shared stream send error,no subscribers closing {}", sanitize_sensitive_info(&streaming_url)); + break; + } + } + } + Some(Err(e)) => { + trace!("Shared stream received error: {e:?}"); + tokio::task::yield_now().await; - match sender.send(arc_data.as_ref().clone()) { - Ok(clients) => { - if clients == 0 { - debug_if_enabled!("No shared stream subscribers closing {}", sanitize_sensitive_info(&streaming_url)); - break; - } - counter += 1; - if counter >= YIELD_COUNTER { - tokio::task::yield_now().await; - counter = 0; - } - } - Err(_e) => { - debug_if_enabled!("Shared stream send error,no subscribers closing {}", sanitize_sensitive_info(&streaming_url)); - break; - } - } - } - Some(Err(e)) => { - trace!("Shared stream received error: {e:?}"); - tokio::task::yield_now().await; - - } - None => { - debug_if_enabled!("Source stream ended. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url)); - break; - } - } - }, - } + } + None => { + debug_if_enabled!("Source stream ended. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url)); + break; + } + } + }, + } } - debug_if_enabled!("Shared stream exhausted. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url)); + debug_if_enabled!( + "Shared stream exhausted. Closing shared provider stream {}", + sanitize_sensitive_info(&streaming_url) + ); shared_streams.unregister(&streaming_url, false).await; }); } @@ -350,7 +356,6 @@ impl SharedStreamState { struct SharedStreamsRegister { by_key: HashMap>, key_by_addr: HashMap, - } pub struct SharedStreamManager { @@ -360,10 +365,7 @@ pub struct SharedStreamManager { impl SharedStreamManager { pub(crate) fn new(provider_manager: Arc) -> Self { - Self { - provider_manager, - shared_streams: RwLock::new(SharedStreamsRegister::default()), - } + Self { provider_manager, shared_streams: RwLock::new(SharedStreamsRegister::default()) } } pub async fn get_shared_state(&self, stream_url: &str) -> Option> { @@ -378,7 +380,8 @@ impl SharedStreamManager { let shared_state_opt = { let mut shared_streams = self.shared_streams.write().await; - let remove_keys: Vec = shared_streams.key_by_addr + let remove_keys: Vec = shared_streams + .key_by_addr .iter() .filter_map(|(addr, url)| if url == stream_url { Some(*addr) } else { None }) .collect(); @@ -427,15 +430,18 @@ impl SharedStreamManager { (tx, is_empty, subs.len()) }; - debug_if_enabled!("Shared stream subscriber removed {}; remaining subscribers={remaining}", sanitize_sensitive_info(&addr.to_string())); + debug_if_enabled!( + "Shared stream subscriber removed {}; remaining subscribers={remaining}", + sanitize_sensitive_info(&addr.to_string()) + ); if is_empty { if let Some(url) = stream_url.as_ref() { debug_if_enabled!( - "No subscribers remain for {} after removing {}", - sanitize_sensitive_info(url), - sanitize_sensitive_info(&addr.to_string()) - ); + "No subscribers remain for {} after removing {}", + sanitize_sensitive_info(url), + sanitize_sensitive_info(&addr.to_string()) + ); self.unregister(url, send_stop_signal).await; } } @@ -463,21 +469,26 @@ impl SharedStreamManager { }; if let Some(shared_state) = shared_state_opt { - debug_if_enabled!("Responding to existing shared client stream {} {}", - sanitize_sensitive_info(&addr.to_string()), sanitize_sensitive_info(stream_url)); + debug_if_enabled!( + "Responding to existing shared client stream {} {}", + sanitize_sensitive_info(&addr.to_string()), + sanitize_sensitive_info(stream_url) + ); Some(shared_state.subscribe(addr, manager).await) } else { None } } - async fn register(&self, addr: &SocketAddr, stream_url: &str, shared_state: Arc) { let mut shared_streams = self.shared_streams.write().await; shared_streams.by_key.insert(stream_url.to_string(), shared_state); shared_streams.key_by_addr.insert(*addr, stream_url.to_string()); - debug_if_enabled!("Registered shared stream {} for initial subscriber {}", - sanitize_sensitive_info(stream_url), sanitize_sensitive_info(&addr.to_string())); + debug_if_enabled!( + "Registered shared stream {} for initial subscriber {}", + sanitize_sensitive_info(stream_url), + sanitize_sensitive_info(&addr.to_string()) + ); } pub(crate) async fn register_shared_stream( @@ -487,9 +498,10 @@ impl SharedStreamManager { addr: &SocketAddr, headers: Vec<(String, String)>, buffer_size: usize, - provider_handle: Option) -> Option<(BoxedProviderStream, Option>)> + provider_handle: Option, + ) -> Option<(BoxedProviderStream, Option>)> where - S: Stream> + Unpin + 'static + Send, + S: Stream> + Unpin + 'static + Send, E: std::fmt::Debug + Send, { let buf_size = CHANNEL_SIZE.max(buffer_size); diff --git a/backend/src/api/model/streams/throttled_stream.rs b/backend/src/api/model/streams/throttled_stream.rs index 8c7227140..f9f8cb480 100644 --- a/backend/src/api/model/streams/throttled_stream.rs +++ b/backend/src/api/model/streams/throttled_stream.rs @@ -1,8 +1,8 @@ use crate::api::model::StreamError; use bytes::Bytes; use futures::Stream; -use std::future::Future; use std::{ + future::Future, pin::Pin, task::{Context, Poll}, time::Duration, @@ -19,18 +19,14 @@ impl ThrottledStream { #[allow(clippy::cast_precision_loss)] pub fn new(inner: S, throttle_kbps: usize) -> Self { assert!(throttle_kbps > 0, "Rate must be greater than 0"); - let rate_bytes_per_sec = (throttle_kbps as f64) * 1000.0 / 8.0; - Self { - inner, - rate_bytes_per_sec, - next_delay: None, - } + let rate_bytes_per_sec = (throttle_kbps as f64) * 1000.0 / 8.0; + Self { inner, rate_bytes_per_sec, next_delay: None } } } impl Stream for ThrottledStream where - S: Stream> + Unpin, + S: Stream> + Unpin, { type Item = Result; @@ -74,4 +70,4 @@ where } } -impl Unpin for ThrottledStream {} \ No newline at end of file +impl Unpin for ThrottledStream {} diff --git a/backend/src/api/model/streams/timed_client_stream.rs b/backend/src/api/model/streams/timed_client_stream.rs index 16fae771d..c57f70e72 100644 --- a/backend/src/api/model/streams/timed_client_stream.rs +++ b/backend/src/api/model/streams/timed_client_stream.rs @@ -1,15 +1,20 @@ -use std::net::SocketAddr; -use crate::api::model::stream_error::StreamError; +use crate::{ + api::model::{stream_error::StreamError, AppState, BoxedProviderStream}, + utils::debug_if_enabled, +}; use bytes::Bytes; use futures::Stream; -use std::pin::Pin; -use std::sync::Arc; -use std::task::Poll; -use std::time::{Duration, Instant}; -use shared::model::VirtualId; -use shared::utils::{default_kick_secs, sanitize_sensitive_info}; -use crate::api::model::{AppState, BoxedProviderStream}; -use crate::utils::debug_if_enabled; +use shared::{ + model::VirtualId, + utils::{default_kick_secs, sanitize_sensitive_info}, +}; +use std::{ + net::SocketAddr, + pin::Pin, + sync::Arc, + task::Poll, + time::{Duration, Instant}, +}; pub struct TimedClientStream { inner: BoxedProviderStream, @@ -20,7 +25,13 @@ pub struct TimedClientStream { } impl TimedClientStream { - pub(crate) fn new(app_state: &Arc, inner: BoxedProviderStream, duration: u32, addr: SocketAddr, virtual_id: VirtualId) -> Self { + pub(crate) fn new( + app_state: &Arc, + inner: BoxedProviderStream, + duration: u32, + addr: SocketAddr, + virtual_id: VirtualId, + ) -> Self { let deadline = Instant::now() + Duration::from_secs(u64::from(duration)); Self { inner, deadline, app_state: Arc::clone(app_state), addr, virtual_id } } @@ -28,14 +39,23 @@ impl TimedClientStream { impl Stream for TimedClientStream { type Item = Result; - fn poll_next(mut self: Pin<&mut Self>,cx: &mut std::task::Context<'_>,) -> Poll> { + fn poll_next(mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>) -> Poll> { if Instant::now() >= self.deadline { - let kick_secs = self.app_state.app_config.config.load().web_ui.as_ref().map_or_else(default_kick_secs, |wc| wc.kick_secs); + let kick_secs = self + .app_state + .app_config + .config + .load() + .web_ui + .as_ref() + .map_or_else(default_kick_secs, |wc| wc.kick_secs); let connection_manager = Arc::clone(&self.app_state.connection_manager); let addr = self.addr; let virtual_id = self.virtual_id; - debug_if_enabled!("TimedClient stream exceeds time limit. Closing stream with virtual_id {virtual_id} for addr: {}", - sanitize_sensitive_info(&addr.to_string())); + debug_if_enabled!( + "TimedClient stream exceeds time limit. Closing stream with virtual_id {virtual_id} for addr: {}", + sanitize_sensitive_info(&addr.to_string()) + ); tokio::spawn(async move { connection_manager.kick_connection(&addr, virtual_id, kick_secs).await; }); @@ -43,4 +63,4 @@ impl Stream for TimedClientStream { } Pin::as_mut(&mut self.inner).poll_next(cx) } -} \ No newline at end of file +} diff --git a/backend/src/api/model/streams/transport_stream_buffer.rs b/backend/src/api/model/streams/transport_stream_buffer.rs index babce0676..b124f9f81 100644 --- a/backend/src/api/model/streams/transport_stream_buffer.rs +++ b/backend/src/api/model/streams/transport_stream_buffer.rs @@ -1,11 +1,13 @@ use bytes::{Bytes, BytesMut}; use futures::task::AtomicWaker; -use std::collections::{HashMap, HashSet}; -use std::sync::Arc; -use std::task::Waker; +use std::{ + collections::{HashMap, HashSet}, + sync::Arc, + task::Waker, +}; -const MAX_PCR: u64 = 1 << 42; // 42 bit PCR cycle -const MAX_PTS_DTS: u64 = 1 << 33; // 33 bit PTS/DTS cycle +const MAX_PCR: u64 = 1 << 42; // 42 bit PCR cycle +const MAX_PTS_DTS: u64 = 1 << 33; // 33 bit PTS/DTS cycle const TS_PACKET_SIZE: usize = 188; const SYNC_BYTE: u8 = 0x47; @@ -196,7 +198,13 @@ pub fn extract_pts_dts_indices_with_continuity(ts_data: &[u8]) -> TsInfoExtracti } /// Replace PTS and DTS timestamps in the TS packet slice -fn replace_pts_dts(packet_slice: &[u8], pts_index: usize, dts_index: Option, new_presentation_ts: u64, new_decoding_ts: u64) -> Vec { +fn replace_pts_dts( + packet_slice: &[u8], + pts_index: usize, + dts_index: Option, + new_presentation_ts: u64, + new_decoding_ts: u64, +) -> Vec { let new_presentation_ts_bytes = encode_timestamp(new_presentation_ts); let mut new_packet = Vec::with_capacity(packet_slice.len()); @@ -266,7 +274,7 @@ pub fn calculate_duration_ticks(buffer: &[u8], packet_indices: &PacketIndices) - let mut last_pts: Option = None; let mut count = 0; - // We already calculated average diff/duration in `extract_pts_dts_indices_with_continuity` + // We already calculated average diff/duration in `extract_pts_dts_indices_with_continuity` // but we didn't expose it. We can re-estimate it here or assume a default. // However, packet_indices stores `diff` in the tuple! `(pts, dts, diff)`. // But only for the first packet of a frame? @@ -418,12 +426,16 @@ impl TransportStreamBuffer { } } - pub fn register_waker(&self, waker: &Waker) { - self.waker.register(waker); - } + pub fn register_waker(&self, waker: &Waker) { self.waker.register(waker); } /// Generates a Discontinuity packet for the given packet/PID state. - fn generate_discontinuity_packet(_pid: u16, new_packet: &[u8], cc: u8, first_pcr: Option, timestamp_offset: u64) -> Vec { + fn generate_discontinuity_packet( + _pid: u16, + new_packet: &[u8], + cc: u8, + first_pcr: Option, + timestamp_offset: u64, + ) -> Vec { let mut pkt = vec![0xFF; TS_PACKET_SIZE]; pkt[0] = SYNC_BYTE; pkt[1] = new_packet[1] & 0x1F; @@ -537,7 +549,13 @@ impl TransportStreamBuffer { *discontinuity_sent = true; - let extra = Self::generate_discontinuity_packet(pid, &new_packet, extra_packet_cc, self.first_pcr, self.timestamp_offset); + let extra = Self::generate_discontinuity_packet( + pid, + &new_packet, + extra_packet_cc, + self.first_pcr, + self.timestamp_offset, + ); bytes.extend_from_slice(&extra); } else { // Payload packet gets current counter (N) @@ -582,7 +600,8 @@ impl TransportStreamBuffer { let orig_presentation_ts = decode_timestamp(&new_packet[pts_offset..pts_offset + 5]); let new_presentation_ts = (orig_presentation_ts + self.timestamp_offset) % MAX_PTS_DTS; - let replaced = replace_pts_dts(&new_packet, pts_offset, dts_offset, new_presentation_ts, new_decoding_ts); + let replaced = + replace_pts_dts(&new_packet, pts_offset, dts_offset, new_presentation_ts, new_decoding_ts); new_packet = replaced; } diff --git a/backend/src/api/model/update_guard.rs b/backend/src/api/model/update_guard.rs index dbe28f067..1d89e1de4 100644 --- a/backend/src/api/model/update_guard.rs +++ b/backend/src/api/model/update_guard.rs @@ -8,42 +8,22 @@ pub struct UpdateGuard { } impl Default for UpdateGuard { - fn default() -> Self { - Self { - playlist: Arc::new(Semaphore::new(1)), - library: Arc::new(Semaphore::new(1)), - } - } + fn default() -> Self { Self { playlist: Arc::new(Semaphore::new(1)), library: Arc::new(Semaphore::new(1)) } } } impl UpdateGuard { - pub fn new() -> Self { - Self::default() - } + pub fn new() -> Self { Self::default() } pub fn try_playlist(&self) -> Option { - self.playlist - .clone() - .try_acquire_owned() - .ok() - .map(|permit| UpdateGuardPermit { _permit: permit }) + self.playlist.clone().try_acquire_owned().ok().map(|permit| UpdateGuardPermit { _permit: permit }) } pub async fn acquire_playlist_lock(&self) -> Option { - self.playlist - .clone() - .acquire_owned() - .await - .ok() - .map(|permit| UpdateGuardPermit { _permit: permit }) + self.playlist.clone().acquire_owned().await.ok().map(|permit| UpdateGuardPermit { _permit: permit }) } pub fn try_library(&self) -> Option { - self.library - .clone() - .try_acquire_owned() - .ok() - .map(|permit| UpdateGuardPermit { _permit: permit }) + self.library.clone().try_acquire_owned().ok().map(|permit| UpdateGuardPermit { _permit: permit }) } } diff --git a/backend/src/api/model/xtream.rs b/backend/src/api/model/xtream.rs index 69b5c3c00..894cf8363 100644 --- a/backend/src/api/model/xtream.rs +++ b/backend/src/api/model/xtream.rs @@ -1,22 +1,23 @@ -use shared::utils::serialize_number_as_string; use crate::model::{ApiProxyServerInfo, ProxyUserCredentials}; use chrono::{Duration, Local}; use serde::{Deserialize, Serialize}; -use shared::model::ProxyUserStatus; -use shared::utils::CONSTANTS; +use shared::{ + model::ProxyUserStatus, + utils::{serialize_number_as_string, CONSTANTS}, +}; #[derive(Serialize, Deserialize, Clone)] pub struct XtreamUserInfoResponse { pub username: String, pub password: String, pub message: String, - pub auth: u16, // 0 | 1 + pub auth: u16, // 0 | 1 pub status: String, // "Active" - #[serde(serialize_with ="serialize_number_as_string")] + #[serde(serialize_with = "serialize_number_as_string")] pub exp_date: i64, //1628755200, pub is_trial: String, // 0 | 1 pub active_cons: String, - #[serde(serialize_with ="serialize_number_as_string")] + #[serde(serialize_with = "serialize_number_as_string")] pub created_at: i64, //1623429679, pub max_connections: String, pub allowed_output_formats: Vec, @@ -42,34 +43,46 @@ pub struct XtreamAuthorizationResponse { } impl XtreamAuthorizationResponse { - pub fn new(server_info: &ApiProxyServerInfo, user: &ProxyUserCredentials, active_connections: u32, access_control: bool) -> Self { + pub fn new( + server_info: &ApiProxyServerInfo, + user: &ProxyUserCredentials, + active_connections: u32, + access_control: bool, + ) -> Self { let now = Local::now(); let created_default = (now - Duration::days(365)).timestamp(); let expired_default = (now + Duration::days(365)).timestamp(); - let (created_at, exp_date, is_trial, max_connections, user_status) = - if access_control { - let exp_date = user.exp_date.as_ref().map_or(expired_default, |d| *d); - let is_expired = (exp_date - now.timestamp()) < 0; - let current_status = user.status.as_ref().unwrap_or(&ProxyUserStatus::Active); - let user_status = match current_status { - ProxyUserStatus::Active | ProxyUserStatus::Trial => if is_expired { &ProxyUserStatus::Expired } else { current_status }, - _ => current_status - }; - (user.created_at.as_ref().map_or(created_default, |d| *d), - exp_date, - user.status.as_ref().map_or("0", |s| if *s == ProxyUserStatus::Trial { "1" } else { "0" }).to_string(), - format!("{}", user.max_connections), - user_status - ) - } else { - (created_default, - expired_default, - "0".to_string(), - if user.max_connections == 0 { "1".to_string() } else { user.max_connections.to_string() }, - &ProxyUserStatus::Active, - ) + let (created_at, exp_date, is_trial, max_connections, user_status) = if access_control { + let exp_date = user.exp_date.as_ref().map_or(expired_default, |d| *d); + let is_expired = (exp_date - now.timestamp()) < 0; + let current_status = user.status.as_ref().unwrap_or(&ProxyUserStatus::Active); + let user_status = match current_status { + ProxyUserStatus::Active | ProxyUserStatus::Trial => { + if is_expired { + &ProxyUserStatus::Expired + } else { + current_status + } + } + _ => current_status, }; + ( + user.created_at.as_ref().map_or(created_default, |d| *d), + exp_date, + user.status.as_ref().map_or("0", |s| if *s == ProxyUserStatus::Trial { "1" } else { "0" }).to_string(), + format!("{}", user.max_connections), + user_status, + ) + } else { + ( + created_default, + expired_default, + "0".to_string(), + if user.max_connections == 0 { "1".to_string() } else { user.max_connections.to_string() }, + &ProxyUserStatus::Active, + ) + }; Self { user_info: XtreamUserInfoResponse { @@ -103,7 +116,7 @@ impl XtreamAuthorizationResponse { timestamp_now: now.timestamp(), time_now: now.format("%Y-%m-%d %H:%M:%S").to_string(), // We don't know what this field is good for, but it is in the response from XtreamCodes, so we will include it. - process: true + process: true, }, } } diff --git a/backend/src/api/panel_api.rs b/backend/src/api/panel_api.rs index d9feebf97..f6fd1cd33 100644 --- a/backend/src/api/panel_api.rs +++ b/backend/src/api/panel_api.rs @@ -1,16 +1,21 @@ -use crate::api::config_file::ConfigFile; -use crate::api::model::{ - create_panel_api_provisioning_stream_with_stop, create_provider_connections_exhausted_stream, - AppState, StreamDetails, -}; -use crate::model::{is_input_expired, ConfigInput, ConfigInputAlias, GracePeriodOptions, PanelApiConfig, PanelApiQueryParam, ProxyUserCredentials}; -use crate::repository::{ - csv_patch_batch_append, csv_patch_batch_remove_expired, csv_patch_batch_sort_by_exp_date, - csv_patch_batch_update_credentials, csv_patch_batch_update_exp_date, get_csv_file_path, -}; -use crate::tools::atomic_once_flag::AtomicOnceFlag; -use crate::utils::{ - debug_if_enabled, format_http_status, persist_source_config, read_sources_file_from_path, +use crate::{ + api::{ + config_file::ConfigFile, + model::{ + create_panel_api_provisioning_stream_with_stop, create_provider_connections_exhausted_stream, AppState, + StreamDetails, + }, + }, + model::{ + is_input_expired, ConfigInput, ConfigInputAlias, GracePeriodOptions, PanelApiConfig, PanelApiQueryParam, + ProxyUserCredentials, + }, + repository::{ + csv_patch_batch_append, csv_patch_batch_remove_expired, csv_patch_batch_sort_by_exp_date, + csv_patch_batch_update_credentials, csv_patch_batch_update_exp_date, get_csv_file_path, + }, + tools::atomic_once_flag::AtomicOnceFlag, + utils::{debug_if_enabled, format_http_status, persist_source_config, read_sources_file_from_path}, }; use axum::http::{Method, StatusCode}; use chrono::{NaiveDateTime, TimeZone}; @@ -19,20 +24,26 @@ use jsonwebtoken::get_current_timestamp; use log::{error, warn}; use serde::{Deserialize, Serialize}; use serde_json::Value; -use shared::concat_string; -use shared::create_bitset; -use shared::error::{info_err, info_err_res, TuliproxError}; -use shared::model::{ - ConfigInputAliasDto, InputType, PanelApiAliasPoolSizeValue, - PanelApiProvisioningMethod, ProxyUserStatus, SourcesConfigDto, VirtualId, +use shared::{ + concat_string, create_bitset, + error::{info_err, info_err_res, TuliproxError}, + model::{ + ConfigInputAliasDto, InputType, PanelApiAliasPoolSizeValue, PanelApiProvisioningMethod, ProxyUserStatus, + SourcesConfigDto, VirtualId, + }, + utils::{ + get_base_url_from_str, get_credentials_from_url, get_credentials_from_url_str, get_i64_from_serde_value, + get_string_from_serde_value, parse_timestamp, sanitize_sensitive_info, Internable, + }, +}; +use std::{ + cmp::Ordering, + collections::{HashMap, HashSet}, + net::SocketAddr, + path::{Path, PathBuf}, + sync::Arc, + time::{Duration, Instant}, }; -use shared::utils::{get_base_url_from_str, get_credentials_from_url, get_credentials_from_url_str, get_i64_from_serde_value, get_string_from_serde_value, parse_timestamp, sanitize_sensitive_info, Internable}; -use std::cmp::Ordering; -use std::collections::{HashMap, HashSet}; -use std::net::SocketAddr; -use std::path::{Path, PathBuf}; -use std::sync::Arc; -use std::time::{Duration, Instant}; use url::Url; #[derive(Debug, Clone)] @@ -79,10 +90,7 @@ fn parse_boolish(value: &Value) -> bool { match value { Value::Bool(b) => *b, Value::Number(n) => n.as_i64().unwrap_or(0) != 0, - Value::String(s) => matches!( - s.trim().to_lowercase().as_str(), - "true" | "1" | "yes" | "y" | "ok" - ), + Value::String(s) => matches!(s.trim().to_lowercase().as_str(), "true" | "1" | "yes" | "y" | "ok"), _ => false, } } @@ -140,11 +148,7 @@ fn parse_panel_expire_with_tz(value: &str, tz: Tz) -> Option { if let Ok(ts) = value.parse::() { return Some(ts); } - let value = if is_date_only_yyyy_mm_dd(value) { - format!("{value} 00:00:00") - } else { - value.to_string() - }; + let value = if is_date_only_yyyy_mm_dd(value) { format!("{value} 00:00:00") } else { value.to_string() }; let dt = NaiveDateTime::parse_from_str(&value, "%Y-%m-%d %H:%M:%S").ok()?; match tz.from_local_datetime(&dt) { chrono::LocalResult::Single(local_dt) => Some(local_dt.timestamp()), @@ -162,10 +166,9 @@ fn normalize_panel_expire(value: &str, ctx: Option<&PanelApiTimeContext>) -> Opt } match ctx.expire_mode { PanelApiExpireMode::UtcString => parse_panel_expire_utc(value), - PanelApiExpireMode::ServerTzString => ctx - .server_tz - .and_then(|tz| parse_panel_expire_with_tz(value, tz)) - .or_else(|| parse_panel_expire_utc(value)), + PanelApiExpireMode::ServerTzString => { + ctx.server_tz.and_then(|tz| parse_panel_expire_with_tz(value, tz)).or_else(|| parse_panel_expire_utc(value)) + } } } @@ -173,8 +176,7 @@ fn is_input_expired_at(exp_date: Option, now: u64) -> bool { let Some(exp_date) = exp_date else { return false; }; - u64::try_from(exp_date) - .map_or(true, |exp_ts| exp_ts <= now) + u64::try_from(exp_date).map_or(true, |exp_ts| exp_ts <= now) } fn is_expiring_with_offset_at(exp_date: Option, offset_secs: u64, now: u64) -> bool { @@ -198,19 +200,9 @@ fn first_json_object(value: &Value) -> Option<&serde_json::Map> { } } -fn extract_username_password_from_json( - obj: &serde_json::Map, -) -> Option<(String, String)> { - let username = obj - .get("username") - .and_then(|v| v.as_str()) - .map(str::trim) - .filter(|s| !s.is_empty()); - let password = obj - .get("password") - .and_then(|v| v.as_str()) - .map(str::trim) - .filter(|s| !s.is_empty()); +fn extract_username_password_from_json(obj: &serde_json::Map) -> Option<(String, String)> { + let username = obj.get("username").and_then(|v| v.as_str()).map(str::trim).filter(|s| !s.is_empty()); + let password = obj.get("password").and_then(|v| v.as_str()).map(str::trim).filter(|s| !s.is_empty()); match (username, password) { (Some(u), Some(p)) => Some((u.to_string(), p.to_string())), _ => None, @@ -218,10 +210,7 @@ fn extract_username_password_from_json( } fn validate_type_is_m3u(params: &[PanelApiQueryParam]) -> Result<(), TuliproxError> { - let typ = params - .iter() - .find(|p| p.key.trim().eq_ignore_ascii_case("type")) - .map(|p| p.value.trim().to_string()); + let typ = params.iter().find(|p| p.key.trim().eq_ignore_ascii_case("type")).map(|p| p.value.trim().to_string()); let typ_str = typ.as_deref().unwrap_or_default(); if typ_str.trim().eq_ignore_ascii_case("m3u") { Ok(()) @@ -233,32 +222,19 @@ fn validate_type_is_m3u(params: &[PanelApiQueryParam]) -> Result<(), TuliproxErr } fn require_api_key_param(params: &[PanelApiQueryParam], section: &str) -> Result<(), TuliproxError> { - let api_key = params - .iter() - .find(|p| p.key.trim().eq_ignore_ascii_case("api_key")); + let api_key = params.iter().find(|p| p.key.trim().eq_ignore_ascii_case("api_key")); let Some(api_key) = api_key else { - return info_err_res!( - "panel_api: {section} must contain query param 'api_key' (use value 'auto')" - ); + return info_err_res!("panel_api: {section} must contain query param 'api_key' (use value 'auto')"); }; if api_key.value.trim().is_empty() { - return info_err_res!( - "panel_api: {section} query param 'api_key' must not be empty (use value 'auto')" - ); + return info_err_res!("panel_api: {section} query param 'api_key' must not be empty (use value 'auto')"); } Ok(()) } -fn require_username_password_params_auto( - params: &[PanelApiQueryParam], - section: &str, -) -> Result<(), TuliproxError> { - let username = params - .iter() - .find(|p| p.key.trim().eq_ignore_ascii_case("username")); - let password = params - .iter() - .find(|p| p.key.trim().eq_ignore_ascii_case("password")); +fn require_username_password_params_auto(params: &[PanelApiQueryParam], section: &str) -> Result<(), TuliproxError> { + let username = params.iter().find(|p| p.key.trim().eq_ignore_ascii_case("username")); + let password = params.iter().find(|p| p.key.trim().eq_ignore_ascii_case("password")); if username.is_none() || password.is_none() { return info_err_res!( "panel_api: {section} must contain query params 'username' and 'password' (use value 'auto')" @@ -277,10 +253,7 @@ fn require_username_password_params_auto( fn validate_client_new_params(params: &[PanelApiQueryParam]) -> Result<(), TuliproxError> { require_api_key_param(params, "query_parameter.client_new")?; validate_type_is_m3u(params)?; - if params - .iter() - .any(|p| p.key.trim().eq_ignore_ascii_case("user")) - { + if params.iter().any(|p| p.key.trim().eq_ignore_ascii_case("user")) { return info_err_res!("panel_api: client_new must not contain query param 'user'"); } Ok(()) @@ -301,47 +274,27 @@ fn validate_client_info_params(params: &[PanelApiQueryParam]) -> Result<(), Tuli fn validate_account_info_params(params: &[PanelApiQueryParam]) -> Result<(), TuliproxError> { require_api_key_param(params, "query_parameter.account_info")?; - let has_user = params - .iter() - .any(|p| p.key.trim().eq_ignore_ascii_case("username")); - let has_pass = params - .iter() - .any(|p| p.key.trim().eq_ignore_ascii_case("password")); + let has_user = params.iter().any(|p| p.key.trim().eq_ignore_ascii_case("username")); + let has_pass = params.iter().any(|p| p.key.trim().eq_ignore_ascii_case("password")); if has_user || has_pass { require_username_password_params_auto(params, "query_parameter.account_info")?; } Ok(()) } -fn validate_client_adult_content_params( - params: &[PanelApiQueryParam], -) -> Result<(), TuliproxError> { +fn validate_client_adult_content_params(params: &[PanelApiQueryParam]) -> Result<(), TuliproxError> { require_api_key_param(params, "query_parameter.client_adult_content")?; - let has_user = params - .iter() - .any(|p| p.key.trim().eq_ignore_ascii_case("username")); - let has_pass = params - .iter() - .any(|p| p.key.trim().eq_ignore_ascii_case("password")); + let has_user = params.iter().any(|p| p.key.trim().eq_ignore_ascii_case("username")); + let has_pass = params.iter().any(|p| p.key.trim().eq_ignore_ascii_case("password")); if has_user || has_pass { require_username_password_params_auto(params, "query_parameter.client_adult_content")?; } Ok(()) } -create_bitset!( - u8, - PanelApiOptionalFlags, - AccountInfo, - ClientNew, - ClientRenew, - AdultContent -); +create_bitset!(u8, PanelApiOptionalFlags, AccountInfo, ClientNew, ClientRenew, AdultContent); -fn resolve_panel_api_optional_flags( - cfg: &PanelApiConfig, - input_name: &str, -) -> PanelApiOptionalFlagsSet { +fn resolve_panel_api_optional_flags(cfg: &PanelApiConfig, input_name: &str) -> PanelApiOptionalFlagsSet { let mut flags = PanelApiOptionalFlagsSet::new(); if !cfg.query_parameter.account_info.is_empty() { flags.set(PanelApiOptionalFlags::AccountInfo); @@ -358,28 +311,16 @@ fn resolve_panel_api_optional_flags( let name = sanitize_sensitive_info(input_name); if !flags.contains(PanelApiOptionalFlags::ClientRenew) { - debug_if_enabled!( - "panel_api request for client_renew disabled due to missing arguments for {}", - name - ); + debug_if_enabled!("panel_api request for client_renew disabled due to missing arguments for {}", name); } if !flags.contains(PanelApiOptionalFlags::ClientNew) { - debug_if_enabled!( - "panel_api request for client_new disabled due to missing arguments for {}", - name - ); + debug_if_enabled!("panel_api request for client_new disabled due to missing arguments for {}", name); } if !flags.contains(PanelApiOptionalFlags::AdultContent) { - debug_if_enabled!( - "panel_api request for client_adult_content disabled due to missing arguments for {}", - name - ); + debug_if_enabled!("panel_api request for client_adult_content disabled due to missing arguments for {}", name); } if !flags.contains(PanelApiOptionalFlags::AccountInfo) { - debug_if_enabled!( - "panel_api request for account_info disabled due to missing arguments for {}", - name - ); + debug_if_enabled!("panel_api request for account_info disabled due to missing arguments for {}", name); } flags } @@ -410,12 +351,9 @@ fn parse_panel_api_provisioning_offset_secs(offset: &str) -> Result Result<(), TuliproxError> { @@ -459,9 +397,7 @@ fn validate_panel_api_config(cfg: &PanelApiConfig) -> Result<(), TuliproxError> let max = max_val.and_then(PanelApiAliasPoolSizeValue::as_number); if let (Some(min), Some(max)) = (min, max) { if min > max { - return info_err_res!( - "panel_api.alias_pool.size.min must be <= panel_api.alias_pool.size.max" - ); + return info_err_res!("panel_api.alias_pool.size.min must be <= panel_api.alias_pool.size.max"); } } if cfg.provisioning.probe_interval_sec == 0 { @@ -488,9 +424,7 @@ fn resolve_query_params( if value.eq_ignore_ascii_case("auto") { if key.eq_ignore_ascii_case("api_key") { let Some(k) = api_key.filter(|s| !s.trim().is_empty()) else { - return info_err_res!( - "panel_api: query param {key} uses 'auto' but panel_api.api_key is missing" - ); + return info_err_res!("panel_api: query param {key} uses 'auto' but panel_api.api_key is missing"); }; value = k.to_string(); } else if key.eq_ignore_ascii_case("username") { @@ -514,12 +448,8 @@ fn resolve_query_params( Ok(out) } -fn build_panel_url( - base_url: &str, - query_params: &[(String, String)], -) -> Result { - let mut url = - Url::parse(base_url).map_err(|e| info_err!("panel_api: invalid url {base_url}: {e}"))?; +fn build_panel_url(base_url: &str, query_params: &[(String, String)]) -> Result { + let mut url = Url::parse(base_url).map_err(|e| info_err!("panel_api: invalid url {base_url}: {e}"))?; { let mut pairs = url.query_pairs_mut(); for (k, v) in query_params { @@ -531,11 +461,9 @@ fn build_panel_url( fn sanitize_panel_api_json_for_log(value: &Value, sanitize_sensitive: bool) -> Value { match value { - Value::Array(arr) => Value::Array( - arr.iter() - .map(|v| sanitize_panel_api_json_for_log(v, sanitize_sensitive)) - .collect(), - ), + Value::Array(arr) => { + Value::Array(arr.iter().map(|v| sanitize_panel_api_json_for_log(v, sanitize_sensitive)).collect()) + } Value::Object(obj) => { let mut out = serde_json::Map::with_capacity(obj.len()); for (k, v) in obj { @@ -554,17 +482,11 @@ fn sanitize_panel_api_json_for_log(value: &Value, sanitize_sensitive: bool) -> V } if k.eq_ignore_ascii_case("url") { if let Some(s) = v.as_str() { - out.insert( - k.clone(), - Value::String(sanitize_sensitive_info(s).into_owned()), - ); + out.insert(k.clone(), Value::String(sanitize_sensitive_info(s).into_owned())); continue; } } - out.insert( - k.clone(), - sanitize_panel_api_json_for_log(v, sanitize_sensitive), - ); + out.insert(k.clone(), sanitize_panel_api_json_for_log(v, sanitize_sensitive)); } Value::Object(out) } @@ -583,19 +505,10 @@ async fn panel_get_json(app_state: &AppState, url: Url) -> Result Result Result<(String, String, Option), TuliproxError> { validate_client_new_params(&cfg.query_parameter.client_new)?; - let params = resolve_query_params( - &cfg.query_parameter.client_new, - cfg.api_key.as_deref(), - None, - )?; + let params = resolve_query_params(&cfg.query_parameter.client_new, cfg.api_key.as_deref(), None)?; let url = build_panel_url(cfg.url.as_ref(), ¶ms)?; let json = panel_get_json(app_state, url).await?; let Some(obj) = first_json_object(&json) else { @@ -680,11 +580,8 @@ async fn panel_client_renew( password: &str, ) -> Result<(), TuliproxError> { validate_client_renew_params(&cfg.query_parameter.client_renew)?; - let params = resolve_query_params( - &cfg.query_parameter.client_renew, - cfg.api_key.as_deref(), - Some((username, password)), - )?; + let params = + resolve_query_params(&cfg.query_parameter.client_renew, cfg.api_key.as_deref(), Some((username, password)))?; let url = build_panel_url(cfg.url.as_ref(), ¶ms)?; let json = panel_get_json(app_state, url).await?; let Some(obj) = first_json_object(&json) else { @@ -704,11 +601,8 @@ async fn panel_client_info_raw( password: &str, ) -> Result, TuliproxError> { validate_client_info_params(&cfg.query_parameter.client_info)?; - let params = resolve_query_params( - &cfg.query_parameter.client_info, - cfg.api_key.as_deref(), - Some((username, password)), - )?; + let params = + resolve_query_params(&cfg.query_parameter.client_info, cfg.api_key.as_deref(), Some((username, password)))?; let url = build_panel_url(cfg.url.as_ref(), ¶ms)?; let json = panel_get_json(app_state, url).await?; let Some(obj) = first_json_object(&json) else { @@ -718,12 +612,7 @@ async fn panel_client_info_raw( if !status_ok { return info_err_res!("panel_api: client_info status=false"); } - let expire = obj - .get("expire") - .and_then(|v| v.as_str()) - .unwrap_or_default() - .trim() - .to_string(); + let expire = obj.get("expire").and_then(|v| v.as_str()).unwrap_or_default().trim().to_string(); if expire.is_empty() { Ok(None) } else { @@ -739,22 +628,18 @@ async fn panel_client_info( time_ctx: Option<&PanelApiTimeContext>, ) -> Result, TuliproxError> { let expire = panel_client_info_raw(app_state, cfg, username, password).await?; - Ok(expire - .as_deref() - .and_then(|value| normalize_panel_expire(value, time_ctx))) + Ok(expire.as_deref().and_then(|value| normalize_panel_expire(value, time_ctx))) } async fn fetch_root_user_api_info( app_state: &Arc, input: &ConfigInput, ) -> Result, TuliproxError> { - - let Some((username, password)) = extract_account_creds_from_input(input) else { - return Ok(None); + return Ok(None); }; - let resolved_url = input.resolve()?; + let resolved_url = input.resolve()?; let base_url = get_base_url_from_str(&resolved_url).unwrap_or_else(|| resolved_url.to_string()); let Ok(mut url) = Url::parse(base_url.as_str()) else { @@ -780,25 +665,15 @@ async fn fetch_root_user_api_info( } }; - let exp_date = json - .get("user_info") - .and_then(|v| v.get("exp_date")) - .and_then(get_i64_from_serde_value); - let server_now_ts = json - .get("server_info") - .and_then(|v| v.get("timestamp_now")) - .and_then(get_i64_from_serde_value); + let exp_date = json.get("user_info").and_then(|v| v.get("exp_date")).and_then(get_i64_from_serde_value); + let server_now_ts = json.get("server_info").and_then(|v| v.get("timestamp_now")).and_then(get_i64_from_serde_value); let server_tz = json .get("server_info") .and_then(|v| v.get("timezone")) .and_then(get_string_from_serde_value) .and_then(|tz| tz.parse::().ok()); - Ok(Some(UserApiAccountInfo { - exp_date, - server_now_ts, - server_tz, - })) + Ok(Some(UserApiAccountInfo { exp_date, server_now_ts, server_tz })) } fn resolve_panel_expire_mode( @@ -815,9 +690,7 @@ fn resolve_panel_expire_mode( let utc_ts = parse_panel_expire_utc(panel_expire); let tz_ts = server_tz.and_then(|tz| parse_panel_expire_with_tz(panel_expire, tz)); let Some(utc_ts) = utc_ts else { - return tz_ts.map_or(PanelApiExpireMode::UtcString, |_| { - PanelApiExpireMode::ServerTzString - }); + return tz_ts.map_or(PanelApiExpireMode::UtcString, |_| PanelApiExpireMode::ServerTzString); }; let Some(tz_ts) = tz_ts else { return PanelApiExpireMode::UtcString; @@ -851,9 +724,7 @@ fn apply_clock_skew(now: u64, skew_secs: i64) -> u64 { fn panel_api_time_cache_path(app_state: &AppState) -> PathBuf { let paths = app_state.app_config.paths.load(); - PathBuf::from(&paths.config_path) - .join("panel-api") - .join("panel_api_time_cache.json") + PathBuf::from(&paths.config_path).join("panel-api").join("panel_api_time_cache.json") } async fn load_panel_api_time_cache(app_state: &AppState, cache_path: &Path) -> PanelApiTimeCache { @@ -875,9 +746,7 @@ async fn persist_panel_api_time_cache( cache: &PanelApiTimeCache, ) -> Result<(), TuliproxError> { if let Some(parent) = cache_path.parent() { - tokio::fs::create_dir_all(parent) - .await - .map_err(|e| info_err!("panel_api: failed to create cache dir: {e}"))?; + tokio::fs::create_dir_all(parent).await.map_err(|e| info_err!("panel_api: failed to create cache dir: {e}"))?; } let content = serde_json::to_string_pretty(cache).map_err(|e| info_err!("panel_api: {e}"))?; let _lock = app_state.app_config.file_locks.write_lock(cache_path).await; @@ -887,9 +756,7 @@ async fn persist_panel_api_time_cache( Ok(()) } -fn parse_cached_tz(tz: Option) -> Option { - tz.and_then(|name| name.parse::().ok()) -} +fn parse_cached_tz(tz: Option) -> Option { tz.and_then(|name| name.parse::().ok()) } async fn panel_account_info( app_state: &AppState, @@ -900,11 +767,7 @@ async fn panel_account_info( return Ok(None); } validate_account_info_params(&cfg.query_parameter.account_info)?; - let params = resolve_query_params( - &cfg.query_parameter.account_info, - cfg.api_key.as_deref(), - creds, - )?; + let params = resolve_query_params(&cfg.query_parameter.account_info, cfg.api_key.as_deref(), creds)?; let url = build_panel_url(cfg.url.as_ref(), ¶ms)?; let json = panel_get_json(app_state, url).await?; let Some(obj) = first_json_object(&json) else { @@ -929,17 +792,11 @@ async fn panel_client_adult_content( return Ok(()); } validate_client_adult_content_params(&cfg.query_parameter.client_adult_content)?; - let params = resolve_query_params( - &cfg.query_parameter.client_adult_content, - cfg.api_key.as_deref(), - creds, - )?; + let params = resolve_query_params(&cfg.query_parameter.client_adult_content, cfg.api_key.as_deref(), creds)?; let url = build_panel_url(cfg.url.as_ref(), ¶ms)?; let json = panel_get_json(app_state, url).await?; let Some(obj) = first_json_object(&json) else { - return info_err_res!( - "panel_api: client_adult_content response is not a JSON object/array" - ); + return info_err_res!("panel_api: client_adult_content response is not a JSON object/array"); }; let status_ok = obj.get("status").is_some_and(parse_boolish); if !status_ok { @@ -957,9 +814,7 @@ fn extract_account_creds_from_input(input: &ConfigInput) -> Option<(String, Stri Url::parse(input.url.as_str()).ok().and_then(|u| { let (uu, pp) = get_credentials_from_url(&u); match (uu, pp) { - (Some(uu), Some(pp)) if !uu.trim().is_empty() && !pp.trim().is_empty() => { - Some((uu, pp)) - } + (Some(uu), Some(pp)) if !uu.trim().is_empty() && !pp.trim().is_empty() => Some((uu, pp)), _ => None, } }) @@ -967,10 +822,7 @@ fn extract_account_creds_from_input(input: &ConfigInput) -> Option<(String, Stri fn alias_pool_limit_values( cfg: &PanelApiConfig, -) -> ( - Option<&PanelApiAliasPoolSizeValue>, - Option<&PanelApiAliasPoolSizeValue>, -) { +) -> (Option<&PanelApiAliasPoolSizeValue>, Option<&PanelApiAliasPoolSizeValue>) { let size = cfg.alias_pool.as_ref().and_then(|p| p.size.as_ref()); let min = size.and_then(|s| s.min.as_ref()); let max = size.and_then(|s| s.max.as_ref()); @@ -982,10 +834,7 @@ fn alias_pool_has_min(cfg: &PanelApiConfig) -> bool { min.is_some() } -fn resolve_alias_pool_limit_value( - value: Option<&PanelApiAliasPoolSizeValue>, - auto_value: Option, -) -> Option { +fn resolve_alias_pool_limit_value(value: Option<&PanelApiAliasPoolSizeValue>, auto_value: Option) -> Option { match value { Some(PanelApiAliasPoolSizeValue::Number(v)) => Some(*v), Some(PanelApiAliasPoolSizeValue::Auto(_)) => auto_value, @@ -1024,18 +873,8 @@ fn count_enabled_proxy_users(app_state: &AppState, input_name: &Arc) -> usi api_proxy .user .iter() - .filter(|target_user| { - target_names - .iter() - .any(|target| target.eq_ignore_ascii_case(&target_user.target)) - }) - .map(|target_user| { - target_user - .credentials - .iter() - .filter(|cred| is_proxy_user_enabled(cred)) - .count() - }) + .filter(|target_user| target_names.iter().any(|target| target.eq_ignore_ascii_case(&target_user.target))) + .map(|target_user| target_user.credentials.iter().filter(|cred| is_proxy_user_enabled(cred)).count()) .sum() } @@ -1047,10 +886,7 @@ fn resolve_alias_pool_auto_value(app_state: &AppState, input_name: &Arc) -> pub(crate) fn target_has_alias_pool_min(app_state: &AppState, target_name: &str) -> bool { let sources = app_state.app_config.sources.load(); for source in &sources.sources { - let target_match = source - .targets - .iter() - .any(|target| target.name.eq_ignore_ascii_case(target_name)); + let target_match = source.targets.iter().any(|target| target.name.eq_ignore_ascii_case(target_name)); if !target_match { continue; } @@ -1089,40 +925,25 @@ fn resolve_alias_pool_limits( }; if let (Some(min), Some(max)) = (min, max) { if min > max { - return info_err_res!( - "panel_api.alias_pool.size.min must be <= panel_api.alias_pool.size.max" - ); + return info_err_res!("panel_api.alias_pool.size.min must be <= panel_api.alias_pool.size.max"); } } Ok((min, max)) } -fn resolve_alias_pool_min( - app_state: &AppState, - input_name: &Arc, - cfg: &PanelApiConfig, -) -> Option { +fn resolve_alias_pool_min(app_state: &AppState, input_name: &Arc, cfg: &PanelApiConfig) -> Option { let (min_val, _) = alias_pool_limit_values(cfg); let min_val = min_val?; - let auto_value = min_val - .is_auto() - .then(|| resolve_alias_pool_auto_value(app_state, input_name)); + let auto_value = min_val.is_auto().then(|| resolve_alias_pool_auto_value(app_state, input_name)); resolve_alias_pool_limit_value(Some(min_val), auto_value) } -fn alias_pool_remove_expired(cfg: &PanelApiConfig) -> bool { - cfg.alias_pool.as_ref().is_some_and(|p| p.remove_expired) -} +fn alias_pool_remove_expired(cfg: &PanelApiConfig) -> bool { cfg.alias_pool.as_ref().is_some_and(|p| p.remove_expired) } fn collect_accounts(input: &ConfigInput) -> Vec { let mut out = Vec::new(); if let Some((u, p)) = extract_account_creds_from_input(input) { - out.push(AccountCredentials { - name: input.name.clone(), - username: u, - password: p, - exp_date: input.exp_date, - }); + out.push(AccountCredentials { name: input.name.clone(), username: u, password: p, exp_date: input.exp_date }); } if let Some(aliases) = input.aliases.as_ref() { for a in aliases { @@ -1163,9 +984,7 @@ fn aliases_need_sort_config(aliases: &[ConfigInputAlias]) -> bool { if aliases.len() < 2 { return false; } - aliases - .windows(2) - .any(|pair| compare_alias_exp_date_config(&pair[0], &pair[1]) == Ordering::Greater) + aliases.windows(2).any(|pair| compare_alias_exp_date_config(&pair[0], &pair[1]) == Ordering::Greater) } fn sort_aliases_by_exp_date(aliases: &mut Vec) -> bool { @@ -1184,11 +1003,8 @@ fn sort_aliases_by_exp_date(aliases: &mut Vec) -> bool { fn sort_account_aliases_keep_root_first(accounts: &mut Vec, root_name: &str) { let root = accounts.iter().find(|acct| acct.name.as_ref() == root_name).cloned(); - let mut aliases: Vec = accounts - .iter() - .filter(|acct| acct.name.as_ref() != root_name) - .cloned() - .collect(); + let mut aliases: Vec = + accounts.iter().filter(|acct| acct.name.as_ref() != root_name).cloned().collect(); aliases.sort_by(compare_account_exp_date); accounts.clear(); if let Some(root) = root { @@ -1197,51 +1013,28 @@ fn sort_account_aliases_keep_root_first(accounts: &mut Vec, accounts.extend(aliases); } -fn is_account_valid(exp_date: Option) -> bool { - exp_date.is_some() && !is_input_expired(exp_date) -} +fn is_account_valid(exp_date: Option) -> bool { exp_date.is_some() && !is_input_expired(exp_date) } fn count_valid_accounts(accounts: &[AccountCredentials]) -> usize { - accounts - .iter() - .filter(|acct| is_account_valid(acct.exp_date)) - .count() + accounts.iter().filter(|acct| is_account_valid(acct.exp_date)).count() } fn root_counts_towards_pool(accounts: &[AccountCredentials], input_name: &Arc) -> bool { - accounts - .iter() - .find(|acct| &acct.name == input_name) - .is_some_and(|acct| is_account_valid(acct.exp_date)) + accounts.iter().find(|acct| &acct.name == input_name).is_some_and(|acct| is_account_valid(acct.exp_date)) } fn count_valid_accounts_at(accounts: &[AccountCredentials], now: u64) -> usize { + accounts.iter().filter(|acct| acct.exp_date.is_some() && !is_input_expired_at(acct.exp_date, now)).count() +} + +fn count_valid_alias_accounts_at(accounts: &[AccountCredentials], input_name: &Arc, now: u64) -> usize { accounts .iter() - .filter(|acct| acct.exp_date.is_some() && !is_input_expired_at(acct.exp_date, now)) + .filter(|acct| &acct.name != input_name && acct.exp_date.is_some() && !is_input_expired_at(acct.exp_date, now)) .count() } -fn count_valid_alias_accounts_at( - accounts: &[AccountCredentials], - input_name: &Arc, - now: u64, -) -> usize { - accounts - .iter() - .filter(|acct| { - &acct.name != input_name - && acct.exp_date.is_some() - && !is_input_expired_at(acct.exp_date, now) - }) - .count() -} - -fn root_counts_towards_pool_at( - accounts: &[AccountCredentials], - input_name: &Arc, - now: u64, -) -> bool { +fn root_counts_towards_pool_at(accounts: &[AccountCredentials], input_name: &Arc, now: u64) -> bool { accounts .iter() .find(|acct| &acct.name == input_name) @@ -1294,10 +1087,7 @@ pub(crate) fn can_provision_on_exhausted(app_state: &AppState, input: &ConfigInp return false; } if let Err(err) = validate_panel_api_config(panel_cfg) { - debug_if_enabled!( - "panel_api config invalid: {}", - sanitize_sensitive_info(err.to_string().as_str()) - ); + debug_if_enabled!("panel_api config invalid: {}", sanitize_sensitive_info(err.to_string().as_str())); return false; } if is_alias_pool_max_reached(app_state, input) { @@ -1306,20 +1096,13 @@ pub(crate) fn can_provision_on_exhausted(app_state: &AppState, input: &ConfigInp true } -pub(crate) fn find_input_by_provider_name( - app_state: &AppState, - provider_name: &str, -) -> Option> { +pub(crate) fn find_input_by_provider_name(app_state: &AppState, provider_name: &str) -> Option> { let sources = app_state.app_config.sources.load(); for input in &sources.inputs { if &*input.name == provider_name { return Some(Arc::clone(input)); } - if input - .aliases - .as_ref() - .is_some_and(|aliases| aliases.iter().any(|alias| &*alias.name == provider_name)) - { + if input.aliases.as_ref().is_some_and(|aliases| aliases.iter().any(|alias| &*alias.name == provider_name)) { return Some(Arc::clone(input)); } } @@ -1364,7 +1147,7 @@ async fn patch_source_yml_add_alias( enabled: true, }; - input.upsert_alias(alias)?; + input.upsert_alias(alias)?; persist_source_config(app_state, Some(source_file_path), sources).await?; Ok(()) @@ -1414,10 +1197,7 @@ fn update_url_query_credentials_if_present(url: &mut String, username: &str, pas let Ok(mut parsed) = Url::parse(url.as_str()) else { return; }; - let mut pairs: Vec<(String, String)> = parsed - .query_pairs() - .map(|(k, v)| (k.to_string(), v.to_string())) - .collect(); + let mut pairs: Vec<(String, String)> = parsed.query_pairs().map(|(k, v)| (k.to_string(), v.to_string())).collect(); let mut has_user = false; let mut has_pass = false; for (k, v) in &mut pairs { @@ -1448,10 +1228,7 @@ fn update_url_query_credentials_if_present(url: &mut String, username: &str, pas } #[allow(clippy::too_many_lines)] -fn apply_sources_yml_patches( - doc: &mut SourcesConfigDto, - patches: &[SourcesYmlPatch], -) -> Result { +fn apply_sources_yml_patches(doc: &mut SourcesConfigDto, patches: &[SourcesYmlPatch]) -> Result { if patches.is_empty() { return Ok(false); } @@ -1477,13 +1254,10 @@ fn apply_sources_yml_patches( for patch in patches { match patch { - SourcesYmlPatch::UpdatePanelApiCredits { - input_name, - credits, - } => { - let idx = *inputs_by_name.get(input_name.as_ref()).ok_or_else(|| { - info_err!("panel_api: could not find input '{input_name}' in source.yml") - })?; + SourcesYmlPatch::UpdatePanelApiCredits { input_name, credits } => { + let idx = *inputs_by_name + .get(input_name.as_ref()) + .ok_or_else(|| info_err!("panel_api: could not find input '{input_name}' in source.yml"))?; let Some(panel_api) = doc.inputs[idx].panel_api.as_mut() else { return Err(info_err!( "panel_api: could not find panel_api for input '{input_name}' in source.yml" @@ -1495,29 +1269,22 @@ fn apply_sources_yml_patches( } } SourcesYmlPatch::SortAliases { input_name } => { - let idx = *inputs_by_name.get(input_name.as_ref()).ok_or_else(|| { - info_err!("panel_api: could not find input '{input_name}' in source.yml") - })?; + let idx = *inputs_by_name + .get(input_name.as_ref()) + .ok_or_else(|| info_err!("panel_api: could not find input '{input_name}' in source.yml"))?; let Some(aliases) = doc.inputs[idx].aliases.as_mut() else { continue; }; if sort_aliases_by_exp_date(aliases) { - alias_indices[idx] = aliases - .iter() - .enumerate() - .map(|(idx, alias)| (alias.name.clone(), idx)) - .collect(); + alias_indices[idx] = + aliases.iter().enumerate().map(|(idx, alias)| (alias.name.clone(), idx)).collect(); changed = true; } } - SourcesYmlPatch::UpdateExpDate { - input_name, - account_name, - exp_date, - } => { - let idx = *inputs_by_name.get(input_name.as_ref()).ok_or_else(|| { - info_err!("panel_api: could not find input '{input_name}' in source.yml") - })?; + SourcesYmlPatch::UpdateExpDate { input_name, account_name, exp_date } => { + let idx = *inputs_by_name + .get(input_name.as_ref()) + .ok_or_else(|| info_err!("panel_api: could not find input '{input_name}' in source.yml"))?; if account_name == input_name { if doc.inputs[idx].exp_date != Some(*exp_date) { doc.inputs[idx].exp_date = Some(*exp_date); @@ -1542,15 +1309,10 @@ fn apply_sources_yml_patches( changed = true; } } - SourcesYmlPatch::UpdateRootCredentials { - input_name, - username, - password, - exp_date, - } => { - let idx = *inputs_by_name.get(input_name.as_ref()).ok_or_else(|| { - info_err!("panel_api: could not find input '{input_name}' in source.yml") - })?; + SourcesYmlPatch::UpdateRootCredentials { input_name, username, password, exp_date } => { + let idx = *inputs_by_name + .get(input_name.as_ref()) + .ok_or_else(|| info_err!("panel_api: could not find input '{input_name}' in source.yml"))?; let input = &mut doc.inputs[idx]; let exp_date_changed = exp_date.is_some() && input.exp_date != *exp_date; if input.username.as_deref() != Some(username.as_str()) @@ -1568,16 +1330,10 @@ fn apply_sources_yml_patches( changed = true; } } - SourcesYmlPatch::UpdateAliasCredentials { - input_name, - alias_name, - username, - password, - exp_date, - } => { - let idx = *inputs_by_name.get(input_name.as_ref()).ok_or_else(|| { - info_err!("panel_api: could not find input '{input_name}' in source.yml") - })?; + SourcesYmlPatch::UpdateAliasCredentials { input_name, alias_name, username, password, exp_date } => { + let idx = *inputs_by_name + .get(input_name.as_ref()) + .ok_or_else(|| info_err!("panel_api: could not find input '{input_name}' in source.yml"))?; let Some(alias_idx) = alias_indices[idx].get(alias_name).copied() else { return Err(info_err!( "panel_api: could not find alias '{alias_name}' under input '{input_name}' in source.yml" @@ -1603,23 +1359,15 @@ fn apply_sources_yml_patches( changed = true; } } - SourcesYmlPatch::AddAlias { - input_name, - alias_name, - base_url, - username, - password, - exp_date, - } => { - let idx = *inputs_by_name.get(input_name).ok_or_else(|| { - info_err!("panel_api: could not find input '{input_name}' in source.yml") - })?; + SourcesYmlPatch::AddAlias { input_name, alias_name, base_url, username, password, exp_date } => { + let idx = *inputs_by_name + .get(input_name) + .ok_or_else(|| info_err!("panel_api: could not find input '{input_name}' in source.yml"))?; let input_type = doc.inputs[idx].input_type; let aliases = doc.inputs[idx].aliases.get_or_insert_with(Vec::new); - let next_index = u16::try_from(aliases.len()).map_err(|_| { - info_err!("panel_api: cannot add alias for '{input_name}': alias id overflow") - })?; + let next_index = u16::try_from(aliases.len()) + .map_err(|_| info_err!("panel_api: cannot add alias for '{input_name}': alias id overflow"))?; let mut alias = ConfigInputAliasDto { id: 0, @@ -1639,20 +1387,17 @@ fn apply_sources_yml_patches( changed = true; } SourcesYmlPatch::RemoveExpiredAliases { input_name } => { - let idx = *inputs_by_name.get(input_name).ok_or_else(|| { - info_err!("panel_api: could not find input '{input_name}' in source.yml") - })?; + let idx = *inputs_by_name + .get(input_name) + .ok_or_else(|| info_err!("panel_api: could not find input '{input_name}' in source.yml"))?; let Some(aliases) = doc.inputs[idx].aliases.as_mut() else { continue; }; let before = aliases.len(); aliases.retain(|a| !is_input_expired(a.exp_date)); if aliases.len() != before { - alias_indices[idx] = aliases - .iter() - .enumerate() - .map(|(idx, alias)| (alias.name.clone(), idx)) - .collect(); + alias_indices[idx] = + aliases.iter().enumerate().map(|(idx, alias)| (alias.name.clone(), idx)).collect(); changed = true; } } @@ -1737,11 +1482,7 @@ fn derive_unique_alias_name(existing: &[Arc], input_name: &Arc, userna base } -fn derive_unique_alias_name_set( - existing: &HashSet>, - input_name: &Arc, - username: &str, -) -> String { +fn derive_unique_alias_name_set(existing: &HashSet>, input_name: &Arc, username: &str) -> String { let base = format!("{input_name}-{username}"); if !existing.contains(base.as_str()) { return base; @@ -1798,16 +1539,11 @@ async fn try_renew_expired_account( let mut candidates = collect_accounts(input); for acct in &mut candidates { if treat_missing_exp_date_as_expired && acct.exp_date.is_none() { - acct.exp_date = panel_client_info( - app_state, - panel_cfg, - acct.username.as_str(), - acct.password.as_str(), - None, - ) - .await - .ok() - .flatten(); + acct.exp_date = + panel_client_info(app_state, panel_cfg, acct.username.as_str(), acct.password.as_str(), None) + .await + .ok() + .flatten(); } } candidates.sort_by_key(|a| a.exp_date.unwrap_or(i64::MAX)); @@ -1819,14 +1555,7 @@ async fn try_renew_expired_account( if !expired { continue; } - match panel_client_renew( - app_state, - panel_cfg, - acct.username.as_str(), - acct.password.as_str(), - ) - .await - { + match panel_client_renew(app_state, panel_cfg, acct.username.as_str(), acct.password.as_str()).await { Ok(()) => { if adult_enabled { if let Err(err) = panel_client_adult_content( @@ -1843,23 +1572,17 @@ async fn try_renew_expired_account( ); } } - let refreshed_exp = panel_client_info( - app_state, - panel_cfg, - acct.username.as_str(), - acct.password.as_str(), - None, - ) - .await - .ok() - .flatten(); + let refreshed_exp = + panel_client_info(app_state, panel_cfg, acct.username.as_str(), acct.password.as_str(), None) + .await + .ok() + .flatten(); if let Some(new_exp) = refreshed_exp.or(acct.exp_date) { if is_batch { let batch_url = input.t_batch_url.as_deref().unwrap_or_default(); if let Ok(csv_path) = get_csv_file_path(batch_url) { - let _csv_lock = - app_state.app_config.file_locks.write_lock(&csv_path).await; + let _csv_lock = app_state.app_config.file_locks.write_lock(&csv_path).await; if let Err(err) = csv_patch_batch_update_exp_date( input.input_type, &csv_path, @@ -1870,31 +1593,16 @@ async fn try_renew_expired_account( ) .await { - debug_if_enabled!( - "panel_api failed to persist renew exp_date to csv: {}", - err - ); + debug_if_enabled!("panel_api failed to persist renew exp_date to csv: {}", err); } } } else { - let _src_lock = app_state - .app_config - .file_locks - .write_lock(sources_path) - .await; - if let Err(err) = patch_source_yml_update_exp_date( - app_state, - sources_path, - &input.name, - &acct.name, - new_exp, - ) - .await + let _src_lock = app_state.app_config.file_locks.write_lock(sources_path).await; + if let Err(err) = + patch_source_yml_update_exp_date(app_state, sources_path, &input.name, &acct.name, new_exp) + .await { - debug_if_enabled!( - "panel_api failed to persist renew exp_date to source.yml: {}", - err - ); + debug_if_enabled!("panel_api failed to persist renew exp_date to source.yml: {}", err); } } } @@ -1937,8 +1645,7 @@ async fn try_create_new_account( match panel_client_new(app_state, panel_cfg).await { Ok((username, password, base_url_from_resp)) => { let base_url = base_url_from_resp.unwrap_or_else(|| input.url.clone()); - let base_url = - get_base_url_from_str(base_url.as_str()).unwrap_or_else(|| base_url.clone()); + let base_url = get_base_url_from_str(base_url.as_str()).unwrap_or_else(|| base_url.clone()); let mut existing_names: Vec> = vec![input.name.clone()]; if let Some(aliases) = input.aliases.as_ref() { @@ -1947,10 +1654,7 @@ async fn try_create_new_account( let alias_name = derive_unique_alias_name(&existing_names, &input.name, &username); if adult_enabled { - if let Err(err) = - panel_client_adult_content(app_state, panel_cfg, Some((&username, &password))) - .await - { + if let Err(err) = panel_client_adult_content(app_state, panel_cfg, Some((&username, &password))).await { debug_if_enabled!( "panel_api client_adult_content failed for {}: {}", sanitize_sensitive_info(&alias_name), @@ -1959,10 +1663,7 @@ async fn try_create_new_account( } } - let exp_date = panel_client_info(app_state, panel_cfg, &username, &password, None) - .await - .ok() - .flatten(); + let exp_date = panel_client_info(app_state, panel_cfg, &username, &password, None).await.ok().flatten(); if is_batch { let batch_url = input.t_batch_url.as_deref().unwrap_or_default(); @@ -1999,11 +1700,7 @@ async fn try_create_new_account( } } } else { - let _src_lock = app_state - .app_config - .file_locks - .write_lock(sources_path) - .await; + let _src_lock = app_state.app_config.file_locks.write_lock(sources_path).await; if let Err(err) = patch_source_yml_add_alias( app_state, sources_path, @@ -2030,10 +1727,7 @@ async fn try_create_new_account( Some(PanelApiProvisionOutcome::Created { username, password }) } Err(err) => { - debug_if_enabled!( - "panel_api client_new failed: {}", - sanitize_sensitive_info(err.to_string().as_str()) - ); + debug_if_enabled!("panel_api client_new failed: {}", sanitize_sensitive_info(err.to_string().as_str())); None } } @@ -2065,17 +1759,11 @@ pub async fn try_provision_account_on_exhausted( return None; } - let _input_lock = app_state - .app_config - .file_locks - .write_lock_str(format!("panel_api:{}", input.name).as_str()) - .await; + let _input_lock = + app_state.app_config.file_locks.write_lock_str(format!("panel_api:{}", input.name).as_str()).await; if let Err(err) = validate_panel_api_config(panel_cfg) { - debug_if_enabled!( - "panel_api config invalid: {}", - sanitize_sensitive_info(err.to_string().as_str()) - ); + debug_if_enabled!("panel_api config invalid: {}", sanitize_sensitive_info(err.to_string().as_str())); return None; } if is_alias_pool_max_reached(app_state, input) { @@ -2089,23 +1777,12 @@ pub async fn try_provision_account_on_exhausted( input.aliases.as_ref().map_or(0, Vec::len) ); - let is_batch = input - .t_batch_url - .as_ref() - .is_some_and(|u| !u.trim().is_empty()); + let is_batch = input.t_batch_url.as_ref().is_some_and(|u| !u.trim().is_empty()); let sources_file_path = app_state.app_config.paths.load().sources_file_path.clone(); let sources_path = PathBuf::from(&sources_file_path); - if let Some(outcome) = try_renew_expired_account( - app_state, - input, - panel_cfg, - is_batch, - sources_path.as_path(), - true, - optional, - ) - .await + if let Some(outcome) = + try_renew_expired_account(app_state, input, panel_cfg, is_batch, sources_path.as_path(), true, optional).await { debug_if_enabled!( "panel_api: provisioning succeeded via client_renew for input {}", @@ -2113,23 +1790,11 @@ pub async fn try_provision_account_on_exhausted( ); return Some(outcome); } - let created = try_create_new_account( - app_state, - input, - panel_cfg, - is_batch, - sources_path.as_path(), - optional, - ) - .await; + let created = try_create_new_account(app_state, input, panel_cfg, is_batch, sources_path.as_path(), optional).await; debug_if_enabled!( "panel_api: provisioning via client_new for input {} => {}", sanitize_sensitive_info(&input.name), - if created.is_some() { - "success" - } else { - "failed" - } + if created.is_some() { "success" } else { "failed" } ); created } @@ -2160,14 +1825,11 @@ async fn ensure_alias_pool_min( let mut changed = false; let mut provisioned = 0_u16; - let max_pool = resolve_alias_pool_limits(app_state.as_ref(), &input.name, panel_cfg) - .ok() - .and_then(|(_, max)| max); + let max_pool = resolve_alias_pool_limits(app_state.as_ref(), &input.name, panel_cfg).ok().and_then(|(_, max)| max); let mut existing_names: HashSet> = accounts.iter().map(|a| a.name.clone()).collect(); let max_attempts = usize::from(min_pool).saturating_add(10); for _ in 0..max_attempts { - let current_valid = - count_valid_alias_accounts_at(accounts, &input.name, effective_now); + let current_valid = count_valid_alias_accounts_at(accounts, &input.name, effective_now); if current_valid >= usize::from(min_pool) { break; } @@ -2180,9 +1842,7 @@ async fn ensure_alias_pool_min( let expired_index = accounts .iter() .enumerate() - .filter(|(_, acct)| { - acct.name != input.name && is_input_expired_at(acct.exp_date, effective_now) - }) + .filter(|(_, acct)| acct.name != input.name && is_input_expired_at(acct.exp_date, effective_now)) .min_by_key(|(_, acct)| acct.exp_date.unwrap_or(i64::MAX)) .map(|(idx, _)| idx); @@ -2192,13 +1852,8 @@ async fn ensure_alias_pool_min( break; }; if renew_enabled { - match panel_client_renew( - app_state.as_ref(), - panel_cfg, - acct.username.as_str(), - acct.password.as_str(), - ) - .await + match panel_client_renew(app_state.as_ref(), panel_cfg, acct.username.as_str(), acct.password.as_str()) + .await { Ok(()) => { provisioned = provisioned.saturating_add(1); @@ -2235,8 +1890,7 @@ async fn ensure_alias_pool_min( acct_mut.exp_date = Some(new_exp); } if let Some(csv_path) = csv_path { - let _csv_lock = - app_state.app_config.file_locks.write_lock(csv_path).await; + let _csv_lock = app_state.app_config.file_locks.write_lock(csv_path).await; if let Err(err) = csv_patch_batch_update_exp_date( input.input_type, csv_path, @@ -2247,10 +1901,7 @@ async fn ensure_alias_pool_min( ) .await { - debug_if_enabled!( - "panel_api failed to persist renew exp_date to csv: {}", - err - ); + debug_if_enabled!("panel_api failed to persist renew exp_date to csv: {}", err); } else { changed = true; } @@ -2282,20 +1933,14 @@ async fn ensure_alias_pool_min( match panel_client_new(app_state.as_ref(), panel_cfg).await { Ok((username, password, base_url_from_resp)) => { let base_url = base_url_from_resp.unwrap_or_else(|| input.url.clone()); - let base_url = - get_base_url_from_str(base_url.as_str()).unwrap_or_else(|| base_url.clone()); + let base_url = get_base_url_from_str(base_url.as_str()).unwrap_or_else(|| base_url.clone()); - let alias_name = - derive_unique_alias_name_set(&existing_names, &input.name, &username); + let alias_name = derive_unique_alias_name_set(&existing_names, &input.name, &username); existing_names.insert(alias_name.clone().into()); if adult_enabled { - if let Err(err) = panel_client_adult_content( - app_state.as_ref(), - panel_cfg, - Some((&username, &password)), - ) - .await + if let Err(err) = + panel_client_adult_content(app_state.as_ref(), panel_cfg, Some((&username, &password))).await { debug_if_enabled!( "panel_api client_adult_content failed for {}: {}", @@ -2305,16 +1950,10 @@ async fn ensure_alias_pool_min( } } - let exp_date = panel_client_info( - app_state.as_ref(), - panel_cfg, - &username, - &password, - time_ctx, - ) - .await - .ok() - .flatten(); + let exp_date = panel_client_info(app_state.as_ref(), panel_cfg, &username, &password, time_ctx) + .await + .ok() + .flatten(); accounts.push(AccountCredentials { name: alias_name.clone().into(), @@ -2361,10 +2000,7 @@ async fn ensure_alias_pool_min( } } Err(err) => { - debug_if_enabled!( - "panel_api client_new failed: {}", - sanitize_sensitive_info(err.to_string().as_str()) - ); + debug_if_enabled!("panel_api client_new failed: {}", sanitize_sensitive_info(err.to_string().as_str())); break; } } @@ -2397,23 +2033,12 @@ async fn sync_panel_api_for_input_on_boot( let optional = resolve_panel_api_optional_flags(panel_cfg, &input.name); let input_name = &input.name; - let _input_lock = app_state - .app_config - .file_locks - .write_lock_str(format!("panel_api:{input_name}").as_str()) - .await; + let _input_lock = app_state.app_config.file_locks.write_lock_str(format!("panel_api:{input_name}").as_str()).await; let mut any_change = false; - let is_batch = input - .t_batch_url - .as_ref() - .is_some_and(|u| !u.trim().is_empty()); + let is_batch = input.t_batch_url.as_ref().is_some_and(|u| !u.trim().is_empty()); let batch_url = input.t_batch_url.as_deref().unwrap_or_default(); - let csv_path = if is_batch { - get_csv_file_path(batch_url).ok() - } else { - None - }; + let csv_path = if is_batch { get_csv_file_path(batch_url).ok() } else { None }; let mut sources_yml_patches: Vec = Vec::new(); let mut pending_sources_yml = false; @@ -2432,14 +2057,9 @@ async fn sync_panel_api_for_input_on_boot( if panel_cfg.alias_pool.is_some() && csv_path.is_none() - && input - .aliases - .as_ref() - .is_some_and(|aliases| aliases_need_sort_config(aliases)) + && input.aliases.as_ref().is_some_and(|aliases| aliases_need_sort_config(aliases)) { - sources_yml_patches.push(SourcesYmlPatch::SortAliases { - input_name: input.name.clone(), - }); + sources_yml_patches.push(SourcesYmlPatch::SortAliases { input_name: input.name.clone() }); pending_sources_yml = true; } @@ -2453,27 +2073,22 @@ async fn sync_panel_api_for_input_on_boot( (None, None, None) }; - let panel_expire_raw = match panel_client_info_raw( - app_state.as_ref(), - panel_cfg, - root_username.as_str(), - root_password.as_str(), - ) - .await - { - Ok(expire) => expire, - Err(err) => { - debug_if_enabled!( - "panel_api root client_info (raw) failed for {}: {}", - sanitize_sensitive_info(&input.name), - sanitize_sensitive_info(err.to_string().as_str()) - ); - None - } - }; + let panel_expire_raw = + match panel_client_info_raw(app_state.as_ref(), panel_cfg, root_username.as_str(), root_password.as_str()) + .await + { + Ok(expire) => expire, + Err(err) => { + debug_if_enabled!( + "panel_api root client_info (raw) failed for {}: {}", + sanitize_sensitive_info(&input.name), + sanitize_sensitive_info(err.to_string().as_str()) + ); + None + } + }; - let mut expire_mode = - resolve_panel_expire_mode(root_exp_date, panel_expire_raw.as_deref(), server_tz); + let mut expire_mode = resolve_panel_expire_mode(root_exp_date, panel_expire_raw.as_deref(), server_tz); let mut server_tz = server_tz; let mut skew_secs = skew_secs.or_else(|| cached_entry.as_ref().and_then(|e| e.skew_secs)); @@ -2489,23 +2104,14 @@ async fn sync_panel_api_for_input_on_boot( } } - time_ctx = Some(PanelApiTimeContext { - expire_mode, - server_tz, - }); + time_ctx = Some(PanelApiTimeContext { expire_mode, server_tz }); effective_now = apply_clock_skew(get_current_timestamp(), skew_secs.unwrap_or_default()); time_cache.inputs.insert( input.name.to_string(), - PanelApiTimeCacheEntry { - expire_mode, - server_tz: server_tz.map(|tz| tz.name().to_string()), - skew_secs, - }, + PanelApiTimeCacheEntry { expire_mode, server_tz: server_tz.map(|tz| tz.name().to_string()), skew_secs }, ); - if let Err(err) = - persist_panel_api_time_cache(app_state.as_ref(), &cache_path, &time_cache).await - { + if let Err(err) = persist_panel_api_time_cache(app_state.as_ref(), &cache_path, &time_cache).await { debug_if_enabled!( "panel_api failed to persist time cache for {}: {}", sanitize_sensitive_info(&input.name), @@ -2523,14 +2129,8 @@ async fn sync_panel_api_for_input_on_boot( ); } else if let Some(cached) = cached_entry.as_ref() { let server_tz = parse_cached_tz(cached.server_tz.clone()); - time_ctx = Some(PanelApiTimeContext { - expire_mode: cached.expire_mode, - server_tz, - }); - effective_now = apply_clock_skew( - get_current_timestamp(), - cached.skew_secs.unwrap_or_default(), - ); + time_ctx = Some(PanelApiTimeContext { expire_mode: cached.expire_mode, server_tz }); + effective_now = apply_clock_skew(get_current_timestamp(), cached.skew_secs.unwrap_or_default()); let server_tz_name = server_tz.as_ref().map_or("none", |tz| Tz::name(*tz)); debug_if_enabled!( "panel_api time context fallback for input {}: expire_mode={:?}, tz={}, skew_secs={}", @@ -2547,25 +2147,20 @@ async fn sync_panel_api_for_input_on_boot( } for acct in &mut accounts { - let new_exp = match panel_client_info( - app_state.as_ref(), - panel_cfg, - &acct.username, - &acct.password, - time_ctx.as_ref(), - ) - .await - { - Ok(v) => v, - Err(err) => { - debug_if_enabled!( - "panel_api client_info failed for {}: {}", - sanitize_sensitive_info(&acct.name), - sanitize_sensitive_info(err.to_string().as_str()) - ); - None - } - }; + let new_exp = + match panel_client_info(app_state.as_ref(), panel_cfg, &acct.username, &acct.password, time_ctx.as_ref()) + .await + { + Ok(v) => v, + Err(err) => { + debug_if_enabled!( + "panel_api client_info failed for {}: {}", + sanitize_sensitive_info(&acct.name), + sanitize_sensitive_info(err.to_string().as_str()) + ); + None + } + }; let Some(new_exp) = new_exp else { continue; }; @@ -2585,10 +2180,7 @@ async fn sync_panel_api_for_input_on_boot( ) .await { - debug_if_enabled!( - "panel_api boot sync failed to persist exp_date to csv: {}", - err - ); + debug_if_enabled!("panel_api boot sync failed to persist exp_date to csv: {}", err); continue; } any_change = true; @@ -2626,19 +2218,16 @@ async fn sync_panel_api_for_input_on_boot( let root_exp_date = accounts[root_idx].exp_date; let root_exp_missing = root_exp_date.is_none(); let root_expired = match root_exp_date { - Some(ts) => u64::try_from(ts) - .map_or(true, |exp_ts| exp_ts <= now), + Some(ts) => u64::try_from(ts).map_or(true, |exp_ts| exp_ts <= now), None => false, }; let root_expiring = match root_exp_date { - Some(ts) => u64::try_from(ts) - .map_or(true, |exp_ts| exp_ts > now && exp_ts <= offset_deadline), + Some(ts) => u64::try_from(ts).map_or(true, |exp_ts| exp_ts > now && exp_ts <= offset_deadline), None => false, }; let should_refresh_root = root_exp_missing || root_expired || root_expiring; - let root_exp_display = - root_exp_date.map_or_else(|| "None".to_string(), |ts| ts.to_string()); + let root_exp_display = root_exp_date.map_or_else(|| "None".to_string(), |ts| ts.to_string()); debug_if_enabled!( "panel_api boot/update root status for input {} (offset={}s): exp_date={}, expired={}, expiring(offset)={}", sanitize_sensitive_info(&input.name), @@ -2654,12 +2243,12 @@ async fn sync_panel_api_for_input_on_boot( let old_password = accounts[root_idx].password.clone(); debug_if_enabled!( - "panel_api boot/update refreshing root account {} for input {} (exp_date={}, offset={}s)", - sanitize_sensitive_info(&old_username), - sanitize_sensitive_info(&input.name), - root_exp_display, - offset_secs - ); + "panel_api boot/update refreshing root account {} for input {} (exp_date={}, offset={}s)", + sanitize_sensitive_info(&old_username), + sanitize_sensitive_info(&input.name), + root_exp_display, + offset_secs + ); let (active_username, active_password, creds_changed) = if renew_enabled { match panel_client_renew( @@ -2686,14 +2275,12 @@ async fn sync_panel_api_for_input_on_boot( provisioned_root = 1; // Variant B: if the old root is still valid but within offset window, // keep it as a new alias entry so we don't lose usable credentials. - let park_old_root_as_alias = root_expiring - && !root_expired - && root_exp_date.is_some(); + let park_old_root_as_alias = + root_expiring && !root_expired && root_exp_date.is_some(); if park_old_root_as_alias { - let base_url = - get_base_url_from_str(input.url.as_str()) - .unwrap_or_else(|| input.url.clone()); + let base_url = get_base_url_from_str(input.url.as_str()) + .unwrap_or_else(|| input.url.clone()); let alias_name = derive_unique_alias_name_set( &existing_names, &input.name, @@ -2702,19 +2289,15 @@ async fn sync_panel_api_for_input_on_boot( existing_names.insert(alias_name.clone().into()); if let Some(csv_path) = csv_path.as_ref() { - let batch_type = - if input.input_type == InputType::Xtream { - InputType::XtreamBatch - } else if input.input_type == InputType::M3u { - InputType::M3uBatch - } else { - input.input_type - }; - let _csv_lock = app_state - .app_config - .file_locks - .write_lock(csv_path) - .await; + let batch_type = if input.input_type == InputType::Xtream { + InputType::XtreamBatch + } else if input.input_type == InputType::M3u { + InputType::M3uBatch + } else { + input.input_type + }; + let _csv_lock = + app_state.app_config.file_locks.write_lock(csv_path).await; if let Err(err) = csv_patch_batch_append( csv_path, batch_type, @@ -2735,16 +2318,14 @@ async fn sync_panel_api_for_input_on_boot( any_change = true; } } else { - sources_yml_patches.push( - SourcesYmlPatch::AddAlias { - input_name: input.name.clone(), - alias_name: alias_name.clone().into(), - base_url, - username: old_username.clone(), - password: old_password.clone(), - exp_date: root_exp_date, - }, - ); + sources_yml_patches.push(SourcesYmlPatch::AddAlias { + input_name: input.name.clone(), + alias_name: alias_name.clone().into(), + base_url, + username: old_username.clone(), + password: old_password.clone(), + exp_date: root_exp_date, + }); pending_sources_yml = true; } @@ -2763,11 +2344,7 @@ async fn sync_panel_api_for_input_on_boot( } if let Some(csv_path) = csv_path.as_ref() { - let _csv_lock = app_state - .app_config - .file_locks - .write_lock(csv_path) - .await; + let _csv_lock = app_state.app_config.file_locks.write_lock(csv_path).await; if let Err(err) = csv_patch_batch_update_credentials( input.input_type, csv_path, @@ -2789,14 +2366,12 @@ async fn sync_panel_api_for_input_on_boot( any_change = true; } } else { - sources_yml_patches.push( - SourcesYmlPatch::UpdateRootCredentials { - input_name: input.name.clone(), - username: new_username.clone(), - password: new_password.clone(), - exp_date: None, - }, - ); + sources_yml_patches.push(SourcesYmlPatch::UpdateRootCredentials { + input_name: input.name.clone(), + username: new_username.clone(), + password: new_password.clone(), + exp_date: None, + }); pending_sources_yml = true; } @@ -2822,17 +2397,13 @@ async fn sync_panel_api_for_input_on_boot( match panel_client_new(app_state.as_ref(), panel_cfg).await { Ok((new_username, new_password, _base_url_from_resp)) => { provisioned_root = 1; - let park_old_root_as_alias = - root_expiring && !root_expired && root_exp_date.is_some(); + let park_old_root_as_alias = root_expiring && !root_expired && root_exp_date.is_some(); if park_old_root_as_alias { - let base_url = get_base_url_from_str(input.url.as_str()) - .unwrap_or_else(|| input.url.clone()); - let alias_name = derive_unique_alias_name_set( - &existing_names, - &input.name, - old_username.as_str(), - ); + let base_url = + get_base_url_from_str(input.url.as_str()).unwrap_or_else(|| input.url.clone()); + let alias_name = + derive_unique_alias_name_set(&existing_names, &input.name, old_username.as_str()); existing_names.insert(alias_name.clone().into()); if let Some(csv_path) = csv_path.as_ref() { @@ -2843,8 +2414,7 @@ async fn sync_panel_api_for_input_on_boot( } else { input.input_type }; - let _csv_lock = - app_state.app_config.file_locks.write_lock(csv_path).await; + let _csv_lock = app_state.app_config.file_locks.write_lock(csv_path).await; if let Err(err) = csv_patch_batch_append( csv_path, batch_type, @@ -2857,10 +2427,10 @@ async fn sync_panel_api_for_input_on_boot( .await { debug_if_enabled!( - "panel_api boot/update failed to park old root as csv alias {}: {}", - sanitize_sensitive_info(&alias_name), - err - ); + "panel_api boot/update failed to park old root as csv alias {}: {}", + sanitize_sensitive_info(&alias_name), + err + ); } else { any_change = true; } @@ -2891,8 +2461,7 @@ async fn sync_panel_api_for_input_on_boot( } if let Some(csv_path) = csv_path.as_ref() { - let _csv_lock = - app_state.app_config.file_locks.write_lock(csv_path).await; + let _csv_lock = app_state.app_config.file_locks.write_lock(csv_path).await; if let Err(err) = csv_patch_batch_update_credentials( input.input_type, csv_path, @@ -2978,9 +2547,9 @@ async fn sync_panel_api_for_input_on_boot( .await; if !ready { debug_if_enabled!( - "panel_api boot/update probe timeout for root {}; skipping exp_date refresh", - sanitize_sensitive_info(&input.name) - ); + "panel_api boot/update probe timeout for root {}; skipping exp_date refresh", + sanitize_sensitive_info(&input.name) + ); } else if let Some(new_exp) = refreshed_exp { if let Some(csv_path) = csv_path.as_ref() { let _csv_lock = app_state.app_config.file_locks.write_lock(csv_path).await; @@ -3009,10 +2578,10 @@ async fn sync_panel_api_for_input_on_boot( }; if let Err(err) = result { debug_if_enabled!( - "panel_api boot/update failed to persist root exp_date to csv for {}: {}", - sanitize_sensitive_info(&input.name), - err - ); + "panel_api boot/update failed to persist root exp_date to csv for {}: {}", + sanitize_sensitive_info(&input.name), + err + ); } else { accounts[root_idx].exp_date = Some(new_exp); any_change = true; @@ -3055,23 +2624,16 @@ async fn sync_panel_api_for_input_on_boot( let offset_deadline = now.saturating_add(offset_secs); let root_valid = root_counts_towards_pool_at(&accounts, &input.name, now); let desired_aliases = min_pool.filter(|m| *m > 0).map_or_else( - || { - u16::try_from(accounts.iter().filter(|a| a.name != input.name).count()) - .unwrap_or(u16::MAX) - }, + || u16::try_from(accounts.iter().filter(|a| a.name != input.name).count()).unwrap_or(u16::MAX), |min_pool| min_pool.saturating_sub(u16::from(root_valid)), ); let expiring_aliases = accounts .iter() - .filter(|a| { - a.name != input.name && is_expiring_with_offset_at(a.exp_date, offset_secs, now) - }) - .count(); - let expired_aliases = accounts - .iter() - .filter(|a| a.name != input.name && is_input_expired_at(a.exp_date, now)) + .filter(|a| a.name != input.name && is_expiring_with_offset_at(a.exp_date, offset_secs, now)) .count(); + let expired_aliases = + accounts.iter().filter(|a| a.name != input.name && is_input_expired_at(a.exp_date, now)).count(); let valid_aliases_beyond_offset = accounts .iter() @@ -3083,8 +2645,7 @@ async fn sync_panel_api_for_input_on_boot( return false; } match a.exp_date { - Some(ts) => u64::try_from(ts) - .is_ok_and(|exp_ts| exp_ts > offset_deadline), + Some(ts) => u64::try_from(ts).is_ok_and(|exp_ts| exp_ts > offset_deadline), None => false, } }) @@ -3108,8 +2669,7 @@ async fn sync_panel_api_for_input_on_boot( } match a.exp_date { None => true, - Some(ts) => u64::try_from(ts) - .map_or(true, |exp_ts| exp_ts <= offset_deadline), + Some(ts) => u64::try_from(ts).map_or(true, |exp_ts| exp_ts <= offset_deadline), } }) .map(|(idx, _)| idx) @@ -3123,28 +2683,17 @@ async fn sync_panel_api_for_input_on_boot( } }); - let valid_aliases_beyond_offset_u16 = - u16::try_from(valid_aliases_beyond_offset).unwrap_or(u16::MAX); - let needed_refresh_aliases_u16 = - desired_aliases_u16.saturating_sub(valid_aliases_beyond_offset_u16); - let planned_refresh_aliases = refresh_candidates - .len() - .min(usize::from(needed_refresh_aliases_u16)); - let refresh_plan: Vec = refresh_candidates - .into_iter() - .take(planned_refresh_aliases) - .collect(); + let valid_aliases_beyond_offset_u16 = u16::try_from(valid_aliases_beyond_offset).unwrap_or(u16::MAX); + let needed_refresh_aliases_u16 = desired_aliases_u16.saturating_sub(valid_aliases_beyond_offset_u16); + let planned_refresh_aliases = refresh_candidates.len().min(usize::from(needed_refresh_aliases_u16)); + let refresh_plan: Vec = refresh_candidates.into_iter().take(planned_refresh_aliases).collect(); (refresh_plan, planned_refresh_aliases) } else { (Vec::new(), 0) }; let log_pool = alias_pool_has_min(panel_cfg); - let enabled_users = if log_pool { - count_enabled_proxy_users(app_state.as_ref(), &input.name) - } else { - 0 - }; + let enabled_users = if log_pool { count_enabled_proxy_users(app_state.as_ref(), &input.name) } else { 0 }; if log_pool { debug_if_enabled!( "panel_api boot/update provisioning aliases for input {} (offset={}s): desired={}, valid_beyond_offset={}, expiring(offset)={}, expired={}, refresh_planned(offset)={}, missing={}", @@ -3189,13 +2738,7 @@ async fn sync_panel_api_for_input_on_boot( ); let (active_username, active_password, creds_changed) = if renew_enabled { - match panel_client_renew( - app_state.as_ref(), - panel_cfg, - old_username.as_str(), - old_password.as_str(), - ) - .await + match panel_client_renew(app_state.as_ref(), panel_cfg, old_username.as_str(), old_password.as_str()).await { Ok(()) => { provisioned_aliases = provisioned_aliases.saturating_add(1); @@ -3210,16 +2753,12 @@ async fn sync_panel_api_for_input_on_boot( if new_enabled { match panel_client_new(app_state.as_ref(), panel_cfg).await { Ok((new_username, new_password, base_url_from_resp)) => { + let base_url = base_url_from_resp.unwrap_or_else(|| input.url.clone()); let base_url = - base_url_from_resp.unwrap_or_else(|| input.url.clone()); - let base_url = get_base_url_from_str(base_url.as_str()) - .unwrap_or_else(|| base_url.clone()); + get_base_url_from_str(base_url.as_str()).unwrap_or_else(|| base_url.clone()); - let alias_name = derive_unique_alias_name_set( - &existing_names, - &input.name, - &new_username, - ); + let alias_name = + derive_unique_alias_name_set(&existing_names, &input.name, &new_username); existing_names.insert(alias_name.clone().into()); if adult_enabled { @@ -3257,8 +2796,7 @@ async fn sync_panel_api_for_input_on_boot( } else { input.input_type }; - let _csv_lock = - app_state.app_config.file_locks.write_lock(csv_path).await; + let _csv_lock = app_state.app_config.file_locks.write_lock(csv_path).await; if let Err(err) = csv_patch_batch_append( csv_path, batch_type, @@ -3316,11 +2854,9 @@ async fn sync_panel_api_for_input_on_boot( match panel_client_new(app_state.as_ref(), panel_cfg).await { Ok((new_username, new_password, base_url_from_resp)) => { let base_url = base_url_from_resp.unwrap_or_else(|| input.url.clone()); - let base_url = get_base_url_from_str(base_url.as_str()) - .unwrap_or_else(|| base_url.clone()); + let base_url = get_base_url_from_str(base_url.as_str()).unwrap_or_else(|| base_url.clone()); - let alias_name = - derive_unique_alias_name_set(&existing_names, &input.name, &new_username); + let alias_name = derive_unique_alias_name_set(&existing_names, &input.name, &new_username); existing_names.insert(alias_name.clone().into()); if adult_enabled { @@ -3570,28 +3106,21 @@ async fn sync_panel_api_for_input_on_boot( match csv_patch_batch_remove_expired(input.input_type, csv_path).await { Ok(true) => any_change = true, Ok(false) => {} - Err(err) => debug_if_enabled!( - "panel_api boot sync failed to remove expired csv accounts: {}", - err - ), + Err(err) => debug_if_enabled!("panel_api boot sync failed to remove expired csv accounts: {}", err), } } else { - sources_yml_patches.push(SourcesYmlPatch::RemoveExpiredAliases { - input_name: input.name.clone(), - }); + sources_yml_patches.push(SourcesYmlPatch::RemoveExpiredAliases { input_name: input.name.clone() }); pending_sources_yml = true; } } if optional.contains(PanelApiOptionalFlags::AccountInfo) { - let creds = accounts - .first() - .map(|acct| (acct.username.as_str(), acct.password.as_str())); + let creds = accounts.first().map(|acct| (acct.username.as_str(), acct.password.as_str())); match panel_account_info(app_state.as_ref(), panel_cfg, creds).await { Ok(Some(credits)) => { let normalized = credits.trim().to_string(); if !normalized.is_empty() - /* && panel_cfg.credits.as_deref().map(str::trim) != Some(normalized.as_str()) */ + /* && panel_cfg.credits.as_deref().map(str::trim) != Some(normalized.as_str()) */ { sources_yml_patches.push(SourcesYmlPatch::UpdatePanelApiCredits { input_name: input.name.clone(), @@ -3627,26 +3156,18 @@ async fn sync_panel_api_for_input_on_boot( } if pending_sources_yml { - let _src_lock = app_state - .app_config - .file_locks - .write_lock(sources_path) - .await; - match persist_sources_yml_with_patches(app_state, sources_path, &sources_yml_patches).await - { + let _src_lock = app_state.app_config.file_locks.write_lock(sources_path).await; + match persist_sources_yml_with_patches(app_state, sources_path, &sources_yml_patches).await { Ok(true) => any_change = true, Ok(false) => {} - Err(err) => debug_if_enabled!( - "panel_api boot sync failed to persist source.yml patches: {}", - err - ), + Err(err) => debug_if_enabled!("panel_api boot sync failed to persist source.yml patches: {}", err), } } any_change } -pub(crate) async fn sync_panel_api_exp_dates_on_boot(app_state: &Arc) { +pub(crate) async fn sync_panel_api_exp_dates(app_state: &Arc) { let sources_file_path = app_state.app_config.paths.load().sources_file_path.clone(); let sources_path = PathBuf::from(&sources_file_path); let mut any_change = false; @@ -3667,20 +3188,18 @@ pub(crate) async fn sync_panel_api_exp_dates_on_boot(app_state: &Arc) } } -pub(crate) async fn sync_panel_api_alias_pool_for_target( - app_state: &Arc, - target_name: &str, -) { +pub(crate) async fn sync_panel_api_exp_dates_on_boot(app_state: &Arc) { + sync_panel_api_exp_dates(app_state).await; +} + +pub(crate) async fn sync_panel_api_alias_pool_for_target(app_state: &Arc, target_name: &str) { let sources_file_path = app_state.app_config.paths.load().sources_file_path.clone(); let sources_path = PathBuf::from(&sources_file_path); let mut any_change = false; let sources = app_state.app_config.sources.load(); for source in &sources.sources { - let target_match = source - .targets - .iter() - .any(|target| target.name.eq_ignore_ascii_case(target_name)); + let target_match = source.targets.iter().any(|target| target.name.eq_ignore_ascii_case(target_name)); if !target_match { continue; } @@ -3728,12 +3247,7 @@ fn provisioning_method_to_reqwest(method: PanelApiProvisioningMethod) -> Method } } -fn build_player_api_action_url( - base_url: &str, - username: &str, - password: &str, - action: &str, -) -> Option { +fn build_player_api_action_url(base_url: &str, username: &str, password: &str, action: &str) -> Option { let url = Url::parse(base_url).ok()?; let host = url.host_str()?; let scheme = url.scheme(); @@ -3768,21 +3282,10 @@ impl PanelApiProbeTarget { } } -fn build_panel_api_probe_targets( - input: &ConfigInput, - username: &str, - password: &str, -) -> Vec { +fn build_panel_api_probe_targets(input: &ConfigInput, username: &str, password: &str) -> Vec { let mut targets = Vec::new(); - for action in [ - "client_info", - "get_live_categories", - "get_series_categories", - "get_vod_categories", - ] { - if let Some(url) = - build_player_api_action_url(input.url.as_str(), username, password, action) - { + for action in ["client_info", "get_live_categories", "get_series_categories", "get_vod_categories"] { + if let Some(url) = build_player_api_action_url(input.url.as_str(), username, password, action) { targets.push(PanelApiProbeTarget::PlayerApi { action, url }); } } @@ -3843,18 +3346,11 @@ async fn probe_panel_api_test_url( ) -> Result { let client = app_state.http_client.load(); let request_method = provisioning_method_to_reqwest(method); - let response = client - .request(request_method, test_url.clone()) - .send() - .await?; + let response = client.request(request_method, test_url.clone()).send().await?; Ok(response.status()) } -async fn apply_provisioning_cooldown( - panel_cfg: &PanelApiConfig, - account_name: &str, - input_name: &Arc, -) { +async fn apply_provisioning_cooldown(panel_cfg: &PanelApiConfig, account_name: &str, input_name: &Arc) { let cooldown_secs = panel_cfg.provisioning.cooldown_sec; if cooldown_secs == 0 { return; @@ -3890,11 +3386,7 @@ async fn wait_for_panel_api_account_ready( return false; } - let targets_list = probe_targets - .iter() - .map(PanelApiProbeTarget::action) - .collect::>() - .join(","); + let targets_list = probe_targets.iter().map(PanelApiProbeTarget::action).collect::>().join(","); debug_if_enabled!( "panel_api probe start for {} (input={} timeout={}s interval={}s method={}) targets={}", sanitize_sensitive_info(account_name), @@ -3912,8 +3404,7 @@ async fn wait_for_panel_api_account_ready( loop { attempt += 1; debug_if_enabled!("panel_api probe attempt {}", attempt); - if probe_panel_api_targets(app_state, probe_method, &probe_targets, &mut done_targets).await - { + if probe_panel_api_targets(app_state, probe_method, &probe_targets, &mut done_targets).await { apply_provisioning_cooldown(panel_cfg, account_name, &input.name).await; return true; } @@ -3926,11 +3417,7 @@ async fn wait_for_panel_api_account_ready( return false; } let remaining = deadline.checked_duration_since(now).unwrap_or_default(); - let sleep_for = if remaining < probe_delay { - remaining - } else { - probe_delay - }; + let sleep_for = if remaining < probe_delay { remaining } else { probe_delay }; tokio::time::sleep(sleep_for).await; } } @@ -3949,10 +3436,7 @@ pub(crate) async fn run_panel_api_provisioning_probe( sanitize_sensitive_info(&input.name) ); stop_signal.notify(); - let _ = app_state - .connection_manager - .kick_connection(&addr, virtual_id, 0) - .await; + let _ = app_state.connection_manager.kick_connection(&addr, virtual_id, 0).await; return Ok(()); }; if !panel_cfg.enabled { @@ -3961,10 +3445,7 @@ pub(crate) async fn run_panel_api_provisioning_probe( sanitize_sensitive_info(&input.name) ); stop_signal.notify(); - let _ = app_state - .connection_manager - .kick_connection(&addr, virtual_id, 0) - .await; + let _ = app_state.connection_manager.kick_connection(&addr, virtual_id, 0).await; return Ok(()); } if panel_cfg.url.trim().is_empty() { @@ -3973,10 +3454,7 @@ pub(crate) async fn run_panel_api_provisioning_probe( sanitize_sensitive_info(&input.name) ); stop_signal.notify(); - let _ = app_state - .connection_manager - .kick_connection(&addr, virtual_id, 0) - .await; + let _ = app_state.connection_manager.kick_connection(&addr, virtual_id, 0).await; return Ok(()); } @@ -4023,10 +3501,7 @@ pub(crate) async fn run_panel_api_provisioning_probe( sanitize_sensitive_info(&input.name), sanitize_sensitive_info(addr.to_string().as_str()) ); - let _ = app_state - .connection_manager - .kick_connection(&addr, virtual_id, 0) - .await; + let _ = app_state.connection_manager.kick_connection(&addr, virtual_id, 0).await; return Ok(()); }; @@ -4041,10 +3516,7 @@ pub(crate) async fn run_panel_api_provisioning_probe( sanitize_sensitive_info(&input.name) ); stop_signal.notify(); - let _ = app_state - .connection_manager - .kick_connection(&addr, virtual_id, 0) - .await; + let _ = app_state.connection_manager.kick_connection(&addr, virtual_id, 0).await; return Ok(()); }; @@ -4086,11 +3558,7 @@ pub(crate) async fn run_panel_api_provisioning_probe( break; } let remaining = deadline.checked_duration_since(now).unwrap_or_default(); - let sleep_for = if remaining < probe_delay { - remaining - } else { - probe_delay - }; + let sleep_for = if remaining < probe_delay { remaining } else { probe_delay }; tokio::time::sleep(sleep_for).await; } @@ -4113,10 +3581,7 @@ pub(crate) async fn run_panel_api_provisioning_probe( sanitize_sensitive_info(&input.name), sanitize_sensitive_info(addr.to_string().as_str()) ); - let _ = app_state - .connection_manager - .kick_connection(&addr, virtual_id, 0) - .await; + let _ = app_state.connection_manager.kick_connection(&addr, virtual_id, 0).await; Ok(()) } @@ -4130,19 +3595,15 @@ pub fn create_panel_api_provisioning_stream_details( ) -> StreamDetails { let stop_signal = Arc::new(AtomicOnceFlag::new()); let headers = [("connection".to_string(), "close".to_string())]; - let (stream, stream_info) = create_panel_api_provisioning_stream_with_stop( - &app_state.app_config, - &headers, - Arc::clone(&stop_signal), - ); + let (stream, stream_info) = + create_panel_api_provisioning_stream_with_stop(&app_state.app_config, &headers, Arc::clone(&stop_signal)); if stream.is_none() { debug_if_enabled!( "panel_api provisioning stream missing; falling back to provider exhausted for input {}", sanitize_sensitive_info(&input.name) ); - let (stream, stream_info) = - create_provider_connections_exhausted_stream(&app_state.app_config, &[]); + let (stream, stream_info) = create_provider_connections_exhausted_stream(&app_state.app_config, &[]); return StreamDetails { stream, stream_info, @@ -4159,14 +3620,9 @@ pub fn create_panel_api_provisioning_stream_details( let input_clone = input.clone(); let stop_clone = Arc::clone(&stop_signal); tokio::spawn(async move { - if let Err(err) = run_panel_api_provisioning_probe( - app_state_clone, - input_clone, - stop_clone, - addr, - virtual_id, - ) - .await { + if let Err(err) = + run_panel_api_provisioning_probe(app_state_clone, input_clone, stop_clone, addr, virtual_id).await + { error!("Error running Probe: {err:?}"); } }); diff --git a/backend/src/api/scheduler.rs b/backend/src/api/scheduler.rs index b7bf3ef11..f4360a93f 100644 --- a/backend/src/api/scheduler.rs +++ b/backend/src/api/scheduler.rs @@ -1,17 +1,18 @@ -use crate::api::model::AppState; -use crate::api::panel_api::sync_panel_api_exp_dates_on_boot; -use crate::api::library_scan::spawn_library_scan; -use crate::model::{AppConfig, ProcessTargets, ScheduleConfig}; -use shared::model::ScheduleTaskType; -use crate::processing::processor::exec_processing; -use crate::utils::exit; +use crate::{ + api::{library_scan::spawn_library_scan, model::AppState}, + model::{AppConfig, ProcessTargets, ScheduleConfig}, + processing::processor::exec_processing, + utils::exit, +}; use chrono::{DateTime, FixedOffset, Local}; use cron::Schedule; -use std::str::FromStr; -use std::sync::Arc; -use std::time::{Duration, Instant, SystemTime}; +use shared::{model::ScheduleTaskType, utils::interner_gc}; +use std::{ + str::FromStr, + sync::Arc, + time::{Duration, Instant, SystemTime}, +}; use tokio_util::sync::CancellationToken; -use shared::utils::interner_gc; pub fn datetime_to_instant(datetime: DateTime) -> Instant { // Convert DateTime to SystemTime @@ -21,23 +22,22 @@ pub fn datetime_to_instant(datetime: DateTime) -> Instant { let now_system_time = SystemTime::now(); // Calculate the duration between now and the target time - let duration_until = target_system_time - .duration_since(now_system_time) - .unwrap_or_else(|_| Duration::from_secs(0)); + let duration_until = target_system_time.duration_since(now_system_time).unwrap_or_else(|_| Duration::from_secs(0)); // Get the current Instant and add the duration to calculate the target Instant Instant::now() + duration_until } -pub fn exec_scheduler(client: &reqwest::Client, app_state: &Arc, targets: &Arc, - cancel: &CancellationToken) { +pub fn exec_scheduler( + client: &reqwest::Client, + app_state: &Arc, + targets: &Arc, + cancel: &CancellationToken, +) { let cfg = &app_state.app_config; let config = cfg.config.load(); - let schedules: Vec = if let Some(schedules) = &config.schedules { - schedules.clone() - } else { - vec![] - }; + let schedules: Vec = + if let Some(schedules) = &config.schedules { schedules.clone() } else { vec![] }; for schedule in schedules { let expression = schedule.schedule.clone(); let task_type = schedule.task_type; @@ -46,13 +46,20 @@ pub fn exec_scheduler(client: &reqwest::Client, app_state: &Arc, targe let http_client = client.clone(); let cancel_token = cancel.clone(); tokio::spawn(async move { - start_scheduler(http_client, expression.as_str(), task_type, app_state_clone, exec_targets, cancel_token).await; + start_scheduler(http_client, expression.as_str(), task_type, app_state_clone, exec_targets, cancel_token) + .await; }); } } -async fn start_scheduler(client: reqwest::Client, expression: &str, task_type: ScheduleTaskType, app_state: Arc, - targets: Arc, cancel: CancellationToken) { +async fn start_scheduler( + client: reqwest::Client, + expression: &str, + task_type: ScheduleTaskType, + app_state: Arc, + targets: Arc, + cancel: CancellationToken, +) { match Schedule::from_str(expression) { Ok(schedule) => { let offset = *Local::now().offset(); @@ -77,7 +84,7 @@ async fn start_scheduler(client: reqwest::Client, expression: &str, task_type: S } } } - Err(err) => exit!("Failed to start scheduler: {err}") + Err(err) => exit!("Failed to start scheduler: {err}"), } } @@ -92,12 +99,12 @@ fn run_playlist_update(client: &reqwest::Client, app_state: &Arc, targ let provider_manager = Arc::clone(&app_state.active_provider); let disabled_headers = app_state.get_disabled_headers(); let metadata_manager = Arc::clone(&app_state.metadata_manager); - sync_panel_api_exp_dates_on_boot(&app_state).await; exec_processing( &client, app_config, targets, Some(event_manager), + Some(app_state.clone()), Some(playlist_state), Some(app_state.update_guard.clone()), disabled_headers, @@ -130,7 +137,11 @@ fn run_library_scan(client: &reqwest::Client, app_state: &Arc) { } } -pub fn get_process_targets(cfg: &Arc, process_targets: &Arc, exec_targets: Option<&Vec>) -> Arc { +pub fn get_process_targets( + cfg: &Arc, + process_targets: &Arc, + exec_targets: Option<&Vec>, +) -> Arc { let sources = cfg.sources.load(); if let Ok(user_targets) = sources.validate_targets(exec_targets) { if user_targets.enabled { @@ -138,24 +149,17 @@ pub fn get_process_targets(cfg: &Arc, process_targets: &Arc = user_targets.inputs.iter() - .filter(|&id| process_targets.inputs.contains(id)) - .copied() - .collect(); - let targets: Vec = user_targets.targets.iter() - .filter(|&id| process_targets.inputs.contains(id)) - .copied() - .collect(); - let target_names: Vec = user_targets.target_names.iter() + let inputs: Vec = + user_targets.inputs.iter().filter(|&id| process_targets.inputs.contains(id)).copied().collect(); + let targets: Vec = + user_targets.targets.iter().filter(|&id| process_targets.inputs.contains(id)).copied().collect(); + let target_names: Vec = user_targets + .target_names + .iter() .filter(|&name| process_targets.target_names.contains(name)) .cloned() .collect(); - return Arc::new(ProcessTargets { - enabled: user_targets.enabled, - inputs, - targets, - target_names, - }); + return Arc::new(ProcessTargets { enabled: user_targets.enabled, inputs, targets, target_names }); } } Arc::clone(process_targets) @@ -184,8 +188,10 @@ mod tests { use crate::api::scheduler::datetime_to_instant; use chrono::Local; use cron::Schedule; - use std::str::FromStr; - use std::sync::atomic::{AtomicU8, Ordering}; + use std::{ + str::FromStr, + sync::atomic::{AtomicU8, Ordering}, + }; #[tokio::test] async fn test_run_scheduler() { diff --git a/backend/src/api/serve.rs b/backend/src/api/serve.rs index 479366f49..9415a8982 100644 --- a/backend/src/api/serve.rs +++ b/backend/src/api/serve.rs @@ -1,70 +1,58 @@ -use axum::body::Body; -use axum::extract::Request; -use axum::response::Response; +use crate::api::model::ConnectionManager; +use axum::{body::Body, extract::Request, response::Response}; use futures::FutureExt; use hyper::body::Incoming; -use hyper_util::rt::{TokioExecutor, TokioIo}; -use hyper_util::server::conn::auto::Builder; -use hyper_util::service::TowerToHyperService; +use hyper_util::{ + rt::{TokioExecutor, TokioIo}, + server::conn::auto::Builder, + service::TowerToHyperService, +}; use log::{debug, error, trace}; use socket2::{SockRef, TcpKeepalive}; -use std::convert::Infallible; -use std::fmt::Debug; -use std::net::SocketAddr; -use std::pin::pin; -use std::sync::Arc; -use std::time::Duration; +use std::{convert::Infallible, fmt::Debug, net::SocketAddr, pin::pin, sync::Arc, time::Duration}; use tokio::sync::watch; use tokio_util::sync::CancellationToken; use tower::{Service, ServiceExt}; -use crate::api::model::{ConnectionManager}; #[derive(Debug)] -struct IncomingStream -{ +struct IncomingStream { remote_addr: SocketAddr, } impl IncomingStream { /// Returns the remote address that this stream is bound to. - pub fn remote_addr(&self) -> &SocketAddr { - &self.remote_addr - } + pub fn remote_addr(&self) -> &SocketAddr { &self.remote_addr } } impl axum::extract::connect_info::Connected for SocketAddr { - fn connect_info(target: IncomingStream) -> SocketAddr { - *target.remote_addr() - } + fn connect_info(target: IncomingStream) -> SocketAddr { *target.remote_addr() } } -pub async fn serve(listener: tokio::net::TcpListener, - router: axum::Router<()>, - cancel_token: Option, - connection_manager: &Arc) { +pub async fn serve( + listener: tokio::net::TcpListener, + router: axum::Router<()>, + cancel_token: Option, + connection_manager: &Arc, +) { let (signal_tx, _signal_rx) = watch::channel(()); let mut make_service = router.into_make_service_with_connect_info::(); match cancel_token { - Some(token) => { - loop { - tokio::select! { - () = token.cancelled() => { - break; - } - accept_result = listener.accept() => { - let Ok((socket, remote_addr)) = accept_result else { continue }; - handle_connection(&mut make_service, &signal_tx, socket, remote_addr, Arc::clone(connection_manager)).await; - } + Some(token) => loop { + tokio::select! { + () = token.cancelled() => { + break; + } + accept_result = listener.accept() => { + let Ok((socket, remote_addr)) = accept_result else { continue }; + handle_connection(&mut make_service, &signal_tx, socket, remote_addr, Arc::clone(connection_manager)).await; } } - } - None => { - loop { - let Ok((socket, remote_addr)) = listener.accept().await else { continue }; - handle_connection(&mut make_service, &signal_tx, socket, remote_addr, Arc::clone(connection_manager)).await; - } - } + }, + None => loop { + let Ok((socket, remote_addr)) = listener.accept().await else { continue }; + handle_connection(&mut make_service, &signal_tx, socket, remote_addr, Arc::clone(connection_manager)).await; + }, } } @@ -74,14 +62,15 @@ async fn handle_connection( socket: tokio::net::TcpStream, remote_addr: SocketAddr, connection_manager: Arc, -) -where - M: for<'a> Service + Send + 'static, +) where + M: for<'a> Service + Send + 'static, for<'a> >::Future: Send, - S: Service + Clone + Send + 'static, + S: Service + Clone + Send + 'static, S::Future: Send, { - let Ok(tcp_stream_std) = socket.into_std() else { return; }; + let Ok(tcp_stream_std) = socket.into_std() else { + return; + }; //tcp_stream_std.set_nonblocking(true).ok(); // this is not necessary // Configure keep alive with socket2 @@ -91,7 +80,8 @@ where let keep_alive_interval = 5; let mut keepalive = TcpKeepalive::new(); - keepalive = keepalive.with_time(Duration::from_secs(keep_alive_first_probe)) // Time until the first keepalive probe (idle time) + keepalive = keepalive + .with_time(Duration::from_secs(keep_alive_first_probe)) // Time until the first keepalive probe (idle time) .with_interval(Duration::from_secs(keep_alive_interval)); // Interval between keep alives #[cfg(not(target_os = "windows"))] { @@ -103,15 +93,14 @@ where error!("Failed to set keepalive for {remote_addr}: {e}"); } - let Ok(socket) = tokio::net::TcpStream::from_std(tcp_stream_std) else { return; }; + let Ok(socket) = tokio::net::TcpStream::from_std(tcp_stream_std) else { + return; + }; let io = TokioIo::new(socket); trace!("connection {remote_addr:?} accepted"); - make_service - .ready() - .await - .unwrap_or_else(|err| match err {}); + make_service.ready().await.unwrap_or_else(|err| match err {}); let tower_service = make_service .call(IncomingStream { diff --git a/backend/src/api/setup_api.rs b/backend/src/api/setup_api.rs index 4780cce15..285d968c9 100644 --- a/backend/src/api/setup_api.rs +++ b/backend/src/api/setup_api.rs @@ -1,33 +1,43 @@ -use crate::api::api_utils::serve_file; -use crate::auth::generate_password_from_input; -use crate::model::validate_library_paths_from_dto; -use crate::utils::{ - file_exists, get_default_path, get_default_web_root_path, read_api_proxy_file, read_config_file, - read_sources_file, read_templates_file, resolve_template_persist_file_path, sanitize_sources_for_persist, +use crate::{ + api::api_utils::serve_file, + auth::generate_password_from_input, + model::validate_library_paths_from_dto, + utils::{ + file_exists, get_default_path, get_default_web_root_path, read_api_proxy_file, read_config_file, + read_sources_file, read_templates_file, resolve_template_persist_file_path, sanitize_sources_for_persist, + }, +}; +use axum::{ + extract::{Path, State}, + http::StatusCode, + response::IntoResponse, + Router, }; -use axum::extract::{Path, State}; -use axum::http::StatusCode; -use axum::response::IntoResponse; -use axum::Router; use core::fmt; use log::{error, info, warn}; use rand::Rng; use serde::{Deserialize, Serialize}; use serde_json::json; -use shared::error::TuliproxError; -use shared::foundation::prepare_templates; -use shared::info_err; -use shared::model::{ - ApiProxyConfigDto, ApiProxyServerInfoDto, AppConfigDto, ConfigApiDto, ConfigDto, ConfigPaths, SourcesConfigDto, - PatternTemplate, TemplateDefinitionDto, TokenResponse, WebAuthConfigDto, WebUiConfigDto, TOKEN_NO_AUTH, +use shared::{ + error::TuliproxError, + foundation::prepare_templates, + info_err, + model::{ + ApiProxyConfigDto, ApiProxyServerInfoDto, AppConfigDto, ConfigApiDto, ConfigDto, ConfigPaths, PatternTemplate, + SourcesConfigDto, TemplateDefinitionDto, TokenResponse, WebAuthConfigDto, WebUiConfigDto, TOKEN_NO_AUTH, + }, + utils::{default_kick_secs, hex_encode, DEFAULT_PORT, DEFAULT_WORKING_DIR, USER_FILE}, +}; +use std::{ + collections::{HashMap, HashSet}, + io::ErrorKind, + net::{SocketAddr, UdpSocket}, + path::{Component, Path as FsPath, PathBuf}, + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }, }; -use shared::utils::{default_kick_secs, hex_encode, DEFAULT_PORT, DEFAULT_WORKING_DIR, USER_FILE}; -use std::collections::{HashMap, HashSet}; -use std::io::ErrorKind; -use std::net::{SocketAddr, UdpSocket}; -use std::path::{Component, Path as FsPath, PathBuf}; -use std::sync::atomic::{AtomicBool, Ordering}; -use std::sync::Arc; use tokio::sync::{oneshot, Mutex, RwLock}; use tower_http::services::ServeDir; @@ -251,11 +261,7 @@ fn ensure_setup_defaults(config: &mut ConfigDto) { if config.working_dir.trim().is_empty() { config.working_dir = get_default_path(DEFAULT_WORKING_DIR).display().to_string(); } - if config - .custom_stream_response_path - .as_ref() - .is_none_or(|path| path.trim().is_empty()) - { + if config.custom_stream_response_path.as_ref().is_none_or(|path| path.trim().is_empty()) { config.custom_stream_response_path = Some(DEFAULT_SETUP_CUSTOM_STREAM_RESPONSE_PATH.to_string()); } @@ -376,13 +382,15 @@ fn restore_redacted_web_auth_secret(app_config: &mut AppConfigDto, draft: &AppCo } } -fn restore_redacted_api_proxy_credentials(api_proxy: &mut ApiProxyConfigDto, draft_api_proxy: Option<&ApiProxyConfigDto>) { +fn restore_redacted_api_proxy_credentials( + api_proxy: &mut ApiProxyConfigDto, + draft_api_proxy: Option<&ApiProxyConfigDto>, +) { let mut credentials_by_username: HashMap)> = HashMap::new(); if let Some(draft) = draft_api_proxy { for target in &draft.user { for user in &target.credentials { - credentials_by_username - .insert(user.username.clone(), (user.password.clone(), user.token.clone())); + credentials_by_username.insert(user.username.clone(), (user.password.clone(), user.token.clone())); } } } @@ -425,8 +433,7 @@ fn has_unresolved_redacted_setup_values(app_config: &AppConfigDto) -> bool { app_config.api_proxy.as_ref().is_some_and(|api_proxy| { api_proxy.user.iter().any(|target| { target.credentials.iter().any(|user| { - is_setup_redacted_value(&user.password) - || user.token.as_deref().is_some_and(is_setup_redacted_value) + is_setup_redacted_value(&user.password) || user.token.as_deref().is_some_and(is_setup_redacted_value) }) }) }) @@ -448,9 +455,7 @@ fn setup_templates_to_persist(app_config: &AppConfigDto) -> Option impl IntoResponse + Send { - StatusCode::NOT_FOUND.into_response() -} +async fn api_not_found() -> impl IntoResponse + Send { StatusCode::NOT_FOUND.into_response() } async fn persist_yaml_file(file_path: &FsPath, payload: &T) -> Result<(), String> { let mut content = String::new(); @@ -781,11 +780,7 @@ async fn setup_complete_inner( match prepare_setup_validation_templates(&req.app_config, template_file_path.to_string_lossy().as_ref()) { Ok(templates) => templates, Err(err) => { - return ( - StatusCode::BAD_REQUEST, - axum::Json(json!({ "error": err.to_string() })), - ) - .into_response(); + return (StatusCode::BAD_REQUEST, axum::Json(json!({ "error": err.to_string() }))).into_response(); } }; @@ -811,12 +806,8 @@ async fn setup_complete_inner( }; let template_definition_to_persist = setup_templates_to_persist(&req.app_config); - let mut persist_paths = vec![ - &state.config_file_path, - &state.source_file_path, - &state.api_proxy_file_path, - &state.user_file_path, - ]; + let mut persist_paths = + vec![&state.config_file_path, &state.source_file_path, &state.api_proxy_file_path, &state.user_file_path]; if template_definition_to_persist.is_some() { persist_paths.push(&template_file_path); } @@ -824,7 +815,9 @@ async fn setup_complete_inner( if let Err(err) = ensure_parent_dir(file_path).await { return ( StatusCode::INTERNAL_SERVER_ERROR, - axum::Json(json!({ "error": format!("Failed to create parent directory for {}: {err}", file_path.display()) })), + axum::Json( + json!({ "error": format!("Failed to create parent directory for {}: {err}", file_path.display()) }), + ), ) .into_response(); } @@ -847,7 +840,9 @@ async fn setup_complete_inner( }; return ( status_code, - axum::Json(json!({ "error": format!("Failed to hash password for user '{username}': {err_message}") })), + axum::Json( + json!({ "error": format!("Failed to hash password for user '{username}': {err_message}") }), + ), ) .into_response(); } @@ -1089,12 +1084,8 @@ mod tests { fn setup_redaction_masks_web_auth_and_api_proxy_credentials() { let redacted = redact_app_config_for_setup(sample_app_config()); - let auth_secret = redacted - .config - .web_ui - .as_ref() - .and_then(|web_ui| web_ui.auth.as_ref()) - .map(|auth| auth.secret.as_str()); + let auth_secret = + redacted.config.web_ui.as_ref().and_then(|web_ui| web_ui.auth.as_ref()).map(|auth| auth.secret.as_str()); assert_eq!(auth_secret, Some(SETUP_REDACTED_SECRET_VALUE)); let creds = &redacted.api_proxy.expect("api_proxy should be present").user[0].credentials[0]; @@ -1109,12 +1100,8 @@ mod tests { restore_redacted_setup_values(&mut submitted, &draft); - let auth_secret = submitted - .config - .web_ui - .as_ref() - .and_then(|web_ui| web_ui.auth.as_ref()) - .map(|auth| auth.secret.as_str()); + let auth_secret = + submitted.config.web_ui.as_ref().and_then(|web_ui| web_ui.auth.as_ref()).map(|auth| auth.secret.as_str()); assert_eq!(auth_secret, Some("very-secret-value")); let creds = &submitted.api_proxy.as_ref().expect("api_proxy should be present").user[0].credentials[0]; diff --git a/backend/src/api/sys_usage.rs b/backend/src/api/sys_usage.rs index 3e88e6a80..74f0080c8 100644 --- a/backend/src/api/sys_usage.rs +++ b/backend/src/api/sys_usage.rs @@ -1,10 +1,7 @@ -use std::sync::Arc; -use sysinfo::{ - MemoryRefreshKind, Pid, ProcessRefreshKind, ProcessesToUpdate, RefreshKind, System -}; +use crate::api::model::AppState; use shared::model::SystemInfo; -use crate::api::model::{AppState}; - +use std::sync::Arc; +use sysinfo::{MemoryRefreshKind, Pid, ProcessRefreshKind, ProcessesToUpdate, RefreshKind, System}; pub fn exec_system_usage(app_state: &Arc) -> tokio::task::JoinHandle<()> { let state = Arc::clone(app_state); @@ -14,11 +11,7 @@ pub fn exec_system_usage(app_state: &Arc) -> tokio::task::JoinHandle<( let refresh_kind = RefreshKind::nothing() .with_memory(MemoryRefreshKind::nothing().with_ram()) - .with_processes( - ProcessRefreshKind::nothing() - .with_cpu() - .with_memory() - ); + .with_processes(ProcessRefreshKind::nothing().with_cpu().with_memory()); let mut sys = System::new_with_specifics(refresh_kind); loop { @@ -37,4 +30,4 @@ pub fn exec_system_usage(app_state: &Arc) -> tokio::task::JoinHandle<( tokio::time::sleep(std::time::Duration::from_secs(1)).await; } }) -} \ No newline at end of file +} diff --git a/backend/src/main.rs b/backend/src/main.rs index a5b668e1b..ed9cedb0e 100644 --- a/backend/src/main.rs +++ b/backend/src/main.rs @@ -348,7 +348,7 @@ async fn start_in_cli_mode(cfg: Arc, targets: Arc) { reqwest::Client::new() }); // In CLI mode, we don't start background managers for events or providers - exec_processing(&client, cfg, targets, None, None, None, None, None, None, None, None).await; + exec_processing(&client, cfg, targets, None, None, None, None, None, None, None, None, None).await; } async fn start_in_server_mode(cfg: Arc, targets: Arc) { diff --git a/backend/src/processing/processor/playlist.rs b/backend/src/processing/processor/playlist.rs index d73f9d6de..f8f861cef 100644 --- a/backend/src/processing/processor/playlist.rs +++ b/backend/src/processing/processor/playlist.rs @@ -1,57 +1,61 @@ -use crate::model::{ - AppConfig, ConfigFavourites, ConfigInput, ConfigRename, MessageContent, ReverseProxyDisabledHeaderConfig, TVGuide, +use crate::{ + api::{ + model::{ + ActiveProviderManager, AppState, EventManager, EventMessage, MetadataUpdateManager, PlaylistStorageState, + ProviderIdType, ResolveReason, ResolveReasonSet, UpdateGuard, UpdateTask, + }, + sync_panel_api_exp_dates, + }, + messaging::send_message, + model::{ + AppConfig, ConfigFavourites, ConfigInput, ConfigInputFlags, ConfigInputOptions, ConfigRename, ConfigTarget, + FetchedPlaylist, Mapping, MessageContent, ProcessTargets, ReverseProxyDisabledHeaderConfig, TVGuide, + }, + processing::{ + input_cache, + input_cache::ClusterState, + parser::xmltv::flatten_tvguide, + playlist_watch::process_group_watch, + processor::{ + epg::process_playlist_epg, library, sort::sort_playlist, trakt::process_trakt_categories_for_target, + xtream_series::playlist_resolve_series, xtream_vod::playlist_resolve_vod, + }, + }, + repository::{ + load_input_playlist, persist_input_playlist, persist_playlist, CategoryKey, MemoryPlaylistSource, + PlaylistSource, + }, + utils::{ + debug_if_enabled, epg, log_memory_snapshot, m3u, trace_if_enabled, xtream, StepMeasure, StepMeasureCallback, + }, }; -use crate::utils::xtream; -use crate::utils::{epg, StepMeasureCallback}; -use crate::utils::{log_memory_snapshot, m3u}; -use std::collections::{HashMap, HashSet}; -use std::path::PathBuf; -use std::sync::{Arc, Weak}; -use tokio::sync::{Mutex, OwnedRwLockWriteGuard, RwLock}; -use tokio::task::JoinSet; - -use crate::api::model::{ - ActiveProviderManager, EventManager, EventMessage, MetadataUpdateManager, PlaylistStorageState, ProviderIdType, - ResolveReason, ResolveReasonSet, UpdateGuard, UpdateTask, -}; - -use crate::messaging::send_message; - -use crate::model::FetchedPlaylist; -use crate::model::Mapping; -use crate::model::{ConfigInputFlags, ConfigInputOptions, ConfigTarget, ProcessTargets}; -use crate::processing::input_cache; -use crate::processing::input_cache::ClusterState; -use crate::processing::parser::xmltv::flatten_tvguide; -use crate::processing::playlist_watch::process_group_watch; -use crate::processing::processor::epg::process_playlist_epg; -use crate::processing::processor::library; -use crate::processing::processor::sort::sort_playlist; -use crate::processing::processor::trakt::process_trakt_categories_for_target; -use crate::processing::processor::xtream_series::playlist_resolve_series; -use crate::processing::processor::xtream_vod::playlist_resolve_vod; -use crate::repository::{load_input_playlist, persist_input_playlist, persist_playlist}; -use crate::repository::{CategoryKey, MemoryPlaylistSource, PlaylistSource}; -use crate::utils::StepMeasure; -use crate::utils::{debug_if_enabled, trace_if_enabled}; use futures::{FutureExt, StreamExt}; use indexmap::IndexMap; use log::{debug, error, info, log_enabled, warn, Level}; -use shared::concat_string; -use shared::error::{get_errors_notify_message, notify_err, TuliproxError}; -use shared::foundation::Filter; -use shared::foundation::{get_field_value, set_field_value, ValueAccessor, ValueProvider}; -use shared::model::xtream_const::XTREAM_CLUSTER; -use shared::model::{ - CounterModifier, FieldGetAccessor, FieldSetAccessor, InputType, ItemField, PlaylistGroup, PlaylistItem, - PlaylistItemType, ProcessingOrder, StreamProperties, XtreamCluster, +use shared::{ + concat_string, + error::{get_errors_notify_message, notify_err, TuliproxError}, + foundation::{get_field_value, set_field_value, Filter, ValueAccessor, ValueProvider}, + model::{ + xtream_const::XTREAM_CLUSTER, CounterModifier, FieldGetAccessor, FieldSetAccessor, InputStats, InputType, + ItemField, PlaylistGroup, PlaylistItem, PlaylistItemType, PlaylistStats, ProcessingOrder, SourceStats, + StreamProperties, TargetStats, UUIDType, XtreamCluster, + }, + utils::{ + create_alias_uuid, default_as_default, default_probe_delay_secs, default_probe_live_interval, interner_gc, + Internable, + }, }; -use shared::model::{InputStats, PlaylistStats, SourceStats, TargetStats, UUIDType}; -use shared::utils::{ - create_alias_uuid, default_as_default, default_probe_delay_secs, default_probe_live_interval, interner_gc, - Internable, +use std::{ + collections::{HashMap, HashSet}, + path::PathBuf, + sync::{Arc, Weak}, + time::Instant, +}; +use tokio::{ + sync::{Mutex, OwnedRwLockWriteGuard, RwLock}, + task::JoinSet, }; -use std::time::Instant; fn is_valid(pli: &PlaylistItem, filter: &Filter, match_as_ascii: bool) -> bool { let provider = ValueProvider { pli, match_as_ascii }; @@ -1206,6 +1210,7 @@ pub async fn exec_processing( app_config: Arc, targets: Arc, event_manager: Option>, + app_state: Option>, playlist_state: Option>, update_guard: Option, disabled_headers: Option, @@ -1230,6 +1235,12 @@ pub async fn exec_processing( None }; + if playlist_guard.is_some() { + if let Some(state) = app_state.as_ref() { + sync_panel_api_exp_dates(state).await; + } + } + // Pause background metadata/probe tasks for the full update lifecycle. let _background_pause_guard = if let Some(manager) = metadata_manager.as_ref() { Some(manager.acquire_update_pause_guard().await) diff --git a/backend/src/processing/processor/stream_probe.rs b/backend/src/processing/processor/stream_probe.rs index 662cde188..272517488 100644 --- a/backend/src/processing/processor/stream_probe.rs +++ b/backend/src/processing/processor/stream_probe.rs @@ -121,6 +121,7 @@ pub async fn update_generic_stream_metadata( analyze_duration, probe_size, ffprobe_timeout, + config.proxy.as_ref(), ).await; if let Some(handle) = acquired_handle { diff --git a/backend/src/processing/processor/xtream.rs b/backend/src/processing/processor/xtream.rs index bb7579d7d..06ceb411e 100644 --- a/backend/src/processing/processor/xtream.rs +++ b/backend/src/processing/processor/xtream.rs @@ -139,6 +139,7 @@ pub async fn update_live_stream_metadata( analyze_duration, probe_size, ffprobe_timeout, + config.proxy.as_ref(), ) .await { diff --git a/backend/src/processing/processor/xtream_series.rs b/backend/src/processing/processor/xtream_series.rs index aa361c3ae..e29a85fd3 100644 --- a/backend/src/processing/processor/xtream_series.rs +++ b/backend/src/processing/processor/xtream_series.rs @@ -869,6 +869,7 @@ pub async fn update_series_metadata( probe_settings.analyze_duration_micros, probe_settings.probe_size_bytes, probe_settings.timeout_secs, + config.proxy.as_ref(), ) .await { diff --git a/backend/src/processing/processor/xtream_vod.rs b/backend/src/processing/processor/xtream_vod.rs index c485bd427..fcf1a3934 100644 --- a/backend/src/processing/processor/xtream_vod.rs +++ b/backend/src/processing/processor/xtream_vod.rs @@ -794,6 +794,7 @@ pub async fn update_vod_metadata( analyze_duration, probe_size, ffprobe_timeout, + config.proxy.as_ref(), ) .await { diff --git a/backend/src/utils/ffmpeg.rs b/backend/src/utils/ffmpeg.rs index 88eb6cf6e..11678d70b 100644 --- a/backend/src/utils/ffmpeg.rs +++ b/backend/src/utils/ffmpeg.rs @@ -1,9 +1,11 @@ +use crate::model::ProxyConfig; use log::{debug, warn}; -use tokio::process::Command; -use std::time::Duration; use serde_json::Value; use shared::model::MediaQuality; use shared::utils::sanitize_sensitive_info; +use std::time::Duration; +use tokio::process::Command; +use url::Url; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ProbeFailureKind { @@ -24,6 +26,44 @@ pub async fn check_ffprobe_availability() -> bool { } } +fn build_ffprobe_proxy_url(proxy_cfg: &ProxyConfig) -> Option { + let mut proxy_url = Url::parse(proxy_cfg.url.as_str()).ok()?; + if let Some(username) = proxy_cfg.username.as_deref() { + let _ = proxy_url.set_username(username); + if let Some(password) = proxy_cfg.password.as_deref() { + let _ = proxy_url.set_password(Some(password)); + } + } + Some(proxy_url.to_string()) +} + +fn apply_proxy_to_ffprobe(command: &mut Command, proxy_cfg: Option<&ProxyConfig>) { + let Some(proxy_cfg) = proxy_cfg else { + return; + }; + + let Some(proxy_url) = build_ffprobe_proxy_url(proxy_cfg) else { + warn!( + "Ignoring invalid ffprobe proxy URL: {}", + sanitize_sensitive_info(proxy_cfg.url.as_str()) + ); + return; + }; + + // ffprobe is an external process and does not consume the app's reqwest proxy config. + // Export proxy env vars explicitly so all probe requests honor the configured upstream proxy. + for key in [ + "HTTP_PROXY", + "HTTPS_PROXY", + "ALL_PROXY", + "http_proxy", + "https_proxy", + "all_proxy", + ] { + command.env(key, proxy_url.as_str()); + } +} + fn is_not_found_probe_error(stderr: &str) -> bool { let normalized = stderr.to_ascii_lowercase(); normalized.contains("404") || normalized.contains("not found") @@ -35,6 +75,7 @@ pub async fn probe_url( analyze_duration: u64, probe_size: u64, timeout_secs: u64, + proxy_cfg: Option<&ProxyConfig>, ) -> ProbeUrlOutcome { // Determine timeout: Ensure it's at least as long as the analyze duration + buffer, // but respect the user setting if it's longer. @@ -54,6 +95,8 @@ pub async fn probe_url( // Optimization for network streams .arg("-analyzeduration").arg(analyze_duration.to_string()) .arg("-probesize").arg(probe_size.to_string()); + + apply_proxy_to_ffprobe(&mut command, proxy_cfg); if let Some(ua) = user_agent { command.arg("-user_agent").arg(ua); @@ -137,3 +180,31 @@ pub async fn probe_url( ProbeUrlOutcome::Failed(ProbeFailureKind::Other) } + +#[cfg(test)] +mod tests { + use super::build_ffprobe_proxy_url; + use crate::model::ProxyConfig; + + #[test] + fn build_ffprobe_proxy_url_injects_credentials() { + let proxy_cfg = ProxyConfig { + url: "http://proxy.local:8080".to_string(), + username: Some("alice".to_string()), + password: Some("secret".to_string()), + }; + let resolved = build_ffprobe_proxy_url(&proxy_cfg).expect("proxy url should parse"); + assert!(resolved.contains("alice:secret@proxy.local:8080")); + } + + #[test] + fn build_ffprobe_proxy_url_keeps_existing_inline_credentials() { + let proxy_cfg = ProxyConfig { + url: "socks5://bob:pass@proxy.local:1080".to_string(), + username: None, + password: None, + }; + let resolved = build_ffprobe_proxy_url(&proxy_cfg).expect("proxy url should parse"); + assert!(resolved.contains("bob:pass@proxy.local:1080")); + } +} diff --git a/backend/src/utils/network/request.rs b/backend/src/utils/network/request.rs index faf0c6e57..7a685988d 100644 --- a/backend/src/utils/network/request.rs +++ b/backend/src/utils/network/request.rs @@ -1,31 +1,44 @@ -use crate::api::model::persist_pipe_stream::tee_dyn_reader; -use crate::api::model::{AppState, STREAM_IDLE_TIMEOUT}; -use crate::model::{ - resolve_provider_scheme_url_with_provider, AppConfig, Config, ConfigInput, ConfigProvider, - InputSource, ResourceRetryConfig, ReverseProxyDisabledHeaderConfig, +use crate::{ + api::model::{persist_pipe_stream::tee_dyn_reader, AppState, STREAM_IDLE_TIMEOUT}, + model::{ + resolve_provider_scheme_url_with_provider, AppConfig, Config, ConfigInput, ConfigProvider, InputSource, + ResourceRetryConfig, ReverseProxyDisabledHeaderConfig, + }, + utils::{ + async_file_reader, async_file_writer, + compression::compression_utils::{is_deflate, is_gzip}, + debug_if_enabled, get_file_path, persist_file, + }, }; -use crate::utils::compression::compression_utils::{is_deflate, is_gzip}; -use crate::utils::{async_file_reader, async_file_writer, debug_if_enabled}; -use crate::utils::{get_file_path, persist_file}; use futures::{StreamExt, TryStreamExt}; use log::{debug, error, log_enabled, trace, warn, Level}; -use reqwest::header::CONTENT_ENCODING; -use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; -use reqwest::redirect::Policy; -use reqwest::StatusCode; -use shared::error::{notify_err_res, string_to_io_error, TuliproxError}; -use shared::model::{format_elapsed_time, InputFetchMethod, DEFAULT_USER_AGENT}; -use shared::utils::{filter_request_header, human_readable_byte_size, sanitize_sensitive_info, CONTENT_TYPE_JSON, ENCODING_DEFLATE, ENCODING_GZIP}; -use std::collections::HashMap; -use std::io::{Error, ErrorKind}; -use std::path::{Path, PathBuf}; -use std::pin::Pin; -use std::sync::{Arc, Once}; -use std::time::Duration; use regex::Regex; -use tokio::fs::File; -use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncReadExt, AsyncWriteExt}; -use tokio::time::sleep; +use reqwest::{ + header::{HeaderMap, HeaderName, HeaderValue, CONTENT_ENCODING}, + redirect::Policy, + StatusCode, +}; +use shared::{ + error::{notify_err_res, string_to_io_error, TuliproxError}, + model::{format_elapsed_time, InputFetchMethod, DEFAULT_USER_AGENT}, + utils::{ + filter_request_header, human_readable_byte_size, sanitize_sensitive_info, CONTENT_TYPE_JSON, ENCODING_DEFLATE, + ENCODING_GZIP, + }, +}; +use std::{ + collections::HashMap, + io::{Error, ErrorKind}, + path::{Path, PathBuf}, + pin::Pin, + sync::{Arc, Once}, + time::Duration, +}; +use tokio::{ + fs::File, + io::{AsyncBufReadExt, AsyncRead, AsyncReadExt, AsyncWriteExt}, + time::sleep, +}; use tokio_util::io::StreamReader; use url::Url; @@ -92,11 +105,9 @@ pub enum MimeCategory { } pub fn classify_content_type(headers: &[(String, String)]) -> MimeCategory { - headers.iter() - .find_map(|(k, v)| { - (k == axum::http::header::CONTENT_TYPE.as_str()).then_some(v) - }) - .map_or(MimeCategory::Unknown, |v| match v.to_lowercase().as_str() { + headers.iter().find_map(|(k, v)| (k == axum::http::header::CONTENT_TYPE.as_str()).then_some(v)).map_or( + MimeCategory::Unknown, + |v| match v.to_lowercase().as_str() { v if v.starts_with("video/") || v == "application/octet-stream" => MimeCategory::Video, v if v.contains("mpegurl") => MimeCategory::M3U8, v if v.starts_with("image/") => MimeCategory::Image, @@ -104,7 +115,8 @@ pub fn classify_content_type(headers: &[(String, String)]) -> MimeCategory { v if v.starts_with("application/xml") || v.ends_with("+xml") || v == "text/xml" => MimeCategory::Xml, v if v.starts_with("text/") => MimeCategory::Text, _ => MimeCategory::Unclassified, - }) + }, + ) } pub fn format_http_status(status: StatusCode) -> String { @@ -127,10 +139,7 @@ pub fn content_type_from_ext(ext: &str) -> &'static str { } } -fn resolve_provider_url_for_attempt( - url: &Url, - provider: Option<&Arc>, -) -> Url { +fn resolve_provider_url_for_attempt(url: &Url, provider: Option<&Arc>) -> Url { let Some(provider) = provider else { return url.clone(); }; @@ -149,7 +158,6 @@ fn resolve_provider_url_for_attempt( } } - #[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss, clippy::cast_precision_loss)] pub fn calculate_retry_backoff(base_delay_ms: u64, multiplier: f64, attempt: u32) -> u64 { let base = base_delay_ms.max(1); @@ -176,19 +184,16 @@ pub async fn send_with_retry_and_provider( mut send: impl FnMut(&Url) -> reqwest::RequestBuilder, ) -> Result { let config = app_config.config.load(); - let (max_attempts, backoff_ms, backoff_multiplier, failover_patterns) = config - .reverse_proxy - .as_ref() - .map_or_else( - || { - let (a, b, c) = ResourceRetryConfig::get_default_retry_values(); - (a, b, c, ResourceRetryConfig::default().failover_redirect_patterns) - }, - |rp| { - let (a, b, c) = rp.resource_retry.get_retry_values(); - (a, b, c, rp.resource_retry.failover_redirect_patterns.clone()) - }, - ); + let (max_attempts, backoff_ms, backoff_multiplier, failover_patterns) = config.reverse_proxy.as_ref().map_or_else( + || { + let (a, b, c) = ResourceRetryConfig::get_default_retry_values(); + (a, b, c, ResourceRetryConfig::default().failover_redirect_patterns) + }, + |rp| { + let (a, b, c) = rp.resource_retry.get_retry_values(); + (a, b, c, rp.resource_retry.failover_redirect_patterns.clone()) + }, + ); drop(config); let idle_timeout = Duration::from_secs(STREAM_IDLE_TIMEOUT); @@ -197,13 +202,11 @@ pub async fn send_with_retry_and_provider( // Record the starting URL index for full-cycle detection. // This allows us to try all URLs even when starting from a non-zero index. - let (start_index, max_provider_attempts) = provider - .as_ref() - .map_or((0, 0), |p| (p.get_current_index(), p.urls.len())); + let (start_index, max_provider_attempts) = + provider.as_ref().map_or((0, 0), |p| (p.get_current_index(), p.urls.len())); let mut provider_attempts = usize::from(max_provider_attempts > 0); 'provider_loop: loop { - // 2. Retry loop for the current URL for attempt in 0..max_attempts { let resolved_url = resolve_provider_url_for_attempt(url, provider); @@ -330,10 +333,12 @@ fn is_failover_redirect(url: &Url, patterns: &[Arc]) -> bool { /// Helper to handle sleep duration for retries, respecting Retry-After headers async fn perform_backoff(attempt: u32, ms: u64, mult: f64, response: &reqwest::Response) { - let wait_dur = response.headers() + let wait_dur = response + .headers() .get(reqwest::header::RETRY_AFTER) .and_then(|h| h.to_str().ok()) - .and_then(|s| s.parse::().ok()).map_or_else(|| Duration::from_millis(calculate_retry_backoff(ms, mult, attempt)), Duration::from_secs); + .and_then(|s| s.parse::().ok()) + .map_or_else(|| Duration::from_millis(calculate_retry_backoff(ms, mult, attempt)), Duration::from_secs); tokio::time::sleep(wait_dur).await; } @@ -353,19 +358,15 @@ pub async fn get_input_epg_content_as_file( sanitize_sensitive_info(url_str) ); if url_str.parse::().is_ok() { - match download_epg_content_as_file( - app_config, - client, - input, - headers, - url_str, - persist_filepath, - ) - .await - { + match download_epg_content_as_file(app_config, client, input, headers, url_str, persist_filepath).await { Ok(content) => Ok(content), Err(e) => { - error!("can't download input {} epg url: {} => {}", input.name, sanitize_sensitive_info(url_str), sanitize_sensitive_info(e.to_string().as_str())); + error!( + "can't download input {} epg url: {} => {}", + input.name, + sanitize_sensitive_info(url_str), + sanitize_sensitive_info(e.to_string().as_str()) + ); notify_err_res!("Failed to download") } } @@ -389,11 +390,14 @@ pub async fn get_input_epg_content_as_file( None => None, }; - result.map_or_else(|| { - let msg = format!("can't read input url: {}", sanitize_sensitive_info(url_str)); - error!("{msg}"); - notify_err_res!("{msg}") - }, Ok) + result.map_or_else( + || { + let msg = format!("can't read input url: {}", sanitize_sensitive_info(url_str)); + error!("{msg}"); + notify_err_res!("{msg}") + }, + Ok, + ) } } @@ -411,19 +415,14 @@ pub async fn get_input_text_content( ); if input.url.parse::().is_ok() { - match download_text_content( - &app_state.app_config, - client, - input, - None, - persist_filepath, - false, - ) - .await - { + match download_text_content(&app_state.app_config, client, input, None, persist_filepath, false).await { Ok((content, _response_url)) => Ok(content), Err(e) => { - error!("Failed to download input '{}': {}", &input.name, sanitize_sensitive_info(e.to_string().as_str())); + error!( + "Failed to download input '{}': {}", + &input.name, + sanitize_sensitive_info(e.to_string().as_str()) + ); notify_err_res!("Failed to download") } } @@ -451,11 +450,14 @@ pub async fn get_input_text_content( } None => None, }; - result.map_or_else(|| { - let msg = format!("can't read input url: {}", sanitize_sensitive_info(&input.url)); - error!("{msg}"); - notify_err_res!("{msg}") - }, Ok) + result.map_or_else( + || { + let msg = format!("can't read input url: {}", sanitize_sensitive_info(&input.url)); + error!("{msg}"); + notify_err_res!("{msg}") + }, + Ok, + ) } } @@ -473,14 +475,7 @@ pub async fn get_input_text_content_as_stream( ); if input.url.parse::().is_ok() { - match download_text_content_as_stream( - app_config, - client, - input, - persist_filepath, - ) - .await - { + match download_text_content_as_stream(app_config, client, input, persist_filepath).await { Ok((content, _response_url)) => Ok(content), Err(e) => { error!( @@ -502,13 +497,10 @@ pub async fn get_input_text_content_as_stream( content, &path, Some(Arc::new(|size| { - debug_if_enabled!( - "Persisted {} bytes", - human_readable_byte_size(size as u64) - ); + debug_if_enabled!("Persisted {} bytes", human_readable_byte_size(size as u64)); })), ) - .await; + .await; Some(tee) } else { Some(content) @@ -524,11 +516,14 @@ pub async fn get_input_text_content_as_stream( } None => None, }; - result.map_or_else(|| { - let msg = format!("can't read input url: {}", sanitize_sensitive_info(&input.url)); - error!("{msg}"); - notify_err_res!("{msg}") - }, Ok) + result.map_or_else( + || { + let msg = format!("can't read input url: {}", sanitize_sensitive_info(&input.url)); + error!("{msg}"); + notify_err_res!("{msg}") + }, + Ok, + ) } } @@ -553,12 +548,7 @@ pub fn get_client_request( client.post(url.clone()).form(¶ms) } }; - let headers = get_request_headers( - headers, - custom_headers, - disabled_headers, - default_user_agent, - ); + let headers = get_request_headers(headers, custom_headers, disabled_headers, default_user_agent); request.headers(headers) } @@ -575,15 +565,11 @@ pub fn get_request_headers( // These should have the highest priority. if let Some(req_headers) = request_headers { for (key, value) in req_headers { - if let (Ok(key), Ok(value)) = ( - HeaderName::from_bytes(key.as_bytes()), - HeaderValue::from_bytes(value.as_bytes()), - ) { + if let (Ok(key), Ok(value)) = + (HeaderName::from_bytes(key.as_bytes()), HeaderValue::from_bytes(value.as_bytes())) + { if filter_request_header(key.as_str()) { - if disabled_headers - .as_ref() - .is_some_and(|d| d.should_remove(key.as_str())) - { + if disabled_headers.as_ref().is_some_and(|d| d.should_remove(key.as_str())) { continue; } if key == axum::http::header::USER_AGENT { @@ -601,16 +587,10 @@ pub fn get_request_headers( for (key, value) in custom { let key_lc = key.to_lowercase(); if filter_request_header(key_lc.as_str()) { - if disabled_headers - .as_ref() - .is_some_and(|d| d.should_remove(key_lc.as_str())) - { + if disabled_headers.as_ref().is_some_and(|d| d.should_remove(key_lc.as_str())) { continue; } - if let (Ok(name), Ok(val)) = ( - HeaderName::from_bytes(key.as_bytes()), - HeaderValue::from_bytes(value), - ) { + if let (Ok(name), Ok(val)) = (HeaderName::from_bytes(key.as_bytes()), HeaderValue::from_bytes(value)) { // Only insert if not already present (config takes precedence) if !headers.contains_key(&name) { if name == axum::http::header::USER_AGENT { @@ -624,15 +604,8 @@ pub fn get_request_headers( } if log_enabled!(Level::Trace) { - let he: HashMap = headers - .iter() - .map(|(k, v)| { - ( - k.to_string(), - String::from_utf8_lossy(v.as_bytes()).to_string(), - ) - }) - .collect(); + let he: HashMap = + headers.iter().map(|(k, v)| (k.to_string(), String::from_utf8_lossy(v.as_bytes()).to_string())).collect(); if !he.is_empty() { trace!("Request headers {he:?}"); } @@ -661,10 +634,7 @@ pub fn get_request_headers( pub async fn get_local_file_content(file_path: &Path) -> Result { // open file let file = File::open(file_path).await.map_err(|err| { - std::io::Error::new( - ErrorKind::NotFound, - format!("Failed to open file: {}, {err:?}", file_path.display()), - ) + std::io::Error::new(ErrorKind::NotFound, format!("Failed to open file: {}, {err:?}", file_path.display())) })?; let mut buf_reader = async_file_reader(file); @@ -693,15 +663,10 @@ pub async fn get_local_file_content(file_path: &Path) -> Result Result { +pub async fn get_local_file_content_as_stream(file_path: &Path) -> Result { // open file let file = File::open(file_path).await.map_err(|err| { - std::io::Error::new( - ErrorKind::NotFound, - format!("Failed to open file: {}, {err:?}", file_path.display()), - ) + std::io::Error::new(ErrorKind::NotFound, format!("Failed to open file: {}, {err:?}", file_path.display())) })?; let mut buf_reader = async_file_reader(file); @@ -712,9 +677,7 @@ pub async fn get_local_file_content_as_stream( if is_gzipped { // use Async Gzip Decoder - Ok(Box::pin( - async_compression::tokio::bufread::GzipDecoder::new(buf_reader), - )) + Ok(Box::pin(async_compression::tokio::bufread::GzipDecoder::new(buf_reader))) } else { Ok(Box::pin(buf_reader)) } @@ -728,11 +691,8 @@ pub async fn get_remote_content_as_file( url: &Url, file_path: &Path, ) -> Result { - let custom_headers = headers.map(|h| { - h.iter() - .map(|(k, v)| (k.as_str().to_string(), v.as_bytes().to_vec())) - .collect::>() - }); + let custom_headers = headers + .map(|h| h.iter().map(|(k, v)| (k.as_str().to_string(), v.as_bytes().to_vec())).collect::>()); let config = app_config.config.load(); let default_user_agent = config.default_user_agent.clone(); @@ -740,33 +700,25 @@ pub async fn get_remote_content_as_file( let provider_config = input.get_resolve_provider(url.as_str()); - let response = send_with_retry_and_provider( - app_config, - url, - provider_config.as_ref(), - false, - |resolved_url| { - get_client_request( - client, - input.method, - Some(&input.headers), - resolved_url, - custom_headers.as_ref(), - None, - default_user_agent.as_deref(), - ) - }, - ) - .await?; + let response = send_with_retry_and_provider(app_config, url, provider_config.as_ref(), false, |resolved_url| { + get_client_request( + client, + input.method, + Some(&input.headers), + resolved_url, + custom_headers.as_ref(), + None, + default_user_agent.as_deref(), + ) + }) + .await?; let start_time = tokio::time::Instant::now(); let mut writer = async_file_writer(File::create(file_path).await?); let mut stream = response.bytes_stream(); while let Some(chunk) = stream.next().await { - let bytes = chunk.map_err(|e| { - string_to_io_error(format!("Failed to read chunk: {e}")) - })?; + let bytes = chunk.map_err(|e| string_to_io_error(format!("Failed to read chunk: {e}")))?; writer.write_all(&bytes).await?; } @@ -813,17 +765,12 @@ pub async fn get_remote_content_as_file( pub type DynReader = Pin>; -async fn build_decoded_stream_reader( - response: reqwest::Response, -) -> Result { +async fn build_decoded_stream_reader(response: reqwest::Response) -> Result { let headers = response.headers(); let header_value = headers.get(CONTENT_ENCODING); - let mut encoding = header_value - .and_then(|h| h.to_str().ok()) - .map(ToString::to_string); + let mut encoding = header_value.and_then(|h| h.to_str().ok()).map(ToString::to_string); - let stream_reader = - StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); + let stream_reader = StreamReader::new(response.bytes_stream().map_err(std::io::Error::other)); let mut buf_reader = async_file_reader(stream_reader); let peek = buf_reader.fill_buf().await?; @@ -836,20 +783,10 @@ async fn build_decoded_stream_reader( } } - let reader: DynReader = if encoding - .as_ref() - .is_some_and(|e| e.eq_ignore_ascii_case(ENCODING_GZIP)) - { - Box::pin(async_compression::tokio::bufread::GzipDecoder::new( - buf_reader, - )) - } else if encoding - .as_ref() - .is_some_and(|e| e.eq_ignore_ascii_case(ENCODING_DEFLATE)) - { - Box::pin(async_compression::tokio::bufread::ZlibDecoder::new( - buf_reader, - )) + let reader: DynReader = if encoding.as_ref().is_some_and(|e| e.eq_ignore_ascii_case(ENCODING_GZIP)) { + Box::pin(async_compression::tokio::bufread::GzipDecoder::new(buf_reader)) + } else if encoding.as_ref().is_some_and(|e| e.eq_ignore_ascii_case(ENCODING_DEFLATE)) { + Box::pin(async_compression::tokio::bufread::ZlibDecoder::new(buf_reader)) } else { Box::pin(buf_reader) }; @@ -857,7 +794,6 @@ async fn build_decoded_stream_reader( Ok(reader) } - #[allow(clippy::implicit_hasher)] pub async fn get_remote_content_as_stream( app_config: &Arc, @@ -866,11 +802,8 @@ pub async fn get_remote_content_as_stream( headers: Option<&HeaderMap>, url: &Url, ) -> Result<(DynReader, String), Error> { - let custom_headers = headers.map(|h| { - h.iter() - .map(|(k, v)| (k.as_str().to_string(), v.as_bytes().to_vec())) - .collect::>() - }); + let custom_headers = headers + .map(|h| h.iter().map(|(k, v)| (k.as_str().to_string(), v.as_bytes().to_vec())).collect::>()); let config = app_config.config.load(); let default_user_agent = config.default_user_agent.clone(); @@ -886,32 +819,21 @@ pub async fn get_remote_content_as_stream( let headers: HashMap = merged .iter() - .map(|(k, v)| { - ( - k.as_str().to_string(), - String::from_utf8_lossy(v.as_bytes()).to_string(), - ) - }) + .map(|(k, v)| (k.as_str().to_string(), String::from_utf8_lossy(v.as_bytes()).to_string())) .collect(); - let response = send_with_retry_and_provider( - app_config, - url, - input.get_provider(), - false, - |resolved_url| { - get_client_request( - client, - input.method, - Some(&headers), - resolved_url, - None, - None, - default_user_agent.as_deref(), - ) - }, - ) - .await?; + let response = send_with_retry_and_provider(app_config, url, input.get_provider(), false, |resolved_url| { + get_client_request( + client, + input.method, + Some(&headers), + resolved_url, + None, + None, + default_user_agent.as_deref(), + ) + }) + .await?; let response_url = response.url().to_string(); @@ -926,13 +848,7 @@ async fn get_remote_content( headers: Option<&HeaderMap>, url: &Url, ) -> Result<(String, String), Error> { - let (mut stream, response_url) = get_remote_content_as_stream( - app_config, - client, - input, - headers, - url, - ) + let (mut stream, response_url) = get_remote_content_as_stream(app_config, client, input, headers, url) .await .map_err(|e| string_to_io_error(format!("Failed to read content: {e}")))?; let mut content = String::new(); @@ -943,6 +859,115 @@ async fn get_remote_content( Ok((content, response_url)) } +async fn get_remote_content_with_manual_redirects( + app_config: &Arc, + client: &reqwest::Client, + input: &InputSource, + headers: Option<&HeaderMap>, + url: &Url, + max_redirects: usize, +) -> Result<(String, String), Error> { + let custom_headers = headers + .map(|h| h.iter().map(|(k, v)| (k.as_str().to_string(), v.as_bytes().to_vec())).collect::>()); + + let config = app_config.config.load(); + let default_user_agent = config.default_user_agent.clone(); + let disabled_headers = config.get_disabled_headers(); + drop(config); + + let merged = get_request_headers( + Some(&input.headers), + custom_headers.as_ref(), + disabled_headers.as_ref(), + default_user_agent.as_deref(), + ); + + let headers: HashMap = merged + .iter() + .map(|(k, v)| (k.as_str().to_string(), String::from_utf8_lossy(v.as_bytes()).to_string())) + .collect(); + + let mut current_url = url.clone(); + let mut current_headers = headers; + let mut remaining_redirects = max_redirects; + loop { + let response = + send_with_retry_and_provider(app_config, ¤t_url, input.get_provider(), true, |resolved_url| { + get_client_request( + client, + input.method, + Some(¤t_headers), + resolved_url, + None, + None, + default_user_agent.as_deref(), + ) + }) + .await?; + let response_base_url = response.url().clone(); + + if response.status().is_redirection() { + if remaining_redirects == 0 { + return Err(string_to_io_error(format!( + "Too many redirects while requesting {}", + sanitize_sensitive_info(url.as_str()) + ))); + } + + let Some(location) = response.headers().get(reqwest::header::LOCATION) else { + return Err(string_to_io_error(format!( + "Redirect response missing location header for {}", + sanitize_sensitive_info(current_url.as_str()) + ))); + }; + let Ok(location_str) = location.to_str() else { + return Err(string_to_io_error(format!( + "Redirect response contains invalid location header for {}", + sanitize_sensitive_info(current_url.as_str()) + ))); + }; + let next_url = + response_base_url.join(location_str).or_else(|_| Url::parse(location_str)).map_err(|_| { + string_to_io_error(format!( + "Redirect response contains invalid location URL for {}", + sanitize_sensitive_info(current_url.as_str()) + )) + })?; + + if !same_origin(&response_base_url, &next_url) { + strip_sensitive_headers_for_cross_origin_redirect(&mut current_headers); + } + current_url = next_url; + remaining_redirects = remaining_redirects.saturating_sub(1); + continue; + } + + let response_url = response.url().to_string(); + let mut stream = build_decoded_stream_reader(response).await?; + let mut content = String::new(); + stream + .read_to_string(&mut content) + .await + .map_err(|e| string_to_io_error(format!("Failed to read content: {e}")))?; + return Ok((content, response_url)); + } +} + +fn same_origin(lhs: &Url, rhs: &Url) -> bool { + lhs.scheme().eq_ignore_ascii_case(rhs.scheme()) + && lhs.host_str() == rhs.host_str() + && lhs.port_or_known_default() == rhs.port_or_known_default() +} + +fn strip_sensitive_headers_for_cross_origin_redirect(headers: &mut HashMap) { + headers.retain(|key, _| { + !key.eq_ignore_ascii_case("authorization") + && !key.eq_ignore_ascii_case("cookie") + && !key.eq_ignore_ascii_case("proxy-authorization") + && !key.eq_ignore_ascii_case("host") + }); +} + async fn download_epg_content_as_file( app_config: &Arc, client: &reqwest::Client, @@ -964,22 +989,15 @@ async fn download_epg_content_as_file( if file_path.exists() { Ok(file_path) } else { - Err(Error::new( - ErrorKind::NotFound, - format!("Unknown file {}", file_path.display()), - )) + Err(Error::new(ErrorKind::NotFound, format!("Unknown file {}", file_path.display()))) } }, ) } else { - get_remote_content_as_file(app_config, client, input, headers, &url, persist_filepath) - .await + get_remote_content_as_file(app_config, client, input, headers, &url, persist_filepath).await } } else { - Err(Error::new( - ErrorKind::Unsupported, - format!("Malformed URL {}", sanitize_sensitive_info(url_str)), - )) + Err(Error::new(ErrorKind::Unsupported, format!("Malformed URL {}", sanitize_sensitive_info(url_str)))) } } @@ -995,23 +1013,11 @@ pub async fn download_text_content( let result = if let Ok(url) = input.url.parse::() { let result = if url.scheme() == "file" { match url.to_file_path() { - Ok(file_path) => get_local_file_content(&file_path) - .await - .map(|c| (c, url.to_string())), - Err(()) => Err(string_to_io_error(format!( - "Unknown file {}", - sanitize_sensitive_info(&input.url) - ))), + Ok(file_path) => get_local_file_content(&file_path).await.map(|c| (c, url.to_string())), + Err(()) => Err(string_to_io_error(format!("Unknown file {}", sanitize_sensitive_info(&input.url)))), } } else { - get_remote_content( - app_config, - client, - input, - headers, - &url, - ) - .await + get_remote_content(app_config, client, input, headers, &url).await }; match result { Ok((content, response_url)) => { @@ -1023,17 +1029,57 @@ pub async fn download_text_content( Err(err) => Err(err), } } else { - Err(string_to_io_error(format!( - "Malformed URL {}", - sanitize_sensitive_info(&input.url) - ))) + Err(string_to_io_error(format!("Malformed URL {}", sanitize_sensitive_info(&input.url)))) }; - let level = if trace_log { - log::Level::Trace + let level = if trace_log { log::Level::Trace } else { log::Level::Debug }; + if log_enabled!(level) { + if let Ok((_content, response_url)) = result.as_ref() { + log::log!( + level, + "Request took: {} {}", + format_elapsed_time(start_time.elapsed().as_secs()), + sanitize_sensitive_info(response_url.as_str()) + ); + } + } + + result +} + +pub async fn download_text_content_with_manual_redirects( + app_config: &Arc, + client: &reqwest::Client, + input: &InputSource, + headers: Option<&HeaderMap>, + persist_filepath: Option, + trace_log: bool, + max_redirects: usize, +) -> Result<(String, String), Error> { + let start_time = tokio::time::Instant::now(); + let result = if let Ok(url) = input.url.parse::() { + let result = if url.scheme() == "file" { + match url.to_file_path() { + Ok(file_path) => get_local_file_content(&file_path).await.map(|c| (c, url.to_string())), + Err(()) => Err(string_to_io_error(format!("Unknown file {}", sanitize_sensitive_info(&input.url)))), + } + } else { + get_remote_content_with_manual_redirects(app_config, client, input, headers, &url, max_redirects).await + }; + match result { + Ok((content, response_url)) => { + if persist_filepath.is_some() { + persist_file(persist_filepath, &content).await; + } + Ok((content, response_url)) + } + Err(err) => Err(err), + } } else { - log::Level::Debug + Err(string_to_io_error(format!("Malformed URL {}", sanitize_sensitive_info(&input.url)))) }; + + let level = if trace_log { log::Level::Trace } else { log::Level::Debug }; if log_enabled!(level) { if let Ok((_content, response_url)) = result.as_ref() { log::log!( @@ -1057,23 +1103,11 @@ pub async fn download_text_content_as_stream( if let Ok(url) = input.url.parse::() { let result = if url.scheme() == "file" { match url.to_file_path() { - Ok(file_path) => get_local_file_content_as_stream(&file_path) - .await - .map(|c| (c, url.to_string())), - Err(()) => Err(string_to_io_error(format!( - "Unknown file {}", - sanitize_sensitive_info(&input.url) - ))), + Ok(file_path) => get_local_file_content_as_stream(&file_path).await.map(|c| (c, url.to_string())), + Err(()) => Err(string_to_io_error(format!("Unknown file {}", sanitize_sensitive_info(&input.url)))), } } else { - get_remote_content_as_stream( - app_config, - client, - input, - None, - &url, - ) - .await + get_remote_content_as_stream(app_config, client, input, None, &url).await }; match result { Ok((content, response_url)) => { @@ -1085,7 +1119,7 @@ pub async fn download_text_content_as_stream( debug!("Persisted {size} bytes"); })), ) - .await; + .await; Ok((tee_reader, response_url)) } else { Ok((content, response_url)) @@ -1094,10 +1128,7 @@ pub async fn download_text_content_as_stream( Err(err) => Err(err), } } else { - Err(string_to_io_error(format!( - "Malformed URL {}", - sanitize_sensitive_info(&input.url) - ))) + Err(string_to_io_error(format!("Malformed URL {}", sanitize_sensitive_info(&input.url)))) } } @@ -1108,20 +1139,8 @@ async fn download_json_content( persist_filepath: Option, trace_log: bool, ) -> Result { - debug_if_enabled!( - "Downloading json content from {}", - sanitize_sensitive_info(&input.url) - ); - match download_text_content( - app_config, - client, - input, - None, - persist_filepath, - trace_log, - ) - .await - { + debug_if_enabled!("Downloading json content from {}", sanitize_sensitive_info(&input.url)); + match download_text_content(app_config, client, input, None, persist_filepath, trace_log).await { Ok((content, _response_url)) => match serde_json::from_str::(&content) { Ok(value) => Ok(value), Err(err) => Err(string_to_io_error(format!("Failed to parse json {err}"))), @@ -1137,15 +1156,7 @@ pub async fn get_input_json_content( persist_filepath: Option, trace_log: bool, ) -> Result { - match download_json_content( - app_config, - client, - input, - persist_filepath, - trace_log, - ) - .await - { + match download_json_content(app_config, client, input, persist_filepath, trace_log).await { Ok(content) => Ok(content), Err(e) => notify_err_res!( "can't download input {}, => {}", @@ -1161,18 +1172,8 @@ async fn download_json_content_as_stream( input: &InputSource, persist_filepath: Option, ) -> Result { - debug_if_enabled!( - "Downloading json content as stream from {}", - sanitize_sensitive_info(&input.url) - ); - match download_text_content_as_stream( - app_config, - client, - input, - persist_filepath, - ) - .await - { + debug_if_enabled!("Downloading json content as stream from {}", sanitize_sensitive_info(&input.url)); + match download_text_content_as_stream(app_config, client, input, persist_filepath).await { Ok((reader, _response_url)) => Ok(reader), Err(err) => Err(err), } @@ -1184,14 +1185,7 @@ pub async fn get_input_json_content_as_stream( input: &InputSource, persist_filepath: Option, ) -> Result { - match download_json_content_as_stream( - app_config, - client, - input, - persist_filepath, - ) - .await - { + match download_json_content_as_stream(app_config, client, input, persist_filepath).await { Ok(stream) => Ok(stream), Err(e) => notify_err_res!( "can't download input {} => {}", @@ -1232,9 +1226,7 @@ pub fn create_client_with_redirect(cfg: &AppConfig, redirect_policy: Policy) -> } "http" | "https" => match reqwest::Proxy::all(url.as_str()) { Ok(p) => { - if let (Some(username), Some(password)) = - (&proxy_cfg.username, &proxy_cfg.password) - { + if let (Some(username), Some(password)) = (&proxy_cfg.username, &proxy_cfg.password) { client = client.proxy(p.basic_auth(username, password)); } else { client = client.proxy(p); @@ -1254,11 +1246,7 @@ pub fn create_client_with_redirect(cfg: &AppConfig, redirect_policy: Policy) -> } if let Some(rp_config) = config.reverse_proxy.as_ref() { - if rp_config - .disabled_header - .as_ref() - .is_some_and(|d| d.referer_header) - { + if rp_config.disabled_header.as_ref().is_some_and(|d| d.referer_header) { client = client.referer(false); } } @@ -1285,14 +1273,14 @@ pub fn parse_range(range: &str) -> Option<(u64, Option)> { Some((start, end)) } -pub fn is_file_url(url: &str) -> bool { - Url::parse(url) - .is_ok_and(|u| u.scheme().eq_ignore_ascii_case("file")) -} +pub fn is_file_url(url: &str) -> bool { Url::parse(url).is_ok_and(|u| u.scheme().eq_ignore_ascii_case("file")) } pub fn is_uri(url: &str) -> bool { - Url::parse(url) - .is_ok_and(|u| u.scheme().eq_ignore_ascii_case("file") || u.scheme().eq_ignore_ascii_case("http") || u.scheme().eq_ignore_ascii_case("https")) + Url::parse(url).is_ok_and(|u| { + u.scheme().eq_ignore_ascii_case("file") + || u.scheme().eq_ignore_ascii_case("http") + || u.scheme().eq_ignore_ascii_case("https") + }) } /// Checks if a status code or error indicates a need for failover @@ -1324,7 +1312,10 @@ pub fn should_trigger_failover(status: StatusCode) -> bool { #[cfg(test)] mod tests { + use super::{same_origin, strip_sensitive_headers_for_cross_origin_redirect}; use shared::utils::{get_base_url_from_str, replace_url_extension, sanitize_sensitive_info}; + use std::collections::HashMap; + use url::Url; #[test] fn test_url_mask() { @@ -1338,22 +1329,10 @@ mod tests { fn test_replace_ext() { let tests = [ ("http://hello.world.com", "http://hello.world.com"), - ( - "http://hello.world.com/123", - "http://hello.world.com/123.mp4", - ), - ( - "http://hello.world.com/123.ts?hello=world", - "http://hello.world.com/123.mp4?hello=world", - ), - ( - "http://hello.world.com/123?hello=world", - "http://hello.world.com/123.mp4?hello=world", - ), - ( - "http://hello.world.com/123#hello=world", - "http://hello.world.com/123.mp4#hello=world", - ), + ("http://hello.world.com/123", "http://hello.world.com/123.mp4"), + ("http://hello.world.com/123.ts?hello=world", "http://hello.world.com/123.mp4?hello=world"), + ("http://hello.world.com/123?hello=world", "http://hello.world.com/123.mp4?hello=world"), + ("http://hello.world.com/123#hello=world", "http://hello.world.com/123.mp4#hello=world"), ]; for (test, expect) in &tests { @@ -1372,50 +1351,65 @@ mod tests { fn test_get_request_headers_prioritization() { use super::{get_request_headers, DEFAULT_USER_AGENT}; use axum::http::header::USER_AGENT; - use std::collections::HashMap; // Case 1: No headers provided -> Default UA - let headers = - get_request_headers::(None, None, None, None); + let headers = get_request_headers::(None, None, None, None); assert_eq!(headers.get(USER_AGENT).unwrap(), DEFAULT_USER_AGENT); // Case 2: No headers provided but config default UA set -> Config default UA - let headers = get_request_headers::( - None, - None, - None, - Some("Config-Default-UA"), - ); + let headers = + get_request_headers::(None, None, None, Some("Config-Default-UA")); assert_eq!(headers.get(USER_AGENT).unwrap(), "Config-Default-UA"); // Case 3: Only client header -> Client UA (overrides config default UA) let mut client_headers = HashMap::new(); client_headers.insert("User-Agent".to_string(), b"Client-UA".to_vec()); - let headers = - get_request_headers(None, Some(&client_headers), None, Some("Config-Default-UA")); + let headers = get_request_headers(None, Some(&client_headers), None, Some("Config-Default-UA")); assert_eq!(headers.get(USER_AGENT).unwrap(), "Client-UA"); // Case 4: Both config and client -> Config UA overrides let mut config_headers = HashMap::new(); config_headers.insert("User-Agent".to_string(), "Config-UA".to_string()); - let headers = get_request_headers( - Some(&config_headers), - Some(&client_headers), - None, - Some("Config-Default-UA"), - ); + let headers = + get_request_headers(Some(&config_headers), Some(&client_headers), None, Some("Config-Default-UA")); assert_eq!(headers.get(USER_AGENT).unwrap(), "Config-UA"); // Case 5: Other headers also prioritized config_headers.insert("X-Test".to_string(), "From-Config".to_string()); let mut client_headers = HashMap::new(); client_headers.insert("X-Test".to_string(), b"From-Client".to_vec()); - let headers = get_request_headers( - Some(&config_headers), - Some(&client_headers), - None, - Some("Config-Default-UA"), - ); + let headers = + get_request_headers(Some(&config_headers), Some(&client_headers), None, Some("Config-Default-UA")); assert_eq!(headers.get("X-Test").unwrap(), "From-Config"); } + + #[test] + fn test_same_origin_checks_scheme_host_and_port() { + let a = Url::parse("https://example.com/path").expect("url parse should work"); + let b = Url::parse("https://example.com/other").expect("url parse should work"); + let c = Url::parse("http://example.com/other").expect("url parse should work"); + let d = Url::parse("https://example.com:8443/other").expect("url parse should work"); + + assert!(same_origin(&a, &b)); + assert!(!same_origin(&a, &c)); + assert!(!same_origin(&a, &d)); + } + + #[test] + fn test_cross_origin_redirect_strips_sensitive_headers() { + let mut headers = HashMap::new(); + headers.insert("Authorization".to_string(), "Bearer test".to_string()); + headers.insert("Cookie".to_string(), "sid=123".to_string()); + headers.insert("Proxy-Authorization".to_string(), "Basic abc".to_string()); + headers.insert("Host".to_string(), "old.host".to_string()); + headers.insert("X-Test".to_string(), "ok".to_string()); + + strip_sensitive_headers_for_cross_origin_redirect(&mut headers); + + assert!(!headers.contains_key("Authorization")); + assert!(!headers.contains_key("Cookie")); + assert!(!headers.contains_key("Proxy-Authorization")); + assert!(!headers.contains_key("Host")); + assert_eq!(headers.get("X-Test").map(String::as_str), Some("ok")); + } }