diff --git a/src/api/api_utils.rs b/src/api/api_utils.rs index f0b690393..559655c83 100644 --- a/src/api/api_utils.rs +++ b/src/api/api_utils.rs @@ -23,6 +23,7 @@ use std::sync::Arc; use async_std::sync::Mutex; use futures::stream::BoxStream; use futures::{TryStreamExt}; +use reqwest::StatusCode; use url::Url; use crate::api::model::model_utils::get_stream_response_with_headers; use crate::api::model::persist_pipe_stream::PersistPipeStream; @@ -72,7 +73,7 @@ async fn create_broadcast_stream( let notify_stream_url = stream_url.to_string(); // Acquire lock and check for existing stream let shared_streams = app_state.shared_streams.lock().await; - if let Some(shared_stream) = shared_streams.get(¬ify_stream_url) { + if let Some((_, shared_stream)) = shared_streams.get(¬ify_stream_url) { Some(shared_stream.get_receiver()) } else { None @@ -86,7 +87,7 @@ pub async fn stream_response(app_state: &AppState, stream_url: &str, let share_stream = is_stream_share_enabled(item_type, target); if share_stream { - if let Some(value) = shared_stream_response(app_state, stream_url, None).await { + if let Some(value) = shared_stream_response(app_state, stream_url).await { return value; } } @@ -116,7 +117,8 @@ pub async fn stream_response(app_state: &AppState, stream_url: &str, if let Some(stream) = stream_opt { let use_buffer = !buffer_enabled || direct_pipe_provider_stream; return if share_stream { - SharedStream::register(app_state, stream_url, stream, use_buffer).await; + let shared_headers = provider_response.as_ref().map_or_else(Vec::new, |(h, _)| h.clone()); + SharedStream::register(app_state, stream_url, stream, use_buffer, shared_headers).await; if let Some(broadcast_stream) = create_broadcast_stream(app_state, stream_url).await { let body_stream = BodyStream::new(broadcast_stream); let mut response_builder = get_stream_response_with_headers(provider_response, stream_url); @@ -134,11 +136,11 @@ pub async fn stream_response(app_state: &AppState, stream_url: &str, HttpResponse::BadRequest().finish() } -async fn shared_stream_response(app_state: &AppState, stream_url: &str, headers: Option<(Vec<(String, String)>, reqwest::StatusCode)>) -> Option { +async fn shared_stream_response(app_state: &AppState, stream_url: &str) -> Option { if let Some(stream) = create_broadcast_stream(app_state, stream_url).await { debug_if_enabled!("Using shared channel {}", mask_sensitive_info(stream_url)); - if app_state.shared_streams.lock().await.get(stream_url).is_some() { - let mut response_builder = get_stream_response_with_headers(headers, stream_url); + if let Some((headers,_)) = app_state.shared_streams.lock().await.get(stream_url) { + let mut response_builder = get_stream_response_with_headers(Some((headers.clone(), StatusCode::OK)), stream_url); let current_date = Utc::now().format("%a, %d %b %Y %H:%M:%S GMT").to_string(); response_builder.insert_header((DATE, current_date.as_bytes())); // response_builder.insert_header((ACCEPT_RANGES, "bytes".as_bytes())); diff --git a/src/api/model/app_state.rs b/src/api/model/app_state.rs index 2735171cb..3775b560a 100644 --- a/src/api/model/app_state.rs +++ b/src/api/model/app_state.rs @@ -6,10 +6,12 @@ use crate::api::model::shared_stream::SharedStream; use crate::model::config::{Config}; use crate::utils::lru_cache::LRUResourceCache; +type SharedStreamState = (Vec<(String, String)>, SharedStream); + pub struct AppState { pub config: Arc, pub downloads: Arc, - pub shared_streams: Arc>>, + pub shared_streams: Arc>>, pub http_client: Arc, pub cache: Arc>> } diff --git a/src/api/model/buffered_stream.rs b/src/api/model/buffered_stream.rs index a725ef0a1..a44bee5a9 100644 --- a/src/api/model/buffered_stream.rs +++ b/src/api/model/buffered_stream.rs @@ -7,7 +7,7 @@ use std::{ use tokio::sync::mpsc::channel; use tokio_stream::wrappers::ReceiverStream; use crate::api::model::stream_error::StreamError; -use crate::utils::atomic_flag::AtomicOnceFlag; +use crate::utils::atomic_once_flag::AtomicOnceFlag; pub(in crate::api::model) struct BufferedStream { stream: ReceiverStream>, @@ -26,7 +26,7 @@ impl BufferedStream { permit.send(Ok(chunk)); } else { // receiver closed. - client_close_signal.disable(); + client_close_signal.notify(); break; } } diff --git a/src/api/model/client_stream.rs b/src/api/model/client_stream.rs index d586eaff0..9bd5079a8 100644 --- a/src/api/model/client_stream.rs +++ b/src/api/model/client_stream.rs @@ -7,7 +7,7 @@ use std::task::{Poll}; use log::debug; use futures::{Stream}; use crate::api::model::stream_error::StreamError; -use crate::utils::atomic_flag::AtomicOnceFlag; +use crate::utils::atomic_once_flag::AtomicOnceFlag; use crate::utils::request_utils::mask_sensitive_info; /// This stream counts the send bytes for reconnecting to the actual position and @@ -45,7 +45,7 @@ impl Stream for ClientStream { return Poll::Ready(Some(Ok(bytes))); } Poll::Ready(None) => { - self.close_signal.disable(); + self.close_signal.notify(); return Poll::Ready(None); } other => return other, @@ -58,6 +58,6 @@ impl Stream for ClientStream { impl Drop for ClientStream { fn drop(&mut self) { debug!("Client disconnected {}", mask_sensitive_info(&self.url)); - self.close_signal.disable(); + self.close_signal.notify(); } } \ No newline at end of file diff --git a/src/api/model/model_utils.rs b/src/api/model/model_utils.rs index 6278bc063..342a6235c 100644 --- a/src/api/model/model_utils.rs +++ b/src/api/model/model_utils.rs @@ -29,9 +29,7 @@ pub fn get_stream_response_with_headers(custom: Option<(Vec<(String, String)>, S let default_headers = vec![ (actix_web::http::header::CONTENT_TYPE, HeaderValue::from_str("application/octet-stream").unwrap()), - // (actix_web::http::header::CONTENT_LENGTH, HeaderValue::from(0)), (actix_web::http::header::CONNECTION, HeaderValue::from_str("keep-alive").unwrap()), - //(actix_web::http::header::CACHE_CONTROL, HeaderValue::from_str("no-cache").unwrap()), (actix_web::http::header::VARY, HeaderValue::from_str("accept-encoding").unwrap()) ]; diff --git a/src/api/model/provider_stream_factory.rs b/src/api/model/provider_stream_factory.rs index 7644389bf..822afd70e 100644 --- a/src/api/model/provider_stream_factory.rs +++ b/src/api/model/provider_stream_factory.rs @@ -19,7 +19,7 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; use std::time::{Duration, Instant}; use url::Url; -use crate::utils::atomic_flag::AtomicOnceFlag; +use crate::utils::atomic_once_flag::AtomicOnceFlag; // TODO make this configurable pub const STREAM_QUEUE_SIZE: usize = 1024; // mpsc channel holding messages. with 8092byte chunks and 2Mbit/s approx 8MB @@ -104,7 +104,7 @@ impl ProviderStreamOptions { #[inline] pub fn cancel_reconnect(&self) { - self.continue_flag.disable(); + self.continue_flag.notify(); } #[inline] @@ -317,12 +317,12 @@ pub async fn create_provider_stream(client: Arc, let stream_options = create_provider_stream_options(stream_url, req, input, &options); let client_stream_factory = |stream, reconnect, range_cnt| { - let stream = ClientStream::new(stream, reconnect, range_cnt, stream_options.get_url().as_str()).boxed(); - if stream_options.is_buffered() { + let stream = if stream_options.is_buffered() { BufferedStream::new(stream, stream_options.get_buffer_size(), stream_options.get_continue_flag_clone(), stream_url.as_str()).boxed() } else { stream - } + }; + ClientStream::new(stream, reconnect, range_cnt, stream_options.get_url().as_str()).boxed() }; match get_initial_stream(Arc::clone(&client), &stream_options).await { diff --git a/src/api/model/shared_stream.rs b/src/api/model/shared_stream.rs index bbe4f03b1..a2bc46cf1 100644 --- a/src/api/model/shared_stream.rs +++ b/src/api/model/shared_stream.rs @@ -29,6 +29,7 @@ impl SharedStream { stream_url: &str, bytes_stream: S, use_buffer: bool, + headers: Vec<(String, String)>, ) where S: Stream> + Unpin + 'static, { @@ -43,10 +44,7 @@ impl SharedStream { .await .insert( stream_url.to_string(), - SharedStream { - // sender: sender.clone(), - receiver: rx - }, + (headers, SharedStream {receiver: rx, }), ); let shared_streams_map = Arc::clone(&app_state.shared_streams); diff --git a/src/utils/atomic_flag.rs b/src/utils/atomic_once_flag.rs similarity index 98% rename from src/utils/atomic_flag.rs rename to src/utils/atomic_once_flag.rs index 49981b25a..2f78b5fb7 100644 --- a/src/utils/atomic_flag.rs +++ b/src/utils/atomic_once_flag.rs @@ -37,7 +37,7 @@ impl AtomicOnceFlag { /// Disables the flag. After calling this method, `is_active()` will always return `false`. /// /// This operation is atomic and uses the specified memory ordering. - pub fn disable(&self) { + pub fn notify(&self) { self.enabled.store(false, self.ordering); } diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 65f8067a8..1275290fb 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -13,7 +13,7 @@ pub mod directed_graph; pub mod lru_cache; pub mod size_utils; pub mod sys; -pub mod atomic_flag; +pub mod atomic_once_flag; #[macro_export] macro_rules! debug_if_enabled {