mirror of
https://github.com/euzu/tuliprox.git
synced 2026-10-08 17:02:22 +02:00
Fix stream session poisoning and correct headers handling
This commit is contained in:
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user