Refactored stream handling

This commit is contained in:
euzu
2025-05-09 17:13:37 +02:00
parent 85b608e68c
commit 9f884703fa
8 changed files with 391 additions and 313 deletions
+90 -38
View File
@@ -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<ProviderName>, StreamUrl),
GracePeriod(Option<ProviderName>, 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<ProviderConnectionGuard>, StreamingOption, Option<HashMap<String, String>>) {
struct StreamingStrategy {
provider_connection_guard: Option<ProviderConnectionGuard>,
provider_stream_state: ProviderStreamState,
input_headers: Option<HashMap<String, String>>,
}
/// 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<u32> {
@@ -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
+1 -1
View File
@@ -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;
+3 -4
View File
@@ -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 {
+4 -4
View File
@@ -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;
+2 -63
View File
@@ -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<String, String>>,
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::<Vec<_>>());
// 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<String, String>>,
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)
}
}
}
+255 -194
View File
@@ -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<AtomicOnceFlag>
pipe_stream: bool,
url: Url,
headers: HeaderMap,
range_bytes: Arc<Option<AtomicUsize>>,
reconnect_flag: Arc<AtomicOnceFlag>,
}
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<String, String>>,
) -> 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<AtomicOnceFlag>,
url: Url,
reconnect: bool,
reconnect_force_secs: u32,
headers: HeaderMap,
range_bytes: Arc<Option<AtomicUsize>>,
}
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<AtomicOnceFlag> {
Arc::clone(&self.continue_flag)
}
// #[inline]
// pub fn get_continue_flag(&self) -> &Arc<AtomicFlag> {
// &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<AtomicOnceFlag>,
// url: Url,
// pipe_stream: bool,
// reconnect: bool,
// reconnect_force_secs: u32,
// headers: HeaderMap,
// range_bytes: Arc<Option<AtomicUsize>>,
// }
//
// 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<AtomicOnceFlag> {
// Arc::clone(&self.continue_flag)
// }
//
// // #[inline]
// // pub fn get_continue_flag(&self) -> &Arc<AtomicFlag> {
// // &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<usize> {
// self.range_bytes.as_ref().as_ref().map(|atomic| atomic.load(Ordering::SeqCst))
// }
//
// // pub fn get_range_bytes(&self) -> &Arc<Option<AtomicUsize>> {
// // &self.range_bytes
// // }
//
// #[inline]
// pub fn get_range_bytes_clone(&self) -> Arc<Option<AtomicUsize>> {
// 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<String, Vec<u8>>) -> Option<usize> {
// 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<String, Vec<u8>>) -> Opti
None
}
fn get_client_stream_request_params(
req_headers: &HeaderMap,
input_headers: Option<&HashMap<String, String>>,
options: &BufferStreamOptions) -> (usize, Option<usize>, 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::<Vec<_>>());
// 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<String> {
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<reqwest::Client>, 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<reqwest::Client>, 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::<Vec<_>>());
let request_builder = request_client.get(url.clone()).headers(headers);
(request_builder, partial)
}
async fn provider_initial_request(cfg: &Config, request_client: Arc<reqwest::Client>, stream_options: &ProviderStreamOptions) -> Result<Option<ProviderStreamFactoryResponse>, StatusCode> {
async fn provider_stream_request(cfg: &Config, request_client: Arc<reqwest::Client>, stream_options: &ProviderStreamFactoryOptions) -> Result<Option<ProviderStreamFactoryResponse>, 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<reqwest::Cli
// debug!("First headers {headers:?} {} {}", sanitize_sensitive_info(url.as_str()));
Some((response_headers, response.status()))
};
return Ok(Some((response.bytes_stream().map_err(|err| {
//error!("Failed to read response body: {err}");
let provider_stream = response.bytes_stream().map_err(|err| {
// error!("Stream error {err}");
StreamError::reqwest(&err)
}).boxed(), response_info)));
}
if let (Some(boxed_provider_stream), response_info) =
create_channel_unavailable_stream(cfg, &get_response_headers(response.headers()), status)
{
}).boxed();
let boxed_provider_stream = if stream_options.get_reconnect_force_secs() > 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<reqwest::Cli
}
}
async fn stream_provider(client: Arc<reqwest::Client>, stream_options: ProviderStreamOptions) -> Option<BoxedProviderStream> {
async fn get_provider_stream(cfg: &Config, client: Arc<reqwest::Client>, stream_options: &ProviderStreamFactoryOptions) -> Result<Option<ProviderStreamFactoryResponse>, 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<reqwest::Client>, stream_options: &ProviderStreamOptions) -> Option<ProviderStreamFactoryResponse> {
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<reqwest::Client>, 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<String, String>>,
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<Config>,
client: Arc<reqwest::Client>,
stream_url: &Url,
req_headers: &HeaderMap,
input_headers: Option<&HashMap<String, String>>,
options: BufferStreamOptions) -> Option<ProviderStreamFactoryResponse> {
let stream_options = create_provider_stream_options(stream_url, req_headers, input_headers, &options);
stream_options: ProviderStreamFactoryOptions) -> Option<ProviderStreamFactoryResponse> {
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
}
}
}
+29 -4
View File
@@ -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,
+7 -5
View File
@@ -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<S: ::std::hash::BuildHasher + Default>(client: &Arc<re
request.headers(headers)
}
pub fn get_request_headers<S: ::std::hash::BuildHasher + Default>(defined_headers: Option<&HashMap<String, String, S>>, custom_headers: Option<&HashMap<String, Vec<u8>, S>>) -> HeaderMap {
pub fn get_request_headers<S: ::std::hash::BuildHasher + Default>(request_headers: Option<&HashMap<String, String, S>>, custom_headers: Option<&HashMap<String, Vec<u8>, 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<S: ::std::hash::BuildHasher + Default>(defined_header
let header_keys: HashSet<String> = 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);