Fix stream session poisoning and correct headers handling

This commit is contained in:
euzu
2026-01-04 10:54:41 +01:00
parent f571e6781b
commit e5e94a5f28
4 changed files with 215 additions and 78 deletions
+58 -10
View File
@@ -760,6 +760,10 @@ pub async fn force_provider_stream_response(
let connection_permission = UserConnectionPermission::Allowed;
let item_type = stream_channel.item_type;
// Release the existing provider connection for this session before acquiring a new one.
// This is critical for users with a connection limit of 1 to avoid "Provider exhausted" or provider-side 502/509 errors during seeking.
app_state.connection_manager.release_provider_connection(&user_session.addr).await;
let stream_details = create_stream_response_details(
app_state,
&stream_options,
@@ -931,13 +935,29 @@ pub async fn stream_response(
axum::http::StatusCode::BAD_REQUEST.into_response()
}
} else {
let session_url = provider_response
.as_ref()
.and_then(|(_, _, u, _)| u.as_ref())
.map_or_else(
|| Cow::Borrowed(stream_url),
|url| Cow::Owned(url.to_string()),
);
// Previously, we would always check if the provider redirected the request.
// If the provider redirected a movie from /movie/... to a temporary /live/... URL,
// We would save that redirected URL in your session.
// When we tried to seek or pause/resume, we would use that saved /live/ URL.
// However, providers often make these redirect links ephemeral or restricted—they
// might not support seeking, or they might trigger a 509 error if accessed again.
// For Movies/Series: We now ignore the redirect and always save the original,
// canonical URL (the one starting with /movie/) in your session.
// This ensures that every time you seek, we start "fresh" with the correct provider handshake,
// preventing the session from being "poisoned" by a temporary redirect.
// For everything else (Live): It continues to work as before, using the redirected URL if available,
// which is often desirable for live streams to stay on the same edge server.
let session_url = if matches!(item_type, PlaylistItemType::Catchup | PlaylistItemType::Video | PlaylistItemType::LocalVideo | PlaylistItemType::Series | PlaylistItemType::LocalSeries) {
Cow::Borrowed(stream_url)
} else {
provider_response
.as_ref()
.and_then(|(_, _, u, _)| u.as_ref())
.map_or_else(
|| Cow::Borrowed(stream_url),
|url| Cow::Owned(url.to_string()),
)
};
if log_enabled!(log::Level::Debug) {
if session_url.eq(&stream_url) {
debug!("Streaming stream request from {}", sanitize_sensitive_info(stream_url)
@@ -1508,9 +1528,6 @@ pub async fn is_seek_request(cluster: XtreamCluster, req_headers: &HeaderMap) ->
.map(ToString::to_string);
if let Some(range) = range {
// if range.starts_with("bytes=0-") {
// return false;
// }
if range.starts_with("bytes=") {
return true;
}
@@ -1542,3 +1559,34 @@ pub fn json_or_bin_response<T: Serialize>(accept: Option<&str>, data: &T) -> imp
pub fn create_session_fingerprint(fingerprint: &str, username: &str, virtual_id: u32) -> String {
format!("{fingerprint}|{username}|{virtual_id}")
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderMap;
use shared::model::XtreamCluster;
#[tokio::test]
async fn test_is_seek_request() {
let mut headers = HeaderMap::new();
// No range header
assert!(!is_seek_request(XtreamCluster::Video, &headers).await);
// Range: bytes=0- (Should be true now to allow session takeover on restart)
headers.insert("range", "bytes=0-".parse().unwrap());
assert!(is_seek_request(XtreamCluster::Video, &headers).await);
// Range: bytes=100- (Should be true)
headers.insert("range", "bytes=100-".parse().unwrap());
assert!(is_seek_request(XtreamCluster::Video, &headers).await);
// Range: bytes=100-200 (Should be true)
headers.insert("range", "bytes=100-200".parse().unwrap());
assert!(is_seek_request(XtreamCluster::Video, &headers).await);
// Live cluster should always return false
headers.insert("range", "bytes=100-".parse().unwrap());
assert!(!is_seek_request(XtreamCluster::Live, &headers).await);
}
}
@@ -39,6 +39,7 @@ pub struct ProviderStreamFactoryOptions {
url: Url,
headers: HeaderMap,
range_bytes: Arc<Option<AtomicUsize>>,
range_requested: bool,
reconnect_flag: Arc<AtomicOnceFlag>,
}
@@ -61,15 +62,18 @@ impl ProviderStreamFactoryOptions {
};
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);
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);
let url = stream_url.clone();
let range_bytes = Arc::new(range_start_bytes.map(AtomicUsize::new));
let range_bytes = if matches!(item_type, PlaylistItemType::Live | PlaylistItemType::LiveUnknown) {
Arc::new(requested_range.map(AtomicUsize::new))
} else {
Arc::new(Some(AtomicUsize::new(requested_range.unwrap_or(0))))
};
Self {
// item_type,
@@ -83,6 +87,7 @@ impl ProviderStreamFactoryOptions {
url,
headers,
range_bytes,
range_requested: requested_range.is_some(),
}
}
@@ -158,6 +163,11 @@ impl ProviderStreamFactoryOptions {
self.reconnect_flag.is_active()
}
#[inline]
pub fn was_range_requested(&self) -> bool {
self.range_requested
}
}
fn get_request_range_start_bytes(req_headers: &HashMap<String, Vec<u8>>) -> Option<usize> {
@@ -221,11 +231,15 @@ fn prepare_client(
}
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);
if range > 0 || stream_options.was_range_requested() {
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
}
true
} else {
false
};
@@ -362,7 +376,7 @@ async fn get_provider_stream(
}
Err(status) => {
debug!("Provider stream response error status response : {status}");
if matches!(status, StatusCode::FORBIDDEN | StatusCode::SERVICE_UNAVAILABLE | StatusCode::UNAUTHORIZED) {
if matches!(status, StatusCode::FORBIDDEN | StatusCode::SERVICE_UNAVAILABLE | StatusCode::UNAUTHORIZED | StatusCode::RANGE_NOT_SATISFIABLE) {
warn!("The stream could be unavailable. ({status}) {}",sanitize_sensitive_info(stream_options.get_url().as_str()));
break;
}
@@ -520,56 +534,83 @@ pub async fn create_provider_stream(
}
}
// #[cfg(test)]
// mod tests {
// use crate::api::model::streams::provider_stream_factory::PlaylistItemType;
// use crate::api::model::streams::provider_stream_factory::{create_provider_stream, BufferStreamOptions};
// use actix_web::test;
// use actix_web::test::TestRequest;
// use actix_web::web;
// use actix_web::App;
// use actix_web::{HttpRequest, HttpResponse};
// use futures::StreamExt;
// use std::sync::Arc;
// use crate::model::Config;
//
// #[tokio::test]
// async fn test_stream() {
// let app = App::new().route("/test", web::get().to(test_stream_handler));
// let server = test::init_service(app).await;
// let req = TestRequest::get().uri("/test").to_request();
// let _response = test::call_service(&server, req).await;
// }
// async fn test_stream_handler(req: axum::http::Request<axum::body::Body>) -> impl axum::response::IntoResponse + Send {
// let cfg = Config::default();
// let mut counter = 5;
// let client = Arc::new(reqwest::Client::new());
// let url = url::Url::parse("https://info.cern.ch/hypertext/WWW/TheProject.html").unwrap();
// let input = None;
//
// let options = BufferStreamOptions::new(PlaylistItemType::Live, true, true, 0, false);
// let value = create_provider_stream(&cfg, &client, &url, &req, input, options);
// let mut values = value.await;
// 'outer: while let Some((ref mut stream, info)) = values.as_mut() {
// if info.is_some() {
// println!("{:?}", info.as_ref().unwrap());
// }
// while let Some(result) = stream.next().await {
// match result {
// Ok(bytes) => {
// println!("Received {} bytes {bytes:?}", bytes.len());
// counter -= 1;
// if counter < 0 {
// break 'outer;
// }
// }
// Err(err) => {
// eprintln!("Error occurred: {}", err);
// break 'outer;
// }
// }
// }
// }
// HttpResponse::Ok().finish()
// }
// }
#[cfg(test)]
mod tests {
use super::*;
use shared::model::PlaylistItemType;
use axum::http::HeaderMap;
#[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 disabled_headers = None;
// Case 1: VOD, no initial range requested
let mut req_headers = HeaderMap::new();
let options = ProviderStreamFactoryOptions::new(
addr,
PlaylistItemType::Video,
false,
&stream_options,
&stream_url,
&req_headers,
None,
disabled_headers,
);
assert!(!options.was_range_requested());
assert_eq!(options.get_total_bytes_send(), Some(0)); // Should track even if not requested
// Case 2: VOD, range requested
req_headers.insert("Range", "bytes=100-".parse().unwrap());
let options = ProviderStreamFactoryOptions::new(
addr,
PlaylistItemType::Video,
false,
&stream_options,
&stream_url,
&req_headers,
None,
disabled_headers,
);
assert!(options.was_range_requested());
assert_eq!(options.get_total_bytes_send(), Some(100));
// Case 3: Live, no initial range requested
let req_headers = HeaderMap::new();
let options = ProviderStreamFactoryOptions::new(
addr,
PlaylistItemType::Live,
false,
&stream_options,
&stream_url,
&req_headers,
None,
disabled_headers,
);
assert!(!options.was_range_requested());
assert_eq!(options.get_total_bytes_send(), None); // Should NOT track
// Case 4: Live, range requested (should be stripped)
let mut req_headers = HeaderMap::new();
req_headers.insert("Range", "bytes=100-".parse().unwrap());
let options = ProviderStreamFactoryOptions::new(
addr,
PlaylistItemType::Live,
false,
&stream_options,
&stream_url,
&req_headers,
None,
disabled_headers,
);
assert!(!options.was_range_requested()); // Stripped by filter
assert_eq!(options.get_total_bytes_send(), None);
}
}
+54 -7
View File
@@ -2,7 +2,7 @@ use futures::{StreamExt, TryStreamExt};
use log::{debug, error, log_enabled, trace, Level};
use reqwest::header::CONTENT_ENCODING;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use std::collections::{HashMap, HashSet};
use std::collections::HashMap;
use std::io::{Error, ErrorKind};
use std::path::{Path, PathBuf};
use std::pin::Pin;
@@ -221,6 +221,10 @@ pub fn get_client_request<S: ::std::hash::BuildHasher + Default>
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>>, disabled_headers: Option<&ReverseProxyDisabledHeaderConfig>) -> HeaderMap {
let mut headers = HeaderMap::default();
let mut has_user_agent = false;
// 1. First, we process the configured request headers (from input config).
// 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())) {
@@ -228,36 +232,49 @@ pub fn get_request_headers<S: ::std::hash::BuildHasher + Default>(request_header
if disabled_headers.as_ref().is_some_and(|d| d.should_remove(key.as_str())) {
continue;
}
if key == axum::http::header::USER_AGENT {
has_user_agent = true;
}
headers.insert(key, value);
}
}
}
}
// 2. Next, we process custom headers (from the client request).
// These are only added if they don't already exist in the headers map (i.e., not overridden by config).
if let Some(custom) = custom_headers {
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 filter_request_header(key_lc.as_str()) {
if disabled_headers.as_ref().is_some_and(|d| d.should_remove(key_lc.as_str())) {
continue;
}
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);
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 {
has_user_agent = true;
}
headers.insert(name, val);
}
}
}
}
}
if log_enabled!(Level::Trace) {
let he: HashMap<String, String> = 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:?}");
}
}
if !headers.contains_key(axum::http::header::USER_AGENT) {
// 3. Finally, if no User-Agent was provided by config OR client, use the default.
if !has_user_agent {
headers.insert(axum::http::header::USER_AGENT, HeaderValue::from_static(DEFAULT_USER_AGENT));
}
headers
}
@@ -671,5 +688,35 @@ mod tests {
let expected = "http://my.provider.com:8080";
assert_eq!(get_base_url_from_str(url).unwrap(), expected);
}
#[test]
fn test_get_request_headers_prioritization() {
use super::{get_request_headers, DEFAULT_USER_AGENT};
use std::collections::HashMap;
use axum::http::header::USER_AGENT;
// Case 1: No headers provided -> Default UA
let headers = get_request_headers::<std::collections::hash_map::RandomState>(None, None, None);
assert_eq!(headers.get(USER_AGENT).unwrap(), DEFAULT_USER_AGENT);
// Case 2: Only client header -> Client 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);
assert_eq!(headers.get(USER_AGENT).unwrap(), "Client-UA");
// Case 3: 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);
assert_eq!(headers.get(USER_AGENT).unwrap(), "Config-UA");
// Case 4: 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);
assert_eq!(headers.get("X-Test").unwrap(), "From-Config");
}
}
+1
View File
@@ -41,6 +41,7 @@ const SUPPORTED_RESPONSE_HEADERS: &[&str] = &[
"access-control-allow-origin",
"access-control-allow-credentials",
"icy-metadata",
"icy-metaint",
"referer",
"last-modified",
"cache-control",