diff --git a/src/api/api_utils.rs b/src/api/api_utils.rs index e82a57014..6e46d5faf 100644 --- a/src/api/api_utils.rs +++ b/src/api/api_utils.rs @@ -186,7 +186,7 @@ pub async fn stream_response(app_state: &AppState, stream_url: &str, fn shared_stream_response(app_state: &AppState, stream_url: &str, log_active_clients: bool, user: &ProxyUserCredentials) -> Option { if let Some(stream) = SharedStreamManager::subscribe_shared_stream(app_state, stream_url) { debug_if_enabled!("Using shared channel {}", sanitize_sensitive_info(stream_url)); - if let Some(headers) = app_state.shared_stream_manager.lock().get_shared_state_headers(stream_url) { + if let Some(headers) = app_state.shared_stream_manager.get_shared_state_headers(stream_url) { let mut response_builder = get_stream_response_with_headers(Some((headers.clone(), StatusCode::OK)), stream_url); let active_clients = Arc::clone(&app_state.active_users); let stream = ActiveClientStream::new(stream, active_clients, user, log_active_clients); diff --git a/src/api/main_api.rs b/src/api/main_api.rs index cb89be28a..aa751b4e7 100644 --- a/src/api/main_api.rs +++ b/src/api/main_api.rs @@ -2,7 +2,7 @@ use actix_cors::Cors; use actix_web::middleware::Logger; use actix_web::web::Data; use actix_web::{web, App, HttpResponse, HttpServer}; -use parking_lot::{Mutex as PlMutex, RwLock as PlRwLock}; +use parking_lot::{Mutex as PlMutex}; use tokio::sync::{RwLock, Mutex}; use log::{error, info}; use std::collections::{VecDeque}; @@ -41,7 +41,7 @@ fn get_web_dir_path(web_ui_enabled: bool, web_root: &str) -> Result,) -> HttpResponse { let ts = chrono::offset::Local::now().format("%Y-%m-%d %H:%M:%S").to_string(); let (active_clients, active_connections) = { - let active_user = app_state.active_users.read(); + let active_user = &app_state.active_users; (active_user.active_users(), active_user.active_connections()) }; HttpResponse::Ok().json(Healthcheck { @@ -75,8 +75,8 @@ fn create_shared_data(cfg: &Arc) -> Data { active: Arc::from(RwLock::new(None)), finished: Arc::from(RwLock::new(Vec::new())), }), - shared_stream_manager: Arc::new(PlMutex::new(SharedStreamManager::new())), - active_users: Arc::new(PlRwLock::new(ActiveUserManager::new())), + shared_stream_manager: Arc::new(SharedStreamManager::new()), + active_users: Arc::new(ActiveUserManager::new()), http_client: Arc::new(reqwest::Client::new()), cache, }) diff --git a/src/api/model/active_user_manager.rs b/src/api/model/active_user_manager.rs index 5ae9675b6..ba41702a3 100644 --- a/src/api/model/active_user_manager.rs +++ b/src/api/model/active_user_manager.rs @@ -1,8 +1,9 @@ use std::collections::HashMap; use std::sync::atomic::{AtomicU32, Ordering}; +use parking_lot::RwLock; pub struct ActiveUserManager { - pub user: HashMap, + pub user: RwLock>, } impl Default for ActiveUserManager { @@ -14,38 +15,46 @@ impl Default for ActiveUserManager { impl ActiveUserManager { pub fn new() -> Self { Self { - user: HashMap::new(), + user: RwLock::new(HashMap::new()), } } pub fn user_connections(&self, username: &str) -> u32 { - if let Some(counter) = self.user.get(username) { - return counter.load(std::sync::atomic::Ordering::Relaxed); + if let Some(counter) = self.user.read().get(username) { + return counter.load(std::sync::atomic::Ordering::SeqCst); } 0 } pub fn active_users(&self) -> usize { - self.user.len() + self.user.read().len() } pub fn active_connections(&self) -> usize { - self.user.values().map(|c| c.load(Ordering::Relaxed) as usize).sum() + self.user.read().values().map(|c| c.load(Ordering::SeqCst) as usize).sum() } - pub fn add_connection(&mut self, username: &str) { - if let Some(counter) = self.user.get(username) { - counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed); - } else { - self.user.insert(username.to_string(), AtomicU32::new(1)); - } - } - - pub fn remove_connection(&mut self, username: &str) { - if let Some(counter) = self.user.get(username) { - if counter.fetch_sub(1, std::sync::atomic::Ordering::Relaxed) == 1 { - self.user.remove(username); + pub fn add_connection(&self, username: &str) -> (usize, usize) { + { + let mut lock = self.user.write(); + if let Some(counter) = lock.get(username) { + counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + } else { + lock.insert(username.to_string(), AtomicU32::new(1)); } } + (self.active_users(), self.active_connections()) + } + + pub fn remove_connection(&self, username: &str) -> (usize, usize) { + { + let mut lock = self.user.write(); + if let Some(counter) = lock.get(username) { + if counter.fetch_sub(1, std::sync::atomic::Ordering::SeqCst) == 1 { + lock.remove(username); + } + } + } + (self.active_users(), self.active_connections()) } } \ No newline at end of file diff --git a/src/api/model/app_state.rs b/src/api/model/app_state.rs index 86971c2de..d85de486a 100644 --- a/src/api/model/app_state.rs +++ b/src/api/model/app_state.rs @@ -1,5 +1,5 @@ use std::sync::{Arc}; -use parking_lot::{Mutex, RwLock}; +use parking_lot::{Mutex}; use crate::api::model::active_user_manager::ActiveUserManager; use crate::api::model::download::DownloadQueue; use crate::api::model::streams::shared_stream_manager::SharedStreamManager; @@ -9,14 +9,14 @@ use crate::tools::lru_cache::LRUResourceCache; pub struct AppState { pub config: Arc, pub downloads: Arc, - pub shared_stream_manager: Arc>, + pub shared_stream_manager: Arc, pub http_client: Arc, pub cache: Arc>>, - pub active_users: Arc>, + pub active_users: Arc, } impl AppState { pub fn get_active_connections_for_user(&self, username: &str) -> u32 { - self.active_users.read().user_connections(username) + self.active_users.user_connections(username) } } diff --git a/src/api/model/stream_error.rs b/src/api/model/stream_error.rs index d170da183..db78d4396 100644 --- a/src/api/model/stream_error.rs +++ b/src/api/model/stream_error.rs @@ -1,10 +1,11 @@ +use tokio_stream::wrappers::errors::BroadcastStreamRecvError; #[derive(Debug, Clone)] pub enum StreamError { Reqwest(String), // StdIo(std::io::Error), // ReceiverClosed, - // ReceiverError(RecvError), + ReceiverError(BroadcastStreamRecvError), } impl StreamError { @@ -24,7 +25,7 @@ impl std::fmt::Display for StreamError { StreamError::Reqwest(e) => write!(f, "Reqwest error: {e}"), // StreamError::StdIo(e) => write!(f, "IO error: {e}"), // StreamError::ReceiverClosed => write!(f, "Receiver closed"), - // StreamError::ReceiverError(e) => write!(f, "Receiver error {e}"), + StreamError::ReceiverError(e) => write!(f, "Receiver error {e}"), } } } \ No newline at end of file diff --git a/src/api/model/streams/active_client_stream.rs b/src/api/model/streams/active_client_stream.rs index 9e58fb7e2..c7ffbff46 100644 --- a/src/api/model/streams/active_client_stream.rs +++ b/src/api/model/streams/active_client_stream.rs @@ -5,27 +5,24 @@ use crate::model::api_proxy::ProxyUserCredentials; use bytes::Bytes; use futures::Stream; use log::info; -use parking_lot::{RwLock}; use std::pin::Pin; use std::sync::Arc; use std::task::Poll; pub(in crate::api) struct ActiveClientStream { inner: ResponseStream, - active_clients: Arc>, + active_clients: Arc, log_active_clients: bool, username: String, } impl ActiveClientStream { - pub(crate) fn new(inner: ResponseStream, active_clients: Arc>, user: &ProxyUserCredentials, log_active_clients: bool) -> Self { - let client_count = { - let mut clients = active_clients.write(); - clients.add_connection(&user.username); - clients.active_users() + pub(crate) fn new(inner: ResponseStream, active_clients: Arc, user: &ProxyUserCredentials, log_active_clients: bool) -> Self { + let (client_count, connection_count) = { + active_clients.add_connection(&user.username) }; if log_active_clients { - info!("Active clients: {client_count}"); + info!("Active clients: {client_count}, active connections {connection_count}"); } Self { inner, active_clients, log_active_clients, username: user.username.clone() } } @@ -44,13 +41,11 @@ impl Stream for ActiveClientStream { impl Drop for ActiveClientStream { fn drop(&mut self) { - let client_count = { - let mut clients = self.active_clients.write(); - clients.remove_connection(&self.username); - clients.active_users() + let (client_count, connection_count) = { + self.active_clients.remove_connection(&self.username) }; if self.log_active_clients { - info!("Active clients: {client_count}"); + info!("Active clients: {client_count}, active connections {connection_count}"); } } } \ No newline at end of file diff --git a/src/api/model/streams/buffered_stream.rs b/src/api/model/streams/buffered_stream.rs index baf746392..7c358a16b 100644 --- a/src/api/model/streams/buffered_stream.rs +++ b/src/api/model/streams/buffered_stream.rs @@ -4,6 +4,7 @@ use std::{ pin::Pin, sync::Arc, }; +use std::time::Duration; use tokio::sync::mpsc::channel; use tokio_stream::wrappers::ReceiverStream; use crate::api::model::stream_error::StreamError; @@ -18,6 +19,7 @@ impl BufferedStream { let (tx, rx) = channel(buffer_size); actix_rt::spawn(async move { let mut stream = stream; + let sleep_duration= Duration::from_millis(100); loop { match stream.next().await { Some(Ok(chunk)) => { @@ -30,7 +32,9 @@ impl BufferedStream { break; } } - Some(Err(_err)) => {} + Some(Err(_err)) => { + actix_web::rt::time::sleep(sleep_duration).await; + } None => { break } diff --git a/src/api/model/streams/client_stream.rs b/src/api/model/streams/client_stream.rs index eec269dae..4a80b912d 100644 --- a/src/api/model/streams/client_stream.rs +++ b/src/api/model/streams/client_stream.rs @@ -39,7 +39,7 @@ impl Stream for ClientStream { } if let Some(counter) = self.total_bytes.as_ref() { - counter.fetch_add(bytes.len(), Ordering::Relaxed); + counter.fetch_add(bytes.len(), Ordering::SeqCst); } return Poll::Ready(Some(Ok(bytes))); diff --git a/src/api/model/streams/persist_pipe_stream.rs b/src/api/model/streams/persist_pipe_stream.rs index b94eb84d2..bba910a95 100644 --- a/src/api/model/streams/persist_pipe_stream.rs +++ b/src/api/model/streams/persist_pipe_stream.rs @@ -52,7 +52,7 @@ where fn on_complete(&mut self) { if !self.completed { self.completed = true; - let size = self.size.load(Ordering::Relaxed); + let size = self.size.load(Ordering::SeqCst); if self.writer.flush().is_ok() { (self.callback)(size); } @@ -61,7 +61,7 @@ where fn on_data(&mut self, data: &Result) { if let Ok(bytes) = data { - self.size.fetch_add(bytes.len(), Ordering::Relaxed); + self.size.fetch_add(bytes.len(), Ordering::SeqCst); let bytes_to_write = bytes.clone(); if let Err(e) = self.writer.write_all(&bytes_to_write) { error!("Error writing to resource file: {e}"); diff --git a/src/api/model/streams/provider_stream_factory.rs b/src/api/model/streams/provider_stream_factory.rs index 0f026b432..31f8bd965 100644 --- a/src/api/model/streams/provider_stream_factory.rs +++ b/src/api/model/streams/provider_stream_factory.rs @@ -125,7 +125,7 @@ impl ProviderStreamOptions { #[inline] pub fn get_total_bytes_send(&self) -> Option { - self.range_bytes.as_ref().as_ref().map(|atomic| atomic.load(Ordering::Relaxed)) + self.range_bytes.as_ref().as_ref().map(|atomic| atomic.load(Ordering::SeqCst)) } // pub fn get_range_bytes(&self) -> &Arc> { @@ -227,7 +227,7 @@ async fn stream_provider(client: Arc, stream_options: ProviderS let url = stream_options.get_url(); let range_start = stream_options.get_total_bytes_send(); let headers = stream_options.get_headers(); - + 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, url, headers, range_start); @@ -386,7 +386,7 @@ mod tests { 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 options = BufferStreamOptions::new(PlaylistItemType::Live, true, true, 0); let value = create_provider_stream(Arc::clone(&client), &url, &req, input, options); let mut values = value.await; 'outer: while let Some((ref mut stream, info)) = values.as_mut() { diff --git a/src/api/model/streams/shared_stream_manager.rs b/src/api/model/streams/shared_stream_manager.rs index 71aa527ad..5234f0ae1 100644 --- a/src/api/model/streams/shared_stream_manager.rs +++ b/src/api/model/streams/shared_stream_manager.rs @@ -3,103 +3,93 @@ use crate::api::model::streams::provider_stream_factory::STREAM_QUEUE_SIZE; use crate::api::model::stream_error::StreamError; use crate::utils::debug_if_enabled; use crate::utils::network::request::sanitize_sensitive_info; -use parking_lot::{FairMutex, Mutex}; +use parking_lot::{FairMutex}; use bytes::Bytes; use futures::stream::BoxStream; use futures::{Stream, StreamExt}; use std::collections::HashMap; use std::sync::Arc; -use tokio::sync::mpsc; -use tokio::sync::mpsc::error::TrySendError; -use tokio::sync::mpsc::{Sender}; -use tokio_stream::wrappers::ReceiverStream; +use tokio_stream::wrappers::BroadcastStream; use std::pin::Pin; use std::task::{Context, Poll}; +use std::time::Duration; -const MIN_STREAM_QUEUE_SIZE: usize = 1024; +const MIN_STREAM_QUEUE_SIZE: usize = 128; /// /// Wraps a `ReceiverStream` as Stream> /// -struct ReceiverStreamWrapper { - stream: S, +struct BroadcastStreamWrapper { + stream: BroadcastStream, } -impl Stream for ReceiverStreamWrapper -where - S: Stream + Unpin, -{ +impl Stream for BroadcastStreamWrapper { type Item = Result; fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { match Pin::new(&mut self.stream).poll_next(cx) { - Poll::Ready(Some(bytes)) => Poll::Ready(Some(Ok(bytes))), + Poll::Ready(Some(Ok(bytes))) => Poll::Ready(Some(Ok(bytes))), + Poll::Ready(Some(Err(_))) | Poll::Pending => Poll::Pending, Poll::Ready(None) => Poll::Ready(None), - Poll::Pending => Poll::Pending, } } } -fn convert_stream(stream: BoxStream) -> BoxStream> { - Box::pin(ReceiverStreamWrapper { stream }.boxed()) +fn convert_stream(stream: BroadcastStream) -> BoxStream<'static, Result> { + Box::pin(BroadcastStreamWrapper { stream }) } /// Represents the state of a shared provider URL. /// /// - `headers`: The initial connection headers used during the setup of the shared stream. -/// - `subscribers`: A list of clients that have subscribed to the shared stream. struct SharedStreamState { headers: Vec<(String, String)>, - buf_size: usize, - subscribers: Arc>>>, + sender: tokio::sync::broadcast::Sender, } impl SharedStreamState { fn new(headers: Vec<(String, String)>, buf_size: usize) -> Self { + let (sender, _) = tokio::sync::broadcast::channel(buf_size); Self { headers, - buf_size, - subscribers: Arc::new(FairMutex::new(Vec::new())), + sender, } } fn subscribe(&self) -> BoxStream<'static, Result> { - let (tx, rx) = mpsc::channel(self.buf_size); - self.subscribers.lock().push(tx); - convert_stream(ReceiverStream::new(rx).boxed()) + let rx = self.sender.subscribe(); + convert_stream(BroadcastStream::new(rx)).boxed() + // .map_err(StreamError::ReceiverError).boxed() } - fn broadcast(&self, stream_url: &str, bytes_stream: S, shared_streams: Arc>) + fn broadcast(&self, stream_url: &str, bytes_stream: S, shared_streams: Arc) where S: Stream> + Unpin + 'static, { let mut source_stream = Box::pin(bytes_stream); - let subscriber = Arc::clone(&self.subscribers); let streaming_url = stream_url.to_string(); + let sender = self.sender.clone(); // Spawn a task to forward items from the source stream to the broadcast channel actix_rt::spawn(async move { - while let Some(item) = source_stream.next().await { - if let Ok(data) = item { - let mut subs = subscriber.lock(); - if subs.len() > 0 { - (*subs).retain(|sender| { - match sender.try_send(data.clone()) { - Err(TrySendError::Closed(_)) => false, - Ok(()) | Err(_) => true, - // Err(TrySendError::Full(_)) => false, // Drop slow consumers - } - }); - } else { - debug_if_enabled!("No active subscribers. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url)); - // Cleanup for removing unused shared streams - shared_streams.lock().unregister(&streaming_url); - return; + let sleep_duration = Duration::from_millis(20); + loop { + match source_stream.next().await { + Some(Ok(data)) => { + if sender.receiver_count() == 0 { + debug_if_enabled!("No active subscribers. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url)); + break; + } + let _ = sender.send(data); + } + None | Some(Err(_)) => { + break; } } + actix_web::rt::time::sleep(sleep_duration).await; } - shared_streams.lock().unregister(&streaming_url); + shared_streams.unregister(&streaming_url); }); } } @@ -139,7 +129,6 @@ impl SharedStreamManager { shared_state.broadcast(stream_url, bytes_stream, Arc::clone(&app_state.shared_stream_manager)); app_state .shared_stream_manager - .lock() .shared_streams .lock() .insert(stream_url.to_string(), shared_state); @@ -153,7 +142,6 @@ impl SharedStreamManager { ) -> Option>> { if let Some(shared_stream) = app_state .shared_stream_manager - .lock() .shared_streams .lock() .get(stream_url) { diff --git a/src/api/scheduler.rs b/src/api/scheduler.rs index a066638f0..df39f32c8 100644 --- a/src/api/scheduler.rs +++ b/src/api/scheduler.rs @@ -54,7 +54,7 @@ mod tests { let expression = "0/1 * * * * * *"; // every second let runs = AtomicU8::new(0); - let run_me = || runs.fetch_add(1, Ordering::Relaxed); + let run_me = || runs.fetch_add(1, Ordering::SeqCst); let start = std::time::Instant::now(); match Schedule::from_str(expression) { @@ -66,7 +66,7 @@ mod tests { actix_web::rt::time::sleep_until(actix_rt::time::Instant::from(datetime_to_instant(datetime))).await; run_me(); } - if runs.load(Ordering::Relaxed) == 6 { + if runs.load(Ordering::SeqCst) == 6 { break; } } @@ -75,7 +75,7 @@ mod tests { }; let duration = start.elapsed(); - assert!(runs.load(Ordering::Relaxed) == 6, "Failed to run"); + assert!(runs.load(Ordering::SeqCst) == 6, "Failed to run"); assert!(duration.as_secs() > 4, "Failed time"); } } \ No newline at end of file diff --git a/src/processing/processor/playlist.rs b/src/processing/processor/playlist.rs index 015e5beb6..31e382138 100644 --- a/src/processing/processor/playlist.rs +++ b/src/processing/processor/playlist.rs @@ -267,7 +267,7 @@ fn map_playlist_counter(target: &ConfigTarget, playlist: &[PlaylistGroup]) { for channel in &plg.channels { let provider = ValueProvider { pli: RefCell::new(channel) }; if counter.filter.filter(&provider, &mut mock_processor) { - let cntval = counter.value.load(core::sync::atomic::Ordering::Relaxed); + let cntval = counter.value.load(core::sync::atomic::Ordering::SeqCst); let new_value = if counter.modifier == CounterModifier::Assign { cntval.to_string() } else { @@ -279,7 +279,7 @@ fn map_playlist_counter(target: &ConfigTarget, playlist: &[PlaylistGroup]) { } }; channel.header.borrow_mut().set_field(&counter.field, new_value.as_str()); - counter.value.fetch_add(1, core::sync::atomic::Ordering::Relaxed); + counter.value.fetch_add(1, core::sync::atomic::Ordering::SeqCst); } } } diff --git a/src/repository/target_id_mapping.rs b/src/repository/target_id_mapping.rs index 227067174..1f95336bd 100644 --- a/src/repository/target_id_mapping.rs +++ b/src/repository/target_id_mapping.rs @@ -98,7 +98,6 @@ impl TargetIdMapping { if let Some(record) = self.by_virtual_id.query(virtual_id) { if record.provider_id == provider_id && (record.item_type != item_type || record.parent_virtual_id != parent_virtual_id) { let new_record = VirtualIdRecord::new(provider_id, *virtual_id, item_type, parent_virtual_id, uuid); - println!("updating record {virtual_id} {record:?} {new_record:?} "); self.by_virtual_id.insert(*virtual_id, new_record); self.dirty = true; } diff --git a/src/tools/atomic_once_flag.rs b/src/tools/atomic_once_flag.rs index 982e4c49d..f13e36162 100644 --- a/src/tools/atomic_once_flag.rs +++ b/src/tools/atomic_once_flag.rs @@ -36,7 +36,7 @@ impl AtomicOnceFlag { /// Creates a new `AtomicOnceFlag` with a default memory ordering of `Relaxed`. pub fn new() -> Self { - Self::with_ordering(Ordering::Relaxed) + Self::with_ordering(Ordering::SeqCst) } /// Disables the flag. After calling this method, `is_active()` will always return `false`.