diff --git a/src/api/api_utils.rs b/src/api/api_utils.rs index bb624a11a..136059485 100644 --- a/src/api/api_utils.rs +++ b/src/api/api_utils.rs @@ -1,15 +1,14 @@ use crate::api::endpoints::xtream_api::{get_xtream_player_api_stream_url, XtreamApiStreamContext}; use crate::api::model::active_provider_manager::{ProviderAllocation, ProviderConnectionGuard}; use crate::api::model::app_state::AppState; -use crate::api::model::model_utils::get_stream_response_with_headers; +use crate::api::model::model_utils::{ get_stream_response_with_headers}; use crate::api::model::request::UserApiRequest; use crate::api::model::stream::{BoxedProviderStream, ProviderStreamInfo, ProviderStreamResponse}; use crate::api::model::stream_error::StreamError; use crate::api::model::streams::active_client_stream::ActiveClientStream; use crate::api::model::streams::persist_pipe_stream::PersistPipeStream; -use crate::api::model::streams::provider_stream; -use crate::api::model::streams::provider_stream::{create_custom_video_stream_response, create_provider_connections_exhausted_stream, CustomVideoStreamType}; -use crate::api::model::streams::provider_stream_factory::BufferStreamOptions; +use crate::api::model::streams::provider_stream::{create_channel_unavailable_stream, create_custom_video_stream_response, create_provider_connections_exhausted_stream, CustomVideoStreamType}; +use crate::api::model::streams::provider_stream_factory::{create_provider_stream, ProviderStreamFactoryOptions}; use crate::api::model::streams::shared_stream_manager::SharedStreamManager; use crate::api::model::streams::throttled_stream::ThrottledStream; use crate::auth::authenticator::Claims; @@ -157,6 +156,25 @@ pub struct StreamOptions { pub pipe_provider_stream: bool, } +/// Constructs a `StreamOptions` object based on the application's reverse proxy configuration. +/// +/// This function retrieves streaming-related settings from the `AppState`: +/// - `stream_retry`: whether retrying the stream is enabled, +/// - `stream_force_retry_secs`: the number of seconds to wait before a forced retry, +/// - `buffer_enabled`: whether stream buffering is enabled, +/// - `buffer_size`: the size of the stream buffer. +/// +/// If the reverse proxy or stream settings are not defined, default values are used: +/// - retry: `false` +/// - forced retry interval: `0` +/// - buffering: `false` +/// - buffer size: `0` +/// +/// Additionally, it computes `pipe_provider_stream`, which is `true` only if +/// both retry and buffering are disabled—indicating that the stream can be piped directly +/// from the provider without additional handling. +/// +/// Returns a `StreamOptions` instance with the resolved configuration. fn get_stream_options(app_state: &AppState) -> StreamOptions { let (stream_retry, stream_force_retry_secs, buffer_enabled, buffer_size) = app_state .config @@ -214,7 +232,7 @@ async fn get_redirect_alternative_url<'a>(app_state: &AppState, redirect_url: &' type StreamUrl = String; type ProviderName = String; -enum StreamingOption { +enum ProviderStreamState { Custom(ProviderStreamResponse), Available(Option, StreamUrl), GracePeriod(Option, StreamUrl), @@ -251,11 +269,31 @@ impl StreamDetails { } } -/** -* If successfully a provider connection is used, do not forget to release if unsuccessfully -*/ -async fn get_streaming_options(app_state: &AppState, stream_url: &str, input: &ConfigInput, force_provider: Option<&str>) - -> (Option, StreamingOption, Option>) { +struct StreamingStrategy { + provider_connection_guard: Option, + provider_stream_state: ProviderStreamState, + input_headers: Option>, +} + +/// Determines the appropriate streaming strategy for the given input and stream URL. +/// +/// This function attempts to acquire a connection to a streaming provider, either using a forced provider +/// (if specified), or based on the input name. It then selects a corresponding `StreamingOption`: +/// +/// - If no connections are available (`Exhausted`), it returns a custom stream indicating exhaustion. +/// - If a connection is available or in a grace period, it constructs a streaming URL accordingly: +/// - If the provider was forced or matches the input, the original URL is reused. +/// - Otherwise, an alternative URL is generated based on the provider and input. +/// +/// The function returns: +/// - an optional `ProviderConnectionGuard` to manage the connection's lifecycle, +/// - a `ProviderStreamState` describing how the stream state is, +/// - and optional HTTP headers to include in the request. +/// +/// This logic helps abstract the decision-making behind provider selection and stream URL resolution. +async fn resolve_streaming_strategy(app_state: &AppState, stream_url: &str, input: &ConfigInput, force_provider: Option<&str>) + -> StreamingStrategy { + // allocate a provider connection let provider_connection_guard = match force_provider { Some(provider) => app_state.active_provider.force_exact_acquire_connection(provider).await, None => app_state.active_provider.acquire_connection(&input.name).await @@ -264,7 +302,7 @@ async fn get_streaming_options(app_state: &AppState, stream_url: &str, input: &C ProviderAllocation::Exhausted => { debug!("Input {} is exhausted. No connections allowed.", input.name); let stream = create_provider_connections_exhausted_stream(&app_state.config, &[]); - StreamingOption::Custom(stream) + ProviderStreamState::Custom(stream) } ProviderAllocation::Available(ref provider) | ProviderAllocation::GracePeriod(ref provider) => { @@ -277,19 +315,23 @@ async fn get_streaming_options(app_state: &AppState, stream_url: &str, input: &C }; if matches!(&*provider_connection_guard, ProviderAllocation::Available(_)) { - StreamingOption::Available(Some(provider), url) + ProviderStreamState::Available(Some(provider), url) } else { - StreamingOption::GracePeriod(Some(provider), url) + ProviderStreamState::GracePeriod(Some(provider), url) } } }; - (Some(provider_connection_guard), stream_response_params, Some(input.headers.clone())) + StreamingStrategy { + provider_connection_guard: Some(provider_connection_guard), + provider_stream_state: stream_response_params, + input_headers: Some(input.headers.clone()) + } } -fn get_grace_period_millis(connection_permission: &UserConnectionPermission, stream_response_params: &StreamingOption, config_grace_period_millis: u64) -> u64 { +fn get_grace_period_millis(connection_permission: &UserConnectionPermission, stream_response_params: &ProviderStreamState, config_grace_period_millis: u64) -> u64 { if config_grace_period_millis > 0 && - (matches!(stream_response_params, StreamingOption::GracePeriod(_, _)) // provider grace period + (matches!(stream_response_params, ProviderStreamState::GracePeriod(_, _)) // provider grace period || connection_permission == &UserConnectionPermission::GracePeriod // user grace period ) { config_grace_period_millis } else { 0 } } @@ -304,13 +346,14 @@ async fn create_stream_response_details(app_state: &AppState, share_stream: bool, connection_permission: UserConnectionPermission, force_provider: Option<&str>) -> StreamDetails { - let (mut provider_connection_guard, stream_response_params, input_headers) = - get_streaming_options(app_state, stream_url, input, force_provider).await; + let mut streaming_strategy = + resolve_streaming_strategy(app_state, stream_url, input, force_provider).await; let config_grace_period_millis = app_state.config.reverse_proxy.as_ref() .and_then(|r| r.stream.as_ref()).map_or_else(default_grace_period_millis, |s| s.grace_period_millis); - let grace_period_millis = get_grace_period_millis(&connection_permission, &stream_response_params, config_grace_period_millis); - match stream_response_params { - StreamingOption::Custom(provider_stream) => { + let grace_period_millis = get_grace_period_millis(&connection_permission, &streaming_strategy.provider_stream_state, config_grace_period_millis); + match streaming_strategy.provider_stream_state { + // custom stream means we display our own stream like connection exhausted, channel unavailable... + ProviderStreamState::Custom(provider_stream) => { let (stream, stream_info) = provider_stream; StreamDetails { stream, @@ -318,28 +361,31 @@ async fn create_stream_response_details(app_state: &AppState, input_name: None, grace_period_millis, reconnect_flag: None, - provider_connection_guard, + provider_connection_guard: streaming_strategy.provider_connection_guard.take(), } } - StreamingOption::Available(provider_name, request_url) | - StreamingOption::GracePeriod(provider_name, request_url) => { + ProviderStreamState::Available(provider_name, request_url) | + ProviderStreamState::GracePeriod(provider_name, request_url) => { let parsed_url = Url::parse(&request_url); let ((stream, stream_info), reconnect_flag) = if let Ok(url) = parsed_url { - if stream_options.pipe_provider_stream { - (provider_stream::get_provider_pipe_stream(app_state, &url, req_headers, input_headers.as_ref(), item_type).await, None) - } else { - let buffer_stream_options = BufferStreamOptions::new(item_type, share_stream, stream_options); - let reconnect_flag = buffer_stream_options.get_reconnect_flag_clone(); - (provider_stream::get_provider_reconnect_buffered_stream(app_state, &url, req_headers, input_headers.as_ref(), buffer_stream_options).await, - Some(reconnect_flag)) - } + let provider_stream_factory_options = ProviderStreamFactoryOptions::new(item_type, share_stream, stream_options, &url, req_headers, streaming_strategy.input_headers.as_ref()); + let reconnect_flag = provider_stream_factory_options.get_reconnect_flag_clone(); + let provider_stream = match create_provider_stream(Arc::clone(&app_state.config), Arc::clone(&app_state.http_client), provider_stream_factory_options).await { + None => (None, None), + Some((stream, info)) => { + (Some(stream), info) + } + }; + (provider_stream, Some(reconnect_flag)) } else { ((None, None), None) }; // if we have no stream we should release the provider if stream.is_none() { - drop(provider_connection_guard.take()); + if let Some(guard) = streaming_strategy.provider_connection_guard.take() { + drop(guard); + } error!("Cant open stream {}", sanitize_sensitive_info(&request_url)); } @@ -360,7 +406,7 @@ async fn create_stream_response_details(app_state: &AppState, input_name: provider_name, grace_period_millis, reconnect_flag, - provider_connection_guard, + provider_connection_guard: streaming_strategy.provider_connection_guard.take(), } } } @@ -480,7 +526,7 @@ const SESSION_COOKIE_NAME: &str = "m3uflt_session="; pub fn create_session_cookie(token: u32) -> String { let cookie = crate::utils::u32_to_base64(token); - format!("{SESSION_COOKIE_NAME}{cookie}; Path=/; HttpOnly; SameSite=Lax") + format!("{SESSION_COOKIE_NAME}{cookie}; Path=/; HttpOnly; SameSite=None") } pub fn read_session_token(headers: &HeaderMap) -> Option { @@ -528,8 +574,7 @@ pub async fn force_provider_stream_response(app_state: &AppState, let stream = ActiveClientStream::new(stream_details, app_state, user, connection_permission).await; let (status_code, header_map) = get_stream_response_with_headers(provider_response); - let mut response = axum::response::Response::builder() - .status(status_code); + let mut response = axum::response::Response::builder().status(status_code); for (key, value) in &header_map { response = response.header(key, value); } @@ -539,7 +584,14 @@ pub async fn force_provider_stream_response(app_state: &AppState, return response.body(body_stream).unwrap().into_response(); } drop(stream_details.provider_connection_guard.take()); - StatusCode::BAD_REQUEST.into_response() + if let (Some(stream), _stream_info) = + create_channel_unavailable_stream(&app_state.config, &vec![], StatusCode::BAD_GATEWAY) + { + debug!("Streaming custom stream"); + axum::response::Response::builder().status(StatusCode::OK).body(Body::from_stream(stream)).unwrap().into_response() + } else { + StatusCode::BAD_REQUEST.into_response() + } } /// # Panics diff --git a/src/api/model/active_user_manager.rs b/src/api/model/active_user_manager.rs index e0fe8cae2..c54787111 100644 --- a/src/api/model/active_user_manager.rs +++ b/src/api/model/active_user_manager.rs @@ -109,7 +109,7 @@ impl ActiveUserManager { let now = get_current_timestamp(); // Check if user already used grace period if connection_data.granted_grace { - if now - connection_data.grace_ts <= self.grace_period_timeout_secs { + if current_connections > connection_data.max_connections && now - connection_data.grace_ts <= self.grace_period_timeout_secs { // Grace timeout still active, deny connection debug!("User access denied, grace exhausted, too many connections: {username}"); return UserConnectionPermission::Exhausted; diff --git a/src/api/model/model_utils.rs b/src/api/model/model_utils.rs index fdf246681..849c5057b 100644 --- a/src/api/model/model_utils.rs +++ b/src/api/model/model_utils.rs @@ -2,11 +2,11 @@ use reqwest::{StatusCode}; use std::collections::{HashSet}; use std::str::FromStr; use reqwest::header::HeaderMap; -use crate::utils::{MEDIA_STREAM_HEADERS}; +use crate::utils::{filter_response_header}; pub fn get_response_headers(headers: &HeaderMap) -> Vec<(String, String)> { let response_headers: Vec<(String, String)> = headers.iter() - .filter(|(key, _)| MEDIA_STREAM_HEADERS.contains(&key.as_str())) + .filter(|(key, _)| filter_response_header(key.as_str())) .map(|(key, value)| (key.to_string(), value.to_str().unwrap().to_string())).collect(); response_headers } @@ -27,8 +27,7 @@ pub fn get_stream_response_with_headers(custom: Option<(Vec<(String, String)>, S } let default_headers = vec![ - ("content-type", "application/octet-stream"), - ("connection", "keep-alive"), + ("content-type", "application/octet-stream") ]; for (key, value) in default_headers { diff --git a/src/api/model/provider_config.rs b/src/api/model/provider_config.rs index 203a8d059..c86bcc92d 100644 --- a/src/api/model/provider_config.rs +++ b/src/api/model/provider_config.rs @@ -94,7 +94,7 @@ impl ProviderConfig { guard.grace_ts = 0; } - if guard.granted_grace { + if guard.granted_grace && guard.current_connections > max { let now = get_current_timestamp(); if now - guard.grace_ts <= grace_period_timeout_secs { // Grace timeout still active, deny connection @@ -133,8 +133,8 @@ impl ProviderConfig { } let now = get_current_timestamp(); - if guard.granted_grace { - if 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; @@ -167,7 +167,7 @@ impl ProviderConfig { let now = get_current_timestamp(); if guard.granted_grace { - if now - guard.grace_ts <= grace_period_timeout_secs { + 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: {}", self.name); return false; diff --git a/src/api/model/streams/provider_stream.rs b/src/api/model/streams/provider_stream.rs index ff6d7cce0..740c67918 100644 --- a/src/api/model/streams/provider_stream.rs +++ b/src/api/model/streams/provider_stream.rs @@ -1,21 +1,11 @@ -use std::collections::HashMap; -use crate::api::api_utils::{get_headers_from_request, HeaderFilter}; -use crate::api::model::model_utils::get_response_headers; -use crate::api::model::stream_error::StreamError; +use crate::api::api_utils::{HeaderFilter}; use crate::api::model::streams::custom_video_stream::CustomVideoStream; -use crate::api::model::streams::provider_stream_factory::{create_provider_stream, BufferStreamOptions}; use crate::model::{Config}; use crate::model::PlaylistItemType; -use crate::utils::debug_if_enabled; -use crate::utils::request::{get_request_headers, sanitize_sensitive_info}; -use futures::TryStreamExt; -use log::{error, trace}; +use log::{trace}; use reqwest::StatusCode; use std::sync::Arc; -use axum::http::HeaderMap; use axum::response::IntoResponse; -use url::Url; -use crate::api::model::app_state::AppState; use crate::api::model::stream::ProviderStreamResponse; pub enum CustomVideoStreamType { @@ -72,54 +62,3 @@ pub fn get_header_filter_for_item_type(item_type: PlaylistItemType) -> HeaderFil _ => None, } } - -pub async fn get_provider_pipe_stream(app_state: &AppState, - stream_url: &Url, - req_headers: &HeaderMap, - input_headers: Option<&HashMap>, - item_type: PlaylistItemType) -> ProviderStreamResponse { - let filter_header = get_header_filter_for_item_type(item_type); - let req_headers = get_headers_from_request(req_headers, &filter_header); - debug_if_enabled!("Stream requested with headers: {:?}", req_headers.iter().map(|header| (header.0, String::from_utf8_lossy(header.1))).collect::>()); - // These are the configured headers for this input. - // The stream url, we need to clone it because of move to async block. - // We merge configured input headers with the headers from the request. - let headers = get_request_headers(input_headers, Some(&req_headers)); - let client = app_state.http_client.get(stream_url.clone()).headers(headers.clone()); - match client.send().await { - Ok(response) => { - let response_headers = get_response_headers(response.headers()); - // TODO hls handling if response gives us a m3u8 file, check content type - let status = response.status(); - if status.is_success() { - (Some(Box::pin(response.bytes_stream().map_err(|err| StreamError::reqwest(&err)))), Some((response_headers, status))) - } else if let (Some(boxed_provider_stream), response_info) = create_channel_unavailable_stream(&app_state.config, &response_headers, status) { - (Some(boxed_provider_stream), response_info) - } else { - (None, Some((response_headers, status))) - } - } - Err(err) => { - let masked_url = sanitize_sensitive_info(stream_url.as_str()); - error!("Failed to open stream {masked_url} {err}"); - if let (Some(boxed_provider_stream), response_info) = create_channel_unavailable_stream(&app_state.config, &get_response_headers(&headers), StatusCode::BAD_GATEWAY) { - (Some(boxed_provider_stream), response_info) - } else { - (None, None) - } - } - } -} - -pub async fn get_provider_reconnect_buffered_stream(app_state: &AppState, - stream_url: &Url, - req_headers: &HeaderMap, - input_headers: Option<&HashMap>, - options: BufferStreamOptions) -> ProviderStreamResponse { - match create_provider_stream(&app_state.config, Arc::clone(&app_state.http_client), stream_url, req_headers, input_headers, options).await { - None => (None, None), - Some((stream, info)) => { - (Some(stream), info) - } - } -} diff --git a/src/api/model/streams/provider_stream_factory.rs b/src/api/model/streams/provider_stream_factory.rs index 6ebd80191..e154c56b3 100644 --- a/src/api/model/streams/provider_stream_factory.rs +++ b/src/api/model/streams/provider_stream_factory.rs @@ -1,11 +1,12 @@ use crate::api::api_utils::{get_headers_from_request, StreamOptions}; use crate::api::model::model_utils::get_response_headers; +use crate::api::model::stream::{BoxedProviderStream, ProviderStreamFactoryResponse}; use crate::api::model::stream_error::StreamError; use crate::api::model::streams::buffered_stream::BufferedStream; use crate::api::model::streams::client_stream::ClientStream; use crate::api::model::streams::provider_stream::{create_channel_unavailable_stream, get_header_filter_for_item_type}; -use crate::api::model::streams::timed_client_stream::{TimeoutClientStream}; -use crate::model::{Config}; +use crate::api::model::streams::timed_client_stream::TimeoutClientStream; +use crate::model::Config; use crate::model::PlaylistItemType; use crate::tools::atomic_once_flag::AtomicOnceFlag; use crate::utils::debug_if_enabled; @@ -18,40 +19,72 @@ use reqwest::StatusCode; use std::collections::HashMap; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; -use std::time::{Duration, Instant}; +use std::time::Instant; use url::Url; -use crate::api::model::stream::{BoxedProviderStream, ProviderStreamFactoryResponse}; // TODO make this configurable pub const STREAM_QUEUE_SIZE: usize = 4096; // mpsc channel holding messages. with possible 8192byte chunks +const RETRY_SECONDS: u64 = 5; +const ERR_MAX_RETRY_COUNT: u32 = 5; -pub struct BufferStreamOptions { - item_type: PlaylistItemType, +#[allow(clippy::struct_excessive_bools)] +#[derive(Debug, Clone)] +pub struct ProviderStreamFactoryOptions { + // item_type: PlaylistItemType, reconnect_enabled: bool, force_reconnect_secs: u32, buffer_enabled: bool, buffer_size: usize, share_stream: bool, - reconnect_flag: Arc + pipe_stream: bool, + url: Url, + headers: HeaderMap, + range_bytes: Arc>, + reconnect_flag: Arc, } -impl BufferStreamOptions { +impl ProviderStreamFactoryOptions { pub(crate) fn new( item_type: PlaylistItemType, share_stream: bool, - stream_options: &StreamOptions + stream_options: &StreamOptions, + stream_url: &Url, + req_headers: &HeaderMap, + input_headers: Option<&HashMap>, ) -> Self { + let buffer_size = if stream_options.buffer_enabled { stream_options.buffer_size } else { STREAM_QUEUE_SIZE }; + let filter_header = get_header_filter_for_item_type(item_type); + let mut req_headers = get_headers_from_request(req_headers, &filter_header); + // we need the range bytes from client request for seek ing to the right position + let range_start_bytes = 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)); + + let url = stream_url.clone(); + let range_bytes = Arc::new(range_start_bytes.map(AtomicUsize::new)); + Self { - item_type, + // item_type, reconnect_enabled: stream_options.stream_retry, force_reconnect_secs: stream_options.stream_force_retry_secs, + pipe_stream: stream_options.pipe_provider_stream, buffer_enabled: stream_options.buffer_enabled, - buffer_size: stream_options.buffer_size, + buffer_size, share_stream, - reconnect_flag: Arc::new(AtomicOnceFlag::new()) + reconnect_flag: Arc::new(AtomicOnceFlag::new()), + url, + headers, + range_bytes, } } + #[inline] + fn is_piped(&self) -> bool { + self.pipe_stream + } + #[inline] fn is_buffer_enabled(&self) -> bool { self.buffer_enabled @@ -62,19 +95,9 @@ impl BufferStreamOptions { self.share_stream } - // #[inline] - // fn get_buffer_size(&self) -> usize { - // self.buffer_size - // } - #[inline] - fn is_reconnect_enabled(&self) -> bool { - self.reconnect_enabled - } - - #[inline] - pub(crate) fn get_stream_buffer_size(&self) -> usize { - if self.buffer_size > 0 { self.buffer_size } else { STREAM_QUEUE_SIZE } + pub(crate) fn get_buffer_size(&self) -> usize { + self.buffer_size } #[inline] @@ -82,42 +105,9 @@ impl BufferStreamOptions { Arc::clone(&self.reconnect_flag) } -} - - -#[derive(Debug, Clone)] -struct ProviderStreamOptions { - buffer_size: usize, - continue_flag: Arc, - url: Url, - reconnect: bool, - reconnect_force_secs: u32, - headers: HeaderMap, - range_bytes: Arc>, -} - -impl ProviderStreamOptions { - #[inline] - pub fn is_buffered(&self) -> bool { - self.buffer_size > 0 - } - #[inline] - pub fn get_buffer_size(&self) -> usize { - self.buffer_size - } - #[inline] - pub fn get_continue_flag_clone(&self) -> Arc { - Arc::clone(&self.continue_flag) - } - - // #[inline] - // pub fn get_continue_flag(&self) -> &Arc { - // &self.continue_flag - // } - #[inline] pub fn cancel_reconnect(&self) { - self.continue_flag.notify(); + self.reconnect_flag.notify(); } #[inline] @@ -125,9 +115,14 @@ impl ProviderStreamOptions { &self.url } + #[inline] + pub fn get_url_as_str(&self) -> &str { + self.url.as_str() + } + #[inline] pub fn should_reconnect(&self) -> bool { - self.reconnect + self.reconnect_enabled } #[inline] @@ -151,10 +146,91 @@ impl ProviderStreamOptions { #[inline] pub fn should_continue(&self) -> bool { - self.continue_flag.is_active() + self.reconnect_flag.is_active() + } + + #[inline] + pub fn get_reconnect_force_secs(&self) -> u32 { + self.force_reconnect_secs } } +// +// #[derive(Debug, Clone)] +// struct ProviderStreamOptions { +// buffer_size: usize, +// continue_flag: Arc, +// url: Url, +// pipe_stream: bool, +// reconnect: bool, +// reconnect_force_secs: u32, +// headers: HeaderMap, +// range_bytes: Arc>, +// } +// +// impl ProviderStreamOptions { +// #[inline] +// pub fn is_piped(&self) -> bool { +// self.pipe_stream +// } +// #[inline] +// pub fn is_buffered(&self) -> bool { +// self.buffer_size > 0 +// } +// #[inline] +// pub fn get_buffer_size(&self) -> usize { +// self.buffer_size +// } +// #[inline] +// pub fn get_continue_flag_clone(&self) -> Arc { +// Arc::clone(&self.continue_flag) +// } +// +// // #[inline] +// // pub fn get_continue_flag(&self) -> &Arc { +// // &self.continue_flag +// // } +// +// #[inline] +// pub fn cancel_reconnect(&self) { +// self.continue_flag.notify(); +// } +// +// #[inline] +// pub fn get_url(&self) -> &Url { +// &self.url +// } +// +// #[inline] +// pub fn should_reconnect(&self) -> bool { +// self.reconnect +// } +// +// #[inline] +// 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::SeqCst)) +// } +// +// // pub fn get_range_bytes(&self) -> &Arc> { +// // &self.range_bytes +// // } +// +// #[inline] +// pub fn get_range_bytes_clone(&self) -> Arc> { +// Arc::clone(&self.range_bytes) +// } +// +// #[inline] +// pub fn should_continue(&self) -> bool { +// self.continue_flag.is_active() +// } +// } + fn get_request_range_start_bytes(req_headers: &HashMap>) -> Option { // range header looks like bytes=1234-5566/2345345 or bytes=0- if let Some(req_range) = req_headers.get(axum::http::header::RANGE.as_str()) { @@ -172,46 +248,55 @@ fn get_request_range_start_bytes(req_headers: &HashMap>) -> Opti None } -fn get_client_stream_request_params( - req_headers: &HeaderMap, - input_headers: Option<&HashMap>, - options: &BufferStreamOptions) -> (usize, Option, bool, u32, HeaderMap) -{ - let stream_buffer_size = if options.is_buffer_enabled() { options.get_stream_buffer_size() } else { 0 }; - let filter_header = get_header_filter_for_item_type(options.item_type); - let mut req_headers = get_headers_from_request(req_headers, &filter_header); - debug_if_enabled!("Stream requested with headers: {:?}", req_headers.iter().map(|header| (header.0, String::from_utf8_lossy(header.1))).collect::>()); - // we need the range bytes from client request for seek ing to the right position - let req_range_start_bytes = 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)); - - (stream_buffer_size, req_range_start_bytes, options.is_reconnect_enabled(), options.force_reconnect_secs, headers) +fn get_host_and_optional_port(url: &Url) -> Option { + let host = url.host_str()?; + match url.port() { + Some(port) => Some(format!("{}:{}", host, port)), + None => Some(host.to_string()), + } } -fn prepare_client(request_client: &Arc, stream_options: &ProviderStreamOptions) -> (reqwest::RequestBuilder, bool) { - let url = stream_options.get_url(); - let range_start = stream_options.get_total_bytes_send(); - let headers = stream_options.get_headers(); - let mut request_builder = request_client.get(url.clone()).headers(headers.clone()); - let (client, partial) = { - if let Some(range) = range_start { - // on reconnect send range header to avoid starting from beginning for vod - let range = format!("bytes={range}-", ); - request_builder = request_builder.header(RANGE, range); - (request_builder, true) // partial content - } else { - (request_builder, false) +fn prepare_client(request_client: &Arc, stream_options: &ProviderStreamFactoryOptions) -> (reqwest::RequestBuilder, bool) { + let url = stream_options.get_url(); + let host = get_host_and_optional_port(url); + let range_start = stream_options.get_total_bytes_send(); + let original_headers = stream_options.get_headers(); + + let mut headers = original_headers.clone(); + + if let Some(host_header) = host { + if let Ok(header_value) = axum::http::header::HeaderValue::from_str(&host_header) { + headers.insert(axum::http::header::HOST, header_value); } + } + + if !headers.contains_key(axum::http::header::TRANSFER_ENCODING) { + headers.insert(axum::http::header::TRANSFER_ENCODING, axum::http::header::HeaderValue::from_static("chunked")); + } + + if !headers.contains_key(axum::http::header::USER_AGENT) { + headers.insert(axum::http::header::USER_AGENT, axum::http::header::HeaderValue::from_static("Mozilla/5.0 (Linux; Android 10)")); + } + + let partial = if let Some(range) = range_start { + let range_header = format!("bytes={range}-"); + if let Ok(header_value) = axum::http::header::HeaderValue::from_str(&range_header) { + headers.insert(RANGE, header_value); + } + true + } else { + false }; - (client, partial) + debug_if_enabled!("Stream requested with headers: {:?}", headers.iter().map(|header| (header.0, String::from_utf8_lossy(header.1.as_ref()))).collect::>()); + + let request_builder = request_client.get(url.clone()).headers(headers); + + (request_builder, partial) } -async fn provider_initial_request(cfg: &Config, request_client: Arc, stream_options: &ProviderStreamOptions) -> Result, StatusCode> { +async fn provider_stream_request(cfg: &Config, request_client: Arc, stream_options: &ProviderStreamFactoryOptions) -> Result, StatusCode> { let (client, _partial_content) = prepare_client(&request_client, stream_options); match client.send().await { Ok(mut response) => { @@ -226,16 +311,33 @@ async fn provider_initial_request(cfg: &Config, request_client: Arc 0 { + TimeoutClientStream::new(provider_stream, stream_options.get_reconnect_force_secs()).boxed() + } else { + provider_stream + }; return Ok(Some((boxed_provider_stream, response_info))); } + + if status.is_client_error() { + debug!("Client error status response : {status}"); + } + if status.is_server_error() { + match status { + StatusCode::INTERNAL_SERVER_ERROR | + StatusCode::BAD_GATEWAY | + StatusCode::SERVICE_UNAVAILABLE | + StatusCode::GATEWAY_TIMEOUT => {} + _ => { + debug!("Server error status response : {status}"); + } + } + } Err(status) } Err(_err) => { @@ -250,79 +352,37 @@ async fn provider_initial_request(cfg: &Config, request_client: Arc, stream_options: ProviderStreamOptions) -> Option { +async fn get_provider_stream(cfg: &Config, client: Arc, stream_options: &ProviderStreamFactoryOptions) -> Result, StatusCode> { let url = stream_options.get_url(); debug_if_enabled!("stream provider {}", sanitize_sensitive_info(url.as_str())); - while stream_options.should_continue() { - debug_if_enabled!("Reconnecting stream {}", sanitize_sensitive_info(url.as_str())); - let (client, _) = prepare_client(&client, &stream_options); - match client.send().await { - Ok(response) => { - let status = response.status(); - if status.is_success() { - let provider_stream = response.bytes_stream().map_err(|err| { - // error!("Stream error {err}"); - StreamError::reqwest(&err) - }).boxed(); - return if stream_options.reconnect_force_secs > 0 { - Some(TimeoutClientStream::new(provider_stream, stream_options.reconnect_force_secs).boxed()) - } else { - Some(provider_stream) - }; - } - if status.is_client_error() { - debug!("Client error status response : {status}"); - return None; - } - if status.is_server_error() { - match status { - StatusCode::INTERNAL_SERVER_ERROR | - StatusCode::BAD_GATEWAY | - StatusCode::SERVICE_UNAVAILABLE | - StatusCode::GATEWAY_TIMEOUT => {} - _ => { - debug!("Server error status response : {status}"); - return None - } - } - } - } - Err(err) => { - debug!("Server connection failed with {err}"); - } - } - if !stream_options.should_continue() { - return None; - } - tokio::time::sleep(Duration::from_millis(100)).await; - } - debug_if_enabled!("Stopped reconnecting stream {}", sanitize_sensitive_info(url.as_str())); - None -} - -const RETRY_SECONDS: u64 = 5; -const ERR_MAX_RETRY_COUNT: u32 = 5; -async fn get_initial_stream(cfg: &Config, client: Arc, stream_options: &ProviderStreamOptions) -> Option { let start = Instant::now(); let mut connect_err: u32 = 1; + while stream_options.should_continue() { - match provider_initial_request(cfg, Arc::clone(&client), stream_options).await { - Ok(Some(value)) => return Some(value), + match provider_stream_request(cfg, Arc::clone(&client), stream_options).await { + Ok(Some(stream_response)) => { + return Ok(Some(stream_response)); + } Ok(None) => { if connect_err > ERR_MAX_RETRY_COUNT { warn!("The stream could be unavailable. {}", sanitize_sensitive_info(stream_options.get_url().as_str())); } } Err(status) => { + debug!("Provider stream response error status response : {status}"); if status == StatusCode::FORBIDDEN || status == StatusCode::SERVICE_UNAVAILABLE || status == StatusCode::UNAUTHORIZED { warn!("The stream could be unavailable. ({status}) {}", sanitize_sensitive_info(stream_options.get_url().as_str())); - break; + stream_options.cancel_reconnect(); + return Err(status); } if connect_err > ERR_MAX_RETRY_COUNT { warn!("The stream could be unavailable. ({status}) {}", sanitize_sensitive_info(stream_options.get_url().as_str())); } } } + if !stream_options.should_continue() { + return Err(StatusCode::SERVICE_UNAVAILABLE); + } if connect_err > ERR_MAX_RETRY_COUNT { break; } @@ -331,70 +391,61 @@ async fn get_initial_stream(cfg: &Config, client: Arc, stream_o break; } connect_err += 1; - tokio::time::sleep(Duration::from_millis(100)).await; + // tokio::time::sleep(Duration::from_millis(50)).await; + debug_if_enabled!("Reconnecting stream {}", sanitize_sensitive_info(url.as_str())); } + debug_if_enabled!("Stopped reconnecting stream {}", sanitize_sensitive_info(url.as_str())); stream_options.cancel_reconnect(); - None + Err(StatusCode::SERVICE_UNAVAILABLE) } -fn create_provider_stream_options(stream_url: &Url, - req_headers: &HeaderMap, - input_headers: Option<&HashMap>, - options: &BufferStreamOptions) -> ProviderStreamOptions { - let (buffer_size, req_range_start_bytes, reconnect, reconnect_force_secs, headers) - = get_client_stream_request_params(req_headers, input_headers, options); - let url = stream_url.clone(); - let range_bytes = Arc::new(req_range_start_bytes.map(AtomicUsize::new)); - ProviderStreamOptions { - buffer_size, - continue_flag: Arc::clone(&options.reconnect_flag), - url, - reconnect, - reconnect_force_secs, - headers, - range_bytes, - } -} - -pub async fn create_provider_stream(cfg: &Config, +pub async fn create_provider_stream(cfg: Arc, client: Arc, - stream_url: &Url, - req_headers: &HeaderMap, - input_headers: Option<&HashMap>, - options: BufferStreamOptions) -> Option { - let stream_options = create_provider_stream_options(stream_url, req_headers, input_headers, &options); - + stream_options: ProviderStreamFactoryOptions) -> Option { let client_stream_factory = |stream, reconnect_flag, range_cnt| { - let stream = if stream_options.is_buffered() && !options.is_shared_stream() { - BufferedStream::new(stream, stream_options.get_buffer_size(), stream_options.get_continue_flag_clone(), stream_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() + ClientStream::new(stream, reconnect_flag, range_cnt, stream_options.get_url_as_str()).boxed() }; - match get_initial_stream(cfg, Arc::clone(&client), &stream_options).await { - Some((init_stream, info)) => { - let is_media_stream = if let Some((headers, _)) = &info { - classify_content_type(headers) == MimeCategory::Video + match get_provider_stream(&cfg, Arc::clone(&client), &stream_options).await { + Ok(Some((init_stream, info))) => { + let is_media_stream_or_not_piped = if let Some((headers, _)) = &info { + // if it is piped or no video stream, then we don't reconnect + !stream_options.pipe_stream && classify_content_type(headers) == MimeCategory::Video } else { - true // don't know what it is but lets assume it is + !stream_options.pipe_stream // don't know what it is but lets assume it is something }; - let continue_signal = stream_options.get_continue_flag_clone(); - if is_media_stream && stream_options.should_reconnect() { + let continue_signal = stream_options.get_reconnect_flag_clone(); + if is_media_stream_or_not_piped && stream_options.should_reconnect() { let continue_client_signal = Arc::clone(&continue_signal); let continue_streaming_signal = continue_client_signal.clone(); let stream_options_provider = stream_options.clone(); + let config = Arc::clone(&cfg); let unfold: BoxedProviderStream = stream::unfold((), move |()| { let client = Arc::clone(&client); let stream_opts = stream_options_provider.clone(); let continue_streaming = continue_streaming_signal.clone(); + let config_clone = Arc::clone(&config); async move { if continue_streaming.is_active() { - let stream = stream_provider(client, stream_opts).await?; - Some((stream, ())) + match get_provider_stream(&config_clone, client, &stream_opts).await { + Ok(Some((stream, _info))) => Some((stream, ())), + Ok(None) => None, + Err(status) => { + if let (Some(boxed_provider_stream), _response_info) = + create_channel_unavailable_stream(&config_clone, &get_response_headers(stream_opts.get_headers()), status) + { + return Some((boxed_provider_stream, ())); + } + None + } + } } else { None } @@ -405,7 +456,17 @@ pub async fn create_provider_stream(cfg: &Config, Some((client_stream_factory(init_stream.boxed(), Arc::clone(&continue_signal), stream_options.get_range_bytes_clone()).boxed(), info)) } } - None => None + Ok(None) => { + None + } + Err(status) => { + if let (Some(boxed_provider_stream), response_info) = + create_channel_unavailable_stream(&cfg, &get_response_headers(stream_options.get_headers()), status) + { + return Some((boxed_provider_stream, response_info)); + } + None + } } } diff --git a/src/utils/constants.rs b/src/utils/constants.rs index 4e44fd350..fdcdb6a62 100644 --- a/src/utils/constants.rs +++ b/src/utils/constants.rs @@ -29,10 +29,35 @@ pub const DASH_EXT_FRAGMENT: &str = ".mpd#"; pub const FILENAME_TRIM_PATTERNS: &[char] = &['.', '-', '_']; -pub const MEDIA_STREAM_HEADERS: &[&str] = &["accept", "content-type", "content-length", "connection", - "accept-ranges", "content-range", "vary", "transfer-encoding", "access-control-allow-origin", - "access-control-allow-credentials", "icy-metadata", "cache-control", "referer", "last-modified", - "etag", "expires"]; +const SUPPORTED_RESPONSE_HEADERS: &[&str] = &[ + "accept", + "accept-ranges", + "content-type", + "content-length", + "content-range", + "vary", + "transfer-encoding", + "connection", + "access-control-allow-origin", + "access-control-allow-credentials", + "icy-metadata", + "referer", + "last-modified", + "cache-control", + "etag", + "expires" +]; + +pub fn filter_response_header(key: &str) -> bool { + SUPPORTED_RESPONSE_HEADERS.contains(&key) +} + +pub fn filter_request_header(key: &str) -> bool { + if key == "host" { + return false; + } + true +} pub struct KodiStyle { pub year: Regex, diff --git a/src/utils/network/request.rs b/src/utils/network/request.rs index 9689d2587..438b97fd0 100644 --- a/src/utils/network/request.rs +++ b/src/utils/network/request.rs @@ -21,7 +21,7 @@ use crate::model::{ConfigInput, ProxyConfig, InputFetchMethod}; use crate::repository::storage::{get_input_storage_path, short_hash}; use crate::repository::storage_const; use crate::utils::compression::compression_utils::{is_deflate, is_gzip}; -use crate::utils::debug_if_enabled; +use crate::utils::{debug_if_enabled, filter_request_header}; use crate::utils::file_utils::{get_file_path, persist_file}; use crate::utils::{CONSTANTS, DASH_EXT, DASH_EXT_FRAGMENT, DASH_EXT_QUERY, ENCODING_DEFLATE, ENCODING_GZIP, HLS_EXT, HLS_EXT_FRAGMENT, HLS_EXT_QUERY}; @@ -142,12 +142,14 @@ pub fn get_client_request(client: &Arc(defined_headers: Option<&HashMap>, custom_headers: Option<&HashMap, S>>) -> HeaderMap { +pub fn get_request_headers(request_headers: Option<&HashMap>, custom_headers: Option<&HashMap, S>>) -> HeaderMap { let mut headers = HeaderMap::default(); - if let Some(def_headers) = defined_headers { + if let Some(def_headers) = request_headers { for (key, value) in def_headers { if let (Ok(key), Ok(value)) = (HeaderName::from_bytes(key.as_bytes()), HeaderValue::from_bytes(value.as_bytes())) { - headers.insert(key, value); + if filter_request_header(key.as_str()) { + headers.insert(key, value); + } } } } @@ -155,7 +157,7 @@ pub fn get_request_headers(defined_header let header_keys: HashSet = headers.keys().map(|k| k.as_str().to_lowercase()).collect(); for (key, value) in custom { let key_lc = key.to_lowercase(); - if "host" == key_lc || header_keys.contains(key_lc.as_str()) { + if header_keys.contains(key_lc.as_str()) { // debug_if_enabled!("Ignoring request header '{}={}'", key_lc, String::from_utf8_lossy(value)); } else if let (Ok(key), Ok(value)) = (HeaderName::from_bytes(key.as_bytes()), HeaderValue::from_bytes(value)) { headers.insert(key, value);