diff --git a/Cargo.lock b/Cargo.lock index cb8e3d3e8..0bf16ce65 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -183,6 +183,7 @@ checksum = "021e862c184ae977658b36c4500f7feac3221ca5da43e3f25bd04ab6c79a29b5" dependencies = [ "axum-core", "axum-macros", + "base64 0.22.1", "bytes", "form_urlencoded", "futures-util", @@ -202,8 +203,10 @@ dependencies = [ "serde_json", "serde_path_to_error", "serde_urlencoded", + "sha1", "sync_wrapper", "tokio", + "tokio-tungstenite", "tower", "tower-layer", "tower-service", @@ -655,6 +658,12 @@ dependencies = [ "parking_lot_core", ] +[[package]] +name = "data-encoding" +version = "2.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a2330da5de22e8a3cb63252ce2abb30116bf5265e89c0e01bc17015ce30a476" + [[package]] name = "deranged" version = "0.4.0" @@ -899,6 +908,8 @@ name = "frontend" version = "0.1.0" dependencies = [ "anyhow", + "bincode 2.0.1", + "bytes", "futures", "futures-signals", "gloo-storage 0.3.0", @@ -3295,6 +3306,17 @@ dependencies = [ "unsafe-libyaml", ] +[[package]] +name = "sha1" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +dependencies = [ + "cfg-if", + "cpufeatures", + "digest", +] + [[package]] name = "sha2" version = "0.10.9" @@ -3311,8 +3333,10 @@ name = "shared" version = "0.1.0" dependencies = [ "base64 0.22.1", + "bincode 2.0.1", "bitflags 2.9.1", "blake3", + "bytes", "chrono", "enum-iterator", "fastrand 2.3.0", @@ -3651,6 +3675,18 @@ dependencies = [ "tokio-util", ] +[[package]] +name = "tokio-tungstenite" +version = "0.26.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a9daff607c6d2bf6c16fd681ccb7eecc83e4e2cdc1ca067ffaadfca5de7f084" +dependencies = [ + "futures-util", + "log", + "tokio", + "tungstenite", +] + [[package]] name = "tokio-util" version = "0.7.15" @@ -3864,6 +3900,23 @@ dependencies = [ "winapi", ] +[[package]] +name = "tungstenite" +version = "0.26.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4793cb5e56680ecbb1d843515b23b6de9a75eb04b66643e256a396d43be33c13" +dependencies = [ + "bytes", + "data-encoding", + "http 1.3.1", + "httparse", + "log", + "rand", + "sha1", + "thiserror 2.0.12", + "utf-8", +] + [[package]] name = "twox-hash" version = "2.1.1" @@ -3929,6 +3982,12 @@ version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "daf8dba3b7eb870caf1ddeed7bc9d2a049f3cfdfae7cb521b087cc33ae4c49da" +[[package]] +name = "utf-8" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" + [[package]] name = "utf8_iter" version = "1.0.4" diff --git a/backend/Cargo.toml b/backend/Cargo.toml index 398c3dca9..a25fc75ce 100644 --- a/backend/Cargo.toml +++ b/backend/Cargo.toml @@ -17,9 +17,9 @@ url = "2.5" reqwest = { version = "0", features = ["blocking", "json", "stream", "rustls-tls", "socks"] } chrono = "0.4" cron = "0.15" -axum = { version = "0" , features = ["macros", "default"]} +axum = { version = "0" , features = ["macros", "default", "ws"]} tower = "0" -tower-http = { version = "0", features = ["cors", "auth", "fs", "compression-full"]} +tower-http = { version = "0", features = ["cors", "auth", "fs", "compression-full", "trace"] } tower_governor = { version = "0.7", features = ["axum"] } jsonwebtoken = "9.3" rust-argon2 = "2.1" @@ -34,12 +34,12 @@ mime = "0.3" log = "0.4" env_logger = "0.11" rustelebot = "0.3" -bincode = { version = "2.0.1", features = ["std", "serde"] } +bincode = { version = "2", features = ["std", "serde"] } rand = "0.9" rpassword = "7.4" flate2 = "1" blake3 = "1.8" -bytes = "1.10" +bytes = "1" tokio-stream = { version = "0.1", features = ["sync"] } tokio = { version = "1.45", features = ["rt-multi-thread", "parking_lot", "fs"] } tokio-util = "0.7" diff --git a/backend/src/api/api_utils.rs b/backend/src/api/api_utils.rs index 2d9967632..a746fc21e 100644 --- a/backend/src/api/api_utils.rs +++ b/backend/src/api/api_utils.rs @@ -240,7 +240,7 @@ pub struct StreamDetails { pub input_name: Option, pub grace_period_millis: u64, pub reconnect_flag: Option>, - pub provider_connection_guard: Option, + pub provider_connection_guard: Option>, } impl StreamDetails { @@ -266,7 +266,7 @@ impl StreamDetails { } struct StreamingStrategy { - provider_connection_guard: Option, + provider_connection_guard: Option>, provider_stream_state: ProviderStreamState, input_headers: Option>, } @@ -294,7 +294,10 @@ async fn resolve_streaming_strategy(app_state: &AppState, stream_url: &str, addr Some(provider) => app_state.active_provider.force_exact_acquire_connection(provider, addr).await, None => app_state.active_provider.acquire_connection(&input.name, addr).await }; - let stream_response_params = match &*provider_connection_guard { + + error!("{:?}", app_state.active_provider.active_connections().await); + + let stream_response_params = match &**provider_connection_guard { ProviderAllocation::Exhausted => { debug!("Input {} is exhausted. No connections allowed.", input.name); let stream = create_provider_connections_exhausted_stream(&app_state.app_config, &[]); @@ -310,7 +313,7 @@ async fn resolve_streaming_strategy(app_state: &AppState, stream_url: &str, addr (provider.name.to_string(), get_stream_alternative_url(stream_url, input, provider)) }; - if matches!(&*provider_connection_guard, ProviderAllocation::Available(_, _)) { + if matches!(&**provider_connection_guard, ProviderAllocation::Available(_, _)) { ProviderStreamState::Available(Some(provider), url) } else { ProviderStreamState::GracePeriod(Some(provider), url) @@ -358,7 +361,7 @@ async fn create_stream_response_details(app_state: &AppState, input_name: None, grace_period_millis, reconnect_flag: None, - provider_connection_guard: streaming_strategy.provider_connection_guard.take(), + provider_connection_guard: streaming_strategy.provider_connection_guard.clone(), } } ProviderStreamState::Available(provider_name, request_url) | @@ -592,7 +595,7 @@ pub async fn stream_response(addr: &str, if stream_details.has_stream() { // let content_length = get_stream_content_length(provider_response.as_ref()); let provider_response = stream_details.stream_info.as_ref().map(|(h, sc, response_url)| (h.clone(), *sc, response_url.clone())); - let provider_name = stream_details.provider_connection_guard.as_ref().and_then(ProviderConnectionGuard::get_provider_name); + let provider_name = stream_details.provider_connection_guard.as_ref().and_then(|guard| guard.get_provider_name()); let provider_guard = if share_stream { stream_details.provider_connection_guard.take() } else { None }; let stream = ActiveClientStream::new(stream_details, app_state, user, connection_permission, addr); @@ -780,10 +783,10 @@ pub fn get_username_from_auth_header( app_state: &Arc, ) -> Option { if let Some(web_auth_config) = &app_state.app_config.config.load().web_ui.as_ref().and_then(|c| c.auth.as_ref()) { - let secret_key: &str = web_auth_config.secret.as_ref(); + let secret_key: &[u8] = web_auth_config.secret.as_ref(); if let Ok(token_data) = decode::( token, - &DecodingKey::from_secret(secret_key.as_bytes()), + &DecodingKey::from_secret(secret_key), &Validation::new(Algorithm::HS256), ) { return Some(token_data.claims.username); diff --git a/backend/src/api/endpoints/mod.rs b/backend/src/api/endpoints/mod.rs index 5b61f0597..1beddb251 100644 --- a/backend/src/api/endpoints/mod.rs +++ b/backend/src/api/endpoints/mod.rs @@ -7,4 +7,5 @@ pub(in crate::api) mod web_index; pub(in crate::api) mod hls_api; mod user_api; pub(in crate::api) mod hdhomerun_api; -mod api_playlist_utils; \ No newline at end of file +mod api_playlist_utils; +pub (in crate::api) mod websocket_api; \ No newline at end of file diff --git a/backend/src/api/endpoints/v1_api.rs b/backend/src/api/endpoints/v1_api.rs index ddd89ba5d..6451a9fd9 100644 --- a/backend/src/api/endpoints/v1_api.rs +++ b/backend/src/api/endpoints/v1_api.rs @@ -289,7 +289,7 @@ async fn create_ipinfo_check(app_state: &Arc) -> Option<(Option) -> StatusCheck { +pub async fn create_status_check(app_state: &Arc) -> StatusCheck { let cache = match app_state.cache.load().as_ref().as_ref() { None => None, Some(lock) => { @@ -339,7 +339,6 @@ async fn ipinfo(axum::extract::State(app_state): axum::extract::State, web_ui_path: &str) -> axum::Router> { let mut router = axum::Router::new(); router = router diff --git a/backend/src/api/endpoints/websocket_api.rs b/backend/src/api/endpoints/websocket_api.rs new file mode 100644 index 000000000..0182881c1 --- /dev/null +++ b/backend/src/api/endpoints/websocket_api.rs @@ -0,0 +1,147 @@ +use std::sync::Arc; +use axum::{ + extract::ws::{WebSocketUpgrade, WebSocket, Message}, + response::IntoResponse, +}; +use axum::extract::ws::CloseFrame; +use log::{error, info}; +use shared::model::{ProtocolHandler, ProtocolMessage, WsCloseCode, PROTOCOL_VERSION}; +use crate::api::endpoints::v1_api::create_status_check; +use crate::api::model::app_state::AppState; +use crate::auth::verify_token; + +// WebSocket upgrade handler +async fn websocket_handler( + axum::extract::State(app_state): axum::extract::State>, + ws: WebSocketUpgrade) -> impl IntoResponse { + info!("Websocket connected"); + ws.on_upgrade(move |socket| handle_socket(socket, app_state, false)) +} + +// WebSocket upgrade handler +async fn websocket_handler_auth( + axum::extract::State(app_state): axum::extract::State>, + ws: WebSocketUpgrade) -> impl IntoResponse { + info!("Websocket connected"); + ws.on_upgrade(move |socket| handle_socket(socket, app_state, true)) +} + +pub fn ws_api_register(web_auth_enabled: bool, web_ui_path: &str) -> axum::Router> { + if web_auth_enabled { + axum::Router::new().route(&format!("{web_ui_path}/ws"), axum::routing::get(websocket_handler_auth)) + } else { + axum::Router::new().route(&format!("{web_ui_path}/ws"), axum::routing::get(websocket_handler)) + } +} + + +// WebSocket communication logic +async fn handle_socket(mut socket: WebSocket, app_state: Arc, auth: bool) { + let secret_key = if auth { + if let Some(web_auth_config) = &app_state.app_config.config.load().web_ui.as_ref().and_then(|c| c.auth.as_ref()) { + let secret_key: &[u8] = web_auth_config.secret.as_ref(); + Some(secret_key.to_vec()) + } else { + None + } + } else { + None + }; + + let verify_auth_token = |auth_token: &str| { + secret_key.as_ref().map(|key| verify_token(auth_token, key.as_slice())) + }; + + let mut active_user_change_rx = app_state.active_users.get_active_user_change_channel(); + let mut active_provider_change_rx = app_state.active_provider.get_active_provider_change_channel(); + + + let mut handler = ProtocolHandler::Version(PROTOCOL_VERSION); + + loop { + tokio::select! { + maybe_msg = socket.recv() => { + match maybe_msg { + Some(Ok(msg)) => { + match handler { + ProtocolHandler::Version(version) => { + let mut version_error = true; + if let Message::Binary(bytes) = msg { + if bytes.len() == 1 { + let client_version = bytes[0]; + if version == client_version { + if socket.send(Message::binary(bytes)).await.is_err() { + error!("Error sending websocket message"); + } else { + version_error = false; + handler = ProtocolHandler::Default; + } + } else { + error!("Version mismatch: server={version}, client={client_version}"); + } + } + } + if version_error { + let _ = socket.send(Message::Close(Some(CloseFrame { + code: WsCloseCode::Protocol.code(), + reason: "Unsupported protocol".into(), + }))).await; + break; + } + } + + ProtocolHandler::Default => { + if let Message::Binary(bytes) = msg { + match ProtocolMessage::from_bytes(bytes) { + Ok(ProtocolMessage::StatusRequest(auth_token)) => { + if !auth || verify_auth_token(&auth_token).is_some() { + let status = create_status_check(&app_state).await; + if let Ok(response) = ProtocolMessage::StatusResponse(status).to_bytes() { + if socket.send(Message::Binary(response)).await.is_err() { + error!("Failed to send websocket status response"); + } + } + } + } + Ok(_) => { + error!("Unexpected protocol message after handshake"); + } + Err(err) => { + error!("Invalid websocket message: {err}"); + } + } + } + } + } + } + Some(Err(err)) => { + error!("WebSocket error: {err}"); + break; + } + None => { + // WebSocket closed + break; + } + } + } + + Ok((user_count, connection_count)) = active_user_change_rx.recv() => { + if let Ok(payload) = ProtocolMessage::ActiveUserResponse(user_count, connection_count).to_bytes() { + if let Err(e) = socket.send(Message::Binary(payload)).await { + error!("Failed to send active user change: {e}"); + break; + } + } + } + + Ok((provider, connection_count)) = active_provider_change_rx.recv() => { + if let Ok(payload) = ProtocolMessage::ActiveProviderResponse(provider, connection_count).to_bytes() { + if let Err(e) = socket.send(Message::Binary(payload)).await { + error!("Failed to send active user change: {e}"); + break; + } + } + } + } + } +} \ No newline at end of file diff --git a/backend/src/api/main_api.rs b/backend/src/api/main_api.rs index bbda2db1a..c915ea2f6 100644 --- a/backend/src/api/main_api.rs +++ b/backend/src/api/main_api.rs @@ -24,6 +24,7 @@ use tokio_util::sync::CancellationToken; use tower_governor::key_extractor::SmartIpKeyExtractor; use crate::api::api_utils::{get_build_time, get_server_time}; use crate::api::config_watch::exec_config_watch; +use crate::api::endpoints::websocket_api::ws_api_register; use crate::api::serve::serve; use crate::VERSION; @@ -130,7 +131,7 @@ pub(in crate::api) fn start_hdhomerun(app_config: &Arc, app_state: &A let router = axum::Router::>::new() .layer(create_cors_layer()) .layer(create_compression_layer()) - // .layer(TraceLayer::new_for_http()) // `Logger::default()` + //.layer(tower_http::trace::TraceLayer::new_for_http()) // `Logger::default()` .merge(hdhr_api_register(basic_auth)); let router: axum::Router<()> = router.with_state(Arc::new(HdHomerunAppState { @@ -199,7 +200,8 @@ pub async fn start_server(app_config: Arc, targets: Arc, targets: Arc = router.with_state(shared_data.clone()); diff --git a/backend/src/api/model/active_provider_manager.rs b/backend/src/api/model/active_provider_manager.rs index 6d20c346e..c1ddb8123 100644 --- a/backend/src/api/model/active_provider_manager.rs +++ b/backend/src/api/model/active_provider_manager.rs @@ -1,19 +1,19 @@ -use crate::api::model::provider_config::{ProviderConfig, ProviderConfigConnection, ProviderConfigWrapper}; +use crate::api::model::provider_config::{ConnectionChangeSender, ProviderConfig, ProviderConfigConnection, ProviderConfigWrapper}; use crate::model::{AppConfig, ConfigInput}; use arc_swap::ArcSwap; use dashmap::DashMap; use log::{debug, log_enabled, trace}; use shared::utils::{default_grace_period_millis, default_grace_period_timeout_secs}; -use std::collections::HashMap; +use std::collections::{HashMap}; use std::ops::Deref; use std::sync::atomic::{AtomicU64, AtomicU8, AtomicUsize, Ordering}; use std::sync::Arc; +use tokio::sync::broadcast::Receiver; const CONNECTION_STATE_ACTIVE: u8 = 0; const CONNECTION_STATE_SHARED: u8 = 1; const CONNECTION_STATE_RELEASED: u8 = 2; -#[derive(Clone)] pub struct ProviderConnectionGuard { allocation: ProviderAllocation, } @@ -36,7 +36,7 @@ impl ProviderConnectionGuard { ProviderAllocation::Available(state, config) | ProviderAllocation::GracePeriod(state, config) => { // we can't release shared state - if state.compare_exchange(CONNECTION_STATE_ACTIVE, CONNECTION_STATE_RELEASED, Ordering::SeqCst, Ordering::SeqCst).is_ok() { + if state.compare_exchange(CONNECTION_STATE_ACTIVE, CONNECTION_STATE_RELEASED, Ordering::SeqCst, Ordering::SeqCst) .is_ok() { let provider_config = Arc::clone(config); trace!("Releasing provider connection {:?}", provider_config.name); tokio::spawn(async move { @@ -47,12 +47,12 @@ impl ProviderConnectionGuard { } } - // wen need to ensure the connections is released + // we need to ensure the connections is released pub(crate) fn force_release(&self) { match &self.allocation { ProviderAllocation::Exhausted => {} - ProviderAllocation::Available(state, config) | - ProviderAllocation::GracePeriod(state, config) => { + ProviderAllocation::Available(state, config) + | ProviderAllocation::GracePeriod(state, config) => { if state.load(Ordering::SeqCst) < CONNECTION_STATE_RELEASED { state.store(CONNECTION_STATE_RELEASED, Ordering::SeqCst); let provider_config = Arc::clone(config); @@ -106,20 +106,20 @@ impl Drop for ProviderConnectionGuard { } } -#[derive(Debug, Clone)] +#[derive(Debug)] pub enum ProviderAllocation { Exhausted, - Available(Arc, Arc), - GracePeriod(Arc, Arc), + Available(AtomicU8, Arc), + GracePeriod(AtomicU8, Arc), } impl ProviderAllocation { pub fn new_available(config: Arc) -> Self { - ProviderAllocation::Available(Arc::new(AtomicU8::new(CONNECTION_STATE_ACTIVE)), config) + ProviderAllocation::Available(AtomicU8::new(CONNECTION_STATE_ACTIVE), config) } pub fn new_grace_period(config: Arc) -> Self { - ProviderAllocation::GracePeriod(Arc::new(AtomicU8::new(CONNECTION_STATE_ACTIVE)), config) + ProviderAllocation::GracePeriod(AtomicU8::new(CONNECTION_STATE_ACTIVE), config) } } @@ -175,12 +175,12 @@ struct SingleProviderLineup { } impl SingleProviderLineup { - fn new<'a, F>(cfg: &ConfigInput, get_connection: Option) -> Self + fn new<'a, F>(cfg: &ConfigInput, get_connection: Option, connection_change_sender: ConnectionChangeSender) -> Self where F: Fn(&str) -> Option<&'a ProviderConfigConnection>, { Self { - provider: ProviderConfigWrapper::new(ProviderConfig::new(cfg, get_connection)), + provider: ProviderConfigWrapper::new(ProviderConfig::new(cfg, get_connection, connection_change_sender)), } } @@ -236,14 +236,14 @@ struct MultiProviderLineup { } impl MultiProviderLineup { - pub fn new<'a, F>(input: &ConfigInput, get_connection: Option) -> Self + pub fn new<'a, F>(input: &ConfigInput, get_connection: Option, connection_change_sender: &ConnectionChangeSender) -> Self where F: Fn(&str) -> Option<&'a ProviderConfigConnection> + Copy, { - let mut inputs = vec![ProviderConfigWrapper::new(ProviderConfig::new(input, get_connection))]; + let mut inputs = vec![ProviderConfigWrapper::new(ProviderConfig::new(input, get_connection, connection_change_sender.clone()))]; if let Some(aliases) = &input.aliases { for alias in aliases { - inputs.push(ProviderConfigWrapper::new(ProviderConfig::new_alias(input, alias, get_connection))); + inputs.push(ProviderConfigWrapper::new(ProviderConfig::new_alias(input, alias, get_connection, connection_change_sender.clone()))); } } let mut providers = HashMap::new(); @@ -457,24 +457,32 @@ struct ProviderLineupManager { grace_period_timeout_secs: AtomicU64, inputs: Arc>>>, providers: Arc>>, + connection_change_tx: ConnectionChangeSender, } + impl ProviderLineupManager { - pub fn new(inputs: Vec>, grace_period_millis: u64, grace_period_timeout_secs: u64) -> Self { - let lineups = inputs.iter().map(|i| Self::create_lineup(i, None)).collect(); + pub fn new(inputs: Vec>, grace_period_millis: u64, grace_period_timeout_secs: u64, connection_change_tx: ConnectionChangeSender) -> Self { + let lineups = inputs.iter().map(|i| Self::create_lineup(i, None, connection_change_tx.clone())).collect(); Self { grace_period_millis: AtomicU64::new(grace_period_millis), grace_period_timeout_secs: AtomicU64::new(grace_period_timeout_secs), inputs: Arc::new(ArcSwap::from_pointee(inputs)), providers: Arc::new(ArcSwap::from_pointee(lineups)), + connection_change_tx } } - fn create_lineup(input: &ConfigInput, provider_connections: Option<&HashMap<&str, ProviderConfigConnection>>) -> ProviderLineup { + + pub fn get_active_provider_change_channel(&self) -> Receiver<(String, usize)> { + self.connection_change_tx.subscribe() + } + + fn create_lineup(input: &ConfigInput, provider_connections: Option<&HashMap<&str, ProviderConfigConnection>>, connection_change_sender: ConnectionChangeSender) -> ProviderLineup { let get_connections = provider_connections.map(|c| |name: &str| c.get(name)); if input.aliases.as_ref().is_some_and(|a| !a.is_empty()) { - ProviderLineup::Multi(MultiProviderLineup::new(input, get_connections)) + ProviderLineup::Multi(MultiProviderLineup::new(input, get_connections, &connection_change_sender)) } else { - ProviderLineup::Single(SingleProviderLineup::new(input, get_connections)) + ProviderLineup::Single(SingleProviderLineup::new(input, get_connections, connection_change_sender)) } } @@ -571,7 +579,7 @@ impl ProviderLineupManager { let mut new_lineups: Vec = Vec::with_capacity(new_inputs.len()); let connections = Some(provider_connections); for input in &new_inputs { - new_lineups.push(Self::create_lineup(input, connections.as_ref())); + new_lineups.push(Self::create_lineup(input, connections.as_ref(), self.connection_change_tx.clone())); } debug!("inputs {new_inputs:?}"); @@ -612,18 +620,18 @@ impl ProviderLineupManager { None } - async fn force_exact_acquire_connection(&self, provider_name: &str) -> ProviderConnectionGuard { + async fn force_exact_acquire_connection(&self, provider_name: &str) -> Arc { let providers = self.providers.load(); let allocation = match Self::get_provider_config(provider_name, &providers) { None => ProviderAllocation::Exhausted, // No Name matched, we don't have this provider Some((_lineup, config)) => config.force_allocate().await, }; - ProviderConnectionGuard::new(allocation) + Arc::new(ProviderConnectionGuard::new(allocation)) } // Returns the next available provider connection - async fn acquire_connection(&self, input_name: &str) -> ProviderConnectionGuard { + async fn acquire_connection(&self, input_name: &str) -> Arc { let providers = self.providers.load(); let allocation = match Self::get_provider_config(input_name, &providers) { None => ProviderAllocation::Exhausted, // No Name matched, we don't have this provider @@ -641,7 +649,7 @@ impl ProviderLineupManager { } } - ProviderConnectionGuard::new(allocation) + Arc::new(ProviderConnectionGuard::new(allocation)) } // This method is used for redirects to cycle through provider @@ -719,15 +727,17 @@ impl ProviderLineupManager { pub struct ActiveProviderManager { providers: ProviderLineupManager, - connections: DashMap, + connections: DashMap>, } impl ActiveProviderManager { pub fn new(cfg: &AppConfig) -> Self { let (grace_period_millis, grace_period_timeout_secs) = Self::get_grace_options(cfg); let inputs = Self::get_config_inputs(cfg); + let (connection_change_tx, _) = tokio::sync::broadcast::channel(10); + Self { - providers: ProviderLineupManager::new(inputs, grace_period_millis, grace_period_timeout_secs), + providers: ProviderLineupManager::new(inputs, grace_period_millis, grace_period_timeout_secs, connection_change_tx), connections: DashMap::new(), } } @@ -750,14 +760,14 @@ impl ActiveProviderManager { self.providers.update_config(inputs, grace_period_millis, grace_period_timeout_secs).await; } - pub async fn force_exact_acquire_connection(&self, provider_name: &str, addr: &str) -> ProviderConnectionGuard { + pub async fn force_exact_acquire_connection(&self, provider_name: &str, addr: &str) -> Arc { let guard = self.providers.force_exact_acquire_connection(provider_name).await; self.register_connection(addr, &guard); guard } // Returns the next available provider connection - pub async fn acquire_connection(&self, input_name: &str, addr: &str) -> ProviderConnectionGuard { + pub async fn acquire_connection(&self, input_name: &str, addr: &str) -> Arc { let guard = self.providers.acquire_connection(input_name).await; self.register_connection(addr, &guard); guard @@ -781,10 +791,10 @@ impl ActiveProviderManager { self.providers.is_over_limit(provider_name).await } - fn register_connection(&self, addr: &str, guard: &ProviderConnectionGuard) { + fn register_connection(&self, addr: &str, guard: &Arc) { if !matches!(guard.allocation, ProviderAllocation::Exhausted) { trace!("Added provider connection {:?}", guard.get_provider_name().unwrap_or_default()); - self.connections.insert(addr.to_string(), guard.clone()); + self.connections.insert(addr.to_string(), Arc::clone(guard)); } } @@ -793,6 +803,10 @@ impl ActiveProviderManager { guard.release(); } } + + pub fn get_active_provider_change_channel(&self) -> tokio::sync::broadcast::Receiver<(String, usize)> { + self.providers.get_active_provider_change_channel() + } } #[cfg(test)] diff --git a/backend/src/api/model/active_user_manager.rs b/backend/src/api/model/active_user_manager.rs index 32c82b6d0..a9aabacb0 100644 --- a/backend/src/api/model/active_user_manager.rs +++ b/backend/src/api/model/active_user_manager.rs @@ -18,10 +18,12 @@ macro_rules! active_user_manager_shared_impl { } fn log_active_user(&self) { + let user = Arc::clone(&self.user); + let user_connection_count = Self::get_active_connections(&user); + let user_count = user.len(); + let _= self.active_user_change_tx.send((user_count, user_connection_count)); if self.is_log_user_enabled() { - let user = Arc::clone(&self.user); - let user_count = user.len(); - let user_connection_count = Self::get_active_connections(&user); + info!("Active Users: {user_count}, Active User Connections: {user_connection_count}"); } } @@ -71,6 +73,7 @@ struct ConnectionGuardUserManager { user_by_addr: Arc>, shared_stream_manager: Arc, provider_manager: Arc, + active_user_change_tx: tokio::sync::broadcast::Sender<(usize, usize)>, } impl ConnectionGuardUserManager { @@ -145,6 +148,7 @@ pub struct ActiveUserManager { close_signal_tx: tokio::sync::broadcast::Sender, shared_stream_manager: Arc, provider_manager: Arc, + active_user_change_tx: tokio::sync::broadcast::Sender<(usize, usize)>, } impl ActiveUserManager { @@ -152,6 +156,7 @@ impl ActiveUserManager { let log_active_user = config.log.as_ref().is_some_and(|l| l.log_active_user); let (grace_period_millis, grace_period_timeout_secs) = get_grace_options(config); let (close_signal_tx, _) = tokio::sync::broadcast::channel(10); + let (active_user_change_tx, _) = tokio::sync::broadcast::channel(10); Self { grace_period_millis: AtomicU64::new(grace_period_millis), grace_period_timeout_secs: AtomicU64::new(grace_period_timeout_secs), @@ -162,6 +167,7 @@ impl ActiveUserManager { close_signal_tx, shared_stream_manager: Arc::clone(shared_stream_manager), provider_manager: Arc::clone(provider_manager), + active_user_change_tx, } } @@ -182,6 +188,7 @@ impl ActiveUserManager { user_by_addr: Arc::clone(&self.user_by_addr), shared_stream_manager: Arc::clone(&self.shared_stream_manager), provider_manager: Arc::clone(&self.provider_manager), + active_user_change_tx: self.active_user_change_tx.clone(), } } @@ -355,6 +362,10 @@ impl ActiveUserManager { self.close_signal_tx.subscribe() } + pub fn get_active_user_change_channel(&self) -> tokio::sync::broadcast::Receiver<(usize, usize)> { + self.active_user_change_tx.subscribe() + } + pub fn get_user_session(&self, username: &str, token: &str) -> Option { self.update_user_session(username, token) } diff --git a/backend/src/api/model/provider_config.rs b/backend/src/api/model/provider_config.rs index de5c9bf07..f905c9174 100644 --- a/backend/src/api/model/provider_config.rs +++ b/backend/src/api/model/provider_config.rs @@ -1,12 +1,14 @@ use crate::api::model::active_provider_manager::ProviderAllocation; use crate::model::{ConfigInput, ConfigInputAlias, InputUserInfo}; use jsonwebtoken::get_current_timestamp; -use log::debug; +use log::{debug}; use std::ops::Deref; use std::sync::Arc; use tokio::sync::RwLock; use shared::model::InputType; +pub type ConnectionChangeSender = tokio::sync::broadcast::Sender<(String, usize)>; + #[derive(Debug, Clone, Copy)] pub enum ProviderConfigAllocation { Exhausted, @@ -39,6 +41,7 @@ pub struct ProviderConfig { max_connections: usize, priority: i16, connection: RwLock, + connection_change_tx: tokio::sync::broadcast::Sender<(String, usize)>, } impl PartialEq for ProviderConfig { @@ -55,8 +58,19 @@ impl PartialEq for ProviderConfig { } } +macro_rules! modify_connections { + ($self:ident, $guard:ident, +1) => {{ + $guard.current_connections += 1; + $self.notify_connection_change($guard.current_connections); + }}; + ($self:ident, $guard:ident, -1) => {{ + $guard.current_connections -= 1; + $self.notify_connection_change($guard.current_connections); + }}; +} + impl ProviderConfig { - pub fn new<'a, F>(cfg: &ConfigInput, get_connection: Option) -> Self + pub fn new<'a, F>(cfg: &ConfigInput, get_connection: Option, connection_change_tx: tokio::sync::broadcast::Sender<(String, usize)>) -> Self where F: Fn(&str) -> Option<&'a ProviderConfigConnection>, { @@ -70,10 +84,11 @@ impl ProviderConfig { max_connections: cfg.max_connections as usize, priority: cfg.priority, connection: RwLock::new(get_connection.and_then(|f| f(cfg.name.as_str())).map_or_else(Default::default, Clone::clone)), + connection_change_tx } } - pub fn new_alias<'a, F>(cfg: &ConfigInput, alias: &ConfigInputAlias, get_connection: Option) -> Self + pub fn new_alias<'a, F>(cfg: &ConfigInput, alias: &ConfigInputAlias, get_connection: Option, connection_change_tx: ConnectionChangeSender) -> Self where F: Fn(&str) -> Option<&'a ProviderConfigConnection>, { @@ -87,6 +102,7 @@ impl ProviderConfig { max_connections: alias.max_connections as usize, priority: alias.priority, connection: RwLock::new(get_connection.and_then(|f| f(alias.name.as_str())).map_or_else(Default::default, Clone::clone)), + connection_change_tx, } } @@ -94,6 +110,10 @@ impl ProviderConfig { InputUserInfo::new(self.input_type, self.username.as_deref(), self.password.as_deref(), &self.url) } + fn notify_connection_change(&self, new_connections: usize) { + self.connection_change_tx.send((self.name.clone(), new_connections)).unwrap(); + } + #[inline] pub async fn is_exhausted(&self) -> bool { let max = self.max_connections; @@ -133,15 +153,16 @@ impl ProviderConfig { // } + async fn force_allocate(&self) { let mut guard = self.connection.write().await; - guard.current_connections += 1; + modify_connections!(self, guard, +1); } async fn try_allocate(&self, grace: bool, grace_period_timeout_secs: u64) -> ProviderConfigAllocation { let mut guard = self.connection.write().await; if self.max_connections == 0 { - guard.current_connections += 1; + modify_connections!(self, guard, +1); return ProviderConfigAllocation::Available; } let connections = guard.current_connections; @@ -149,7 +170,7 @@ impl ProviderConfig { if connections < self.max_connections { guard.granted_grace = false; guard.grace_ts = 0; - guard.current_connections += 1; + modify_connections!(self, guard, +1); return ProviderConfigAllocation::Available; } @@ -166,7 +187,7 @@ impl ProviderConfig { } guard.granted_grace = true; guard.grace_ts = now; - guard.current_connections += 1; + modify_connections!(self, guard, +1); return ProviderConfigAllocation::GracePeriod; } ProviderConfigAllocation::Exhausted @@ -205,7 +226,7 @@ impl ProviderConfig { pub async fn release(&self) { let mut guard = self.connection.write().await; if guard.current_connections > 0 { - guard.current_connections -= 1; + modify_connections!(self, guard, -1); } if guard.current_connections == 0 || guard.current_connections < self.max_connections { diff --git a/backend/src/api/model/streams/active_client_stream.rs b/backend/src/api/model/streams/active_client_stream.rs index ca6689c88..79fe46fd5 100644 --- a/backend/src/api/model/streams/active_client_stream.rs +++ b/backend/src/api/model/streams/active_client_stream.rs @@ -1,7 +1,6 @@ use crate::api::api_utils::StreamDetails; use crate::api::model::active_provider_manager::{ActiveProviderManager, ProviderConnectionGuard}; -use crate::api::model::active_user_manager::ActiveUserManager; -use crate::api::model::active_user_manager::UserConnectionGuard; +use crate::api::model::active_user_manager::{ActiveUserManager, UserConnectionGuard}; use crate::api::model::app_state::AppState; use crate::api::model::stream::BoxedProviderStream; use crate::api::model::stream_error::StreamError; @@ -29,7 +28,7 @@ pub(in crate::api) struct ActiveClientStream { #[allow(unused)] user_connection_guard: Option, #[allow(dead_code)] - provider_connection_guard: Option, + provider_connection_guard: Option>, custom_video: (Option, Option), waker: Arc>>, } @@ -51,7 +50,7 @@ impl ActiveClientStream { let cfg = &app_state.app_config; let waker = Arc::new(Mutex::new(None)); let waker_clone = Arc::clone(&waker); - let grace_stop_flag = Self::stream_grace_period(&stream_details, grant_user_grace_period, user, &active_user, &active_provider, &waker_clone); + let grace_stop_flag = Self::stream_grace_period(&stream_details, grant_user_grace_period, user, addr, &active_user, &active_provider, &waker_clone); let custom_response = cfg.custom_stream_response.load(); let custom_video = custom_response.as_ref() .map_or((None, None), |c| @@ -86,6 +85,7 @@ impl ActiveClientStream { fn stream_grace_period(stream_details: &StreamDetails, user_grace_period: bool, user: &ProxyUserCredentials, + addr: &str, active_user: &Arc, active_provider: &Arc, waker: &Arc>>) -> Option> { @@ -113,6 +113,7 @@ impl ActiveClientStream { let waker_copy = Arc::clone(waker); let grace_period_millis = stream_details.grace_period_millis; + let address = addr.to_string(); tokio::spawn(async move { tokio::time::sleep(tokio::time::Duration::from_millis(grace_period_millis)).await; @@ -135,6 +136,7 @@ impl ActiveClientStream { if let Some((provider_name, provider_manager, reconnect_flag)) = provider_grace_check { if provider_manager.is_over_limit(&provider_name).await { info!("Provider connections exhausted for active clients: {provider_name}"); + provider_manager.release_connection(&address); stream_strategy_flag_copy.store(PROVIDER_EXHAUSTED_STREAM, std::sync::atomic::Ordering::SeqCst); if let Some(flag) = reconnect_flag { info!("Stopped reconnecting, provider connections exhausted {provider_name}"); diff --git a/backend/src/api/model/streams/shared_stream_manager.rs b/backend/src/api/model/streams/shared_stream_manager.rs index b3169f5e5..03ad5cfc5 100644 --- a/backend/src/api/model/streams/shared_stream_manager.rs +++ b/backend/src/api/model/streams/shared_stream_manager.rs @@ -10,7 +10,7 @@ use tokio::sync::mpsc::Sender; use crate::api::model::stream::BoxedProviderStream; use dashmap::DashMap; -use log::trace; +use log::{trace}; use shared::utils::sanitize_sensitive_info; use std::pin::Pin; use std::task::{Context, Poll}; @@ -53,13 +53,21 @@ type SubscriberId = String; struct SharedStreamState { headers: Vec<(String, String)>, buf_size: usize, - provider_guard: Option, + provider_guard: Option>, subscribers: Arc>>, //Arc>>>, } +impl Drop for SharedStreamState { + fn drop(&mut self) { + if let Some(guard) = self.provider_guard.as_ref() { + guard.force_release(); + } + } +} + impl SharedStreamState { fn new(headers: Vec<(String, String)>, buf_size: usize, - provider_guard: Option) -> Self { + provider_guard: Option>) -> Self { if let Some(guard) = &provider_guard { guard.disable_release(); } @@ -211,8 +219,13 @@ impl SharedStreamManager { } fn subscribe_stream(&self, stream_url: &str, addr: &str) -> Option { - let stream_data = self.shared_streams.get(stream_url)?.subscribe(addr); - Some(stream_data) + match self.shared_streams.get(stream_url) { + None => None, + Some(stream_state) => { + debug_if_enabled!("Responding to existing shared client stream {}", sanitize_sensitive_info(stream_url)); + Some(stream_state.subscribe(addr)) + } + } } fn register(&self, stream_url: &str, shared_state: SharedStreamState) { @@ -226,7 +239,7 @@ impl SharedStreamManager { bytes_stream: S, headers: Vec<(String, String)>, buffer_size: usize, - provider_guard: Option) -> Option + provider_guard: Option>) -> Option where S: Stream> + Unpin + 'static + Send, E: std::fmt::Debug + Send, @@ -245,7 +258,6 @@ impl SharedStreamManager { stream_url: &str, addr: &str, ) -> Option { - debug_if_enabled!("Responding existing shared client stream {}", sanitize_sensitive_info(stream_url)); app_state.shared_stream_manager.subscribe_stream(stream_url, addr) } } \ No newline at end of file diff --git a/backend/src/processing/playlist_watch.rs b/backend/src/processing/playlist_watch.rs index 6ef0a6a90..01cd46bfe 100644 --- a/backend/src/processing/playlist_watch.rs +++ b/backend/src/processing/playlist_watch.rs @@ -3,10 +3,10 @@ use std::path::{Path}; use std::sync::Arc; use log::{error, info}; use shared::model::{MsgKind, PlaylistGroup}; +use shared::utils::{bincode_deserialize, bincode_serialize}; use crate::messaging::{send_message}; use crate::model::Config; use crate::utils; -use crate::utils::{bincode_deserialize, bincode_serialize}; pub fn process_group_watch(client: &Arc, cfg: &Config, target_name: &str, pl: &PlaylistGroup) { let mut new_tree = BTreeSet::new(); diff --git a/backend/src/processing/processor/xtream_series.rs b/backend/src/processing/processor/xtream_series.rs index b3250c88e..daf3a5fff 100644 --- a/backend/src/processing/processor/xtream_series.rs +++ b/backend/src/processing/processor/xtream_series.rs @@ -17,9 +17,9 @@ use std::io::{BufWriter, Write}; use std::sync::Arc; use std::time::Instant; use log::{info, log_enabled, Level}; +use shared::utils::bincode_serialize; use crate::model::{XtreamSeriesEpisode, XtreamSeriesInfoEpisode}; use crate::utils; -use crate::utils::bincode_serialize; use crate::processing::processor::xtream::normalize_json_content; create_resolve_options_function_for_xtream_target!(series); diff --git a/backend/src/repository/bplustree.rs b/backend/src/repository/bplustree.rs index 4c53a76c1..d5e532224 100644 --- a/backend/src/repository/bplustree.rs +++ b/backend/src/repository/bplustree.rs @@ -9,8 +9,8 @@ use ruzstd::decoding::StreamingDecoder; use ruzstd::encoding::{compress_to_vec, CompressionLevel}; use serde::{Deserialize, Serialize}; use tempfile::NamedTempFile; +use shared::utils::{bincode_deserialize, bincode_serialize}; use crate::utils; -use crate::utils::{bincode_deserialize, bincode_serialize}; const BLOCK_SIZE: usize = 4096; const BINCODE_OVERHEAD: usize = 8; diff --git a/backend/src/repository/indexed_document.rs b/backend/src/repository/indexed_document.rs index ae2c17335..f393e26a7 100644 --- a/backend/src/repository/indexed_document.rs +++ b/backend/src/repository/indexed_document.rs @@ -9,8 +9,8 @@ use log::error; use serde::{Deserialize, Serialize}; use tempfile::NamedTempFile; use shared::error::{str_to_io_error, to_io_error}; +use shared::utils::{bincode_deserialize, bincode_serialize}; use crate::utils; -use crate::utils::{bincode_deserialize, bincode_serialize}; const BLOCK_SIZE: usize = 4096; const LEN_SIZE: usize = 4; diff --git a/backend/src/repository/xtream_repository.rs b/backend/src/repository/xtream_repository.rs index e9b68b527..342450fe6 100644 --- a/backend/src/repository/xtream_repository.rs +++ b/backend/src/repository/xtream_repository.rs @@ -10,7 +10,6 @@ use crate::repository::storage::{get_input_storage_path, get_target_id_mapping_f use crate::repository::storage_const; use crate::repository::target_id_mapping::VirtualIdRecord; use crate::repository::xtream_playlist_iterator::XtreamPlaylistJsonIterator; -use crate::utils::bincode_deserialize; use crate::utils::FileReadGuard; use crate::utils::file_reader; use crate::utils::open_readonly_file; @@ -26,7 +25,7 @@ use std::fs::File; use std::io::{BufReader, BufWriter, Error, ErrorKind, Read, Write}; use std::path::{Path, PathBuf}; use shared::model::{PlaylistEntry, PlaylistGroup, PlaylistItem, PlaylistItemType, XtreamCluster, XtreamPlaylistItem}; -use shared::utils::{generate_playlist_uuid, get_u32_from_serde_value, hex_encode, json_iter_array}; +use shared::utils::{bincode_deserialize, generate_playlist_uuid, get_u32_from_serde_value, hex_encode, json_iter_array}; macro_rules! cant_write_result { ($path:expr, $err:expr) => { diff --git a/backend/src/utils/mod.rs b/backend/src/utils/mod.rs index b235b5461..e5174e78a 100644 --- a/backend/src/utils/mod.rs +++ b/backend/src/utils/mod.rs @@ -2,7 +2,6 @@ mod sys_utils; mod compression; mod file; mod network; -mod bincode_utils; mod crypto_utils; mod step_measure; mod logging; @@ -64,6 +63,5 @@ pub use self::sys_utils::*; pub use self::compression::*; pub use self::file::*; pub use self::network::*; -pub use self::bincode_utils::*; pub use self::crypto_utils::*; pub use self::step_measure::*; diff --git a/shared/Cargo.toml b/shared/Cargo.toml index 0e5d2b653..db7fadd55 100644 --- a/shared/Cargo.toml +++ b/shared/Cargo.toml @@ -20,3 +20,5 @@ blake3 = "1" fastrand = "2" zeroize = "1" chrono = "0.4.41" +bytes = "1" +bincode = "2" diff --git a/shared/src/model/mod.rs b/shared/src/model/mod.rs index a69af716a..b142859c3 100644 --- a/shared/src/model/mod.rs +++ b/shared/src/model/mod.rs @@ -11,6 +11,7 @@ pub mod xtream_const; mod auth; mod status_check; mod ip_check; +mod web_socket; pub use self::cluster_flags::*; pub use self::playlist::*; @@ -24,3 +25,4 @@ pub use self::mapping::*; pub use self::auth::*; pub use self::status_check::*; pub use self::ip_check::*; +pub use self::web_socket::*; diff --git a/shared/src/model/status_check.rs b/shared/src/model/status_check.rs index 8ead50762..2c2b36784 100644 --- a/shared/src/model/status_check.rs +++ b/shared/src/model/status_check.rs @@ -15,4 +15,20 @@ pub struct StatusCheck { pub active_user_connections: usize, #[serde(skip_serializing_if = "Option::is_none")] pub active_provider_connections: Option>, +} + +impl Default for StatusCheck { + fn default() -> Self { + Self { + status: "n/a".to_string(), + version: "n/a".to_string(), + build_time: None, + server_time: "n/a".to_string(), + memory: "n/a".to_string(), + cache: None, + active_users: 0, + active_user_connections: 0, + active_provider_connections: None, + } + } } \ No newline at end of file diff --git a/shared/src/model/web_socket.rs b/shared/src/model/web_socket.rs new file mode 100644 index 000000000..acfa1f9ec --- /dev/null +++ b/shared/src/model/web_socket.rs @@ -0,0 +1,79 @@ +use std::io; +use bytes::Bytes; +use crate::model::StatusCheck; +use crate::utils::{bincode_deserialize, bincode_serialize}; +use serde::{Deserialize, Serialize}; + +pub const PROTOCOL_VERSION: u8 = 1; + +pub enum ProtocolHandler { + Version(u8), + Default, +} + +pub enum WsCloseCode { + // Normal, + // Away, + Protocol, + // Unsupported, + // Abnormal, + // Invalid, + // Policy, + // Size, + // Extension, + // Error, + // Restart, + // Again, + // Tls, +} + +impl WsCloseCode { + pub fn code(&self) -> u16 { + match self { + // WsCloseCode::Normal => 1000, + // WsCloseCode::Away => 1001, + WsCloseCode::Protocol => 1002, + // WsCloseCode::Unsupported => 1003, + // WsCloseCode::Abnormal => 1006, + // WsCloseCode::Invalid => 1007, + // WsCloseCode::Policy => 1008, + // WsCloseCode::Size => 1009, + // WsCloseCode::Extension => 1010, + // WsCloseCode::Error => 1011, + // WsCloseCode::Restart => 1012, + // WsCloseCode::Again => 1013, + // WsCloseCode::Tls => 1015, + } + } +} + +#[derive(Serialize, Deserialize, Debug)] +pub enum ProtocolMessage { + Version(u8), + StatusRequest(String), + StatusResponse(StatusCheck), + ActiveUserResponse(usize, usize), // user_count, connection count + ActiveProviderResponse(String, usize) +} + +impl ProtocolMessage { + pub fn to_bytes(&self) -> io::Result { + match self { + ProtocolMessage::Version(version) => { + Ok(Bytes::from(vec![*version])) + } + _ => { + let encoded = bincode_serialize(self)?; + Ok(Bytes::from(encoded)) + } + } + } + + pub fn from_bytes(bytes: Bytes) -> io::Result { + if bytes.len() == 1 { + Ok(ProtocolMessage::Version(bytes[0])) + } else { + bincode_deserialize::(bytes.as_ref()) + } + } +} \ No newline at end of file diff --git a/backend/src/utils/bincode_utils.rs b/shared/src/utils/bincode_utils.rs similarity index 94% rename from backend/src/utils/bincode_utils.rs rename to shared/src/utils/bincode_utils.rs index cb80df784..ae6a2b75c 100644 --- a/backend/src/utils/bincode_utils.rs +++ b/shared/src/utils/bincode_utils.rs @@ -1,5 +1,5 @@ use std::io; -use shared::error::to_io_error; +use crate::error::to_io_error; #[inline] pub fn bincode_serialize(value: &T) -> io::Result> diff --git a/shared/src/utils/mod.rs b/shared/src/utils/mod.rs index f1a70407e..debc43dc6 100644 --- a/shared/src/utils/mod.rs +++ b/shared/src/utils/mod.rs @@ -8,7 +8,7 @@ mod directed_graph; mod hash_utils; mod json_utils; mod serde_utils; - +mod bincode_utils; pub use self::default_utils::*; pub use self::time_utils::*; @@ -20,3 +20,4 @@ pub use self::directed_graph::*; pub use self::hash_utils::*; pub use self::json_utils::*; pub use self::serde_utils::*; +pub use self::bincode_utils::*; diff --git a/webui/Cargo.toml b/webui/Cargo.toml index e4ca307b9..b32e64c83 100644 --- a/webui/Cargo.toml +++ b/webui/Cargo.toml @@ -25,6 +25,8 @@ futures-signals = "0.3" futures = "0.3" prost = "0" wasm-bindgen-futures = "0" +bytes = "1" +bincode = { version = "2", features = ["std", "serde"] } [dependencies.web-sys] version = "0.3" diff --git a/webui/Trunk.toml b/webui/Trunk.toml index 1e33c8ee5..65dac1f4f 100644 --- a/webui/Trunk.toml +++ b/webui/Trunk.toml @@ -9,3 +9,7 @@ backend = "http://localhost:8901/api/v1/" [[proxy]] backend = "http://localhost:8901/auth/" + +[[proxy]] +backend = "ws://localhost:8901/ws" +ws = true \ No newline at end of file diff --git a/webui/src/app/components/authentication.rs b/webui/src/app/components/authentication.rs index f1ffd4c41..4324bf1ac 100644 --- a/webui/src/app/components/authentication.rs +++ b/webui/src/app/components/authentication.rs @@ -23,6 +23,9 @@ pub fn Authentication(props: &AuthenticationProps) -> Html { services_ctx.auth.auth_subscribe( &mut |success| { authenticated_state.set(success); + if success { + services_ctx.websocket.connect_ws(); + } future::ready(()) } ).await diff --git a/webui/src/app/components/dashboard/stats_view.rs b/webui/src/app/components/dashboard/stats_view.rs index a1414d8fe..f1bb436ec 100644 --- a/webui/src/app/components/dashboard/stats_view.rs +++ b/webui/src/app/components/dashboard/stats_view.rs @@ -8,9 +8,18 @@ pub fn StatsView() -> Html { let status_ctx = use_context::().expect("Status context not found"); let render_active_provider_connections = || -> Html { - match &status_ctx.status { - Some(stats) => { - if let Some(map) = &stats.active_provider_connections { + let empty_card = || html! { + + + + }; + match &status_ctx.status { + Some(stats) => { + if let Some(map) = &stats.active_provider_connections { + if map.len() > 0 { let cards = map.iter().map(|(provider, connections)| { html! { @@ -25,25 +34,14 @@ pub fn StatsView() -> Html { cards } else { - html! { - - - - } + empty_card() } + } else { + empty_card() } - None => html! { - - - - } } + None => empty_card() + } }; html! { diff --git a/webui/src/app/components/home.rs b/webui/src/app/components/home.rs index f2a1f78a2..40f708c09 100644 --- a/webui/src/app/components/home.rs +++ b/webui/src/app/components/home.rs @@ -1,7 +1,8 @@ +use std::cell::RefCell; +use std::collections::BTreeMap; use std::future; use std::rc::Rc; -use gloo_timers::callback::Interval; -use wasm_bindgen_futures::spawn_local; +use yew::platform::spawn_local; use yew::prelude::*; use yew::suspense::use_future; use shared::model::{AppConfigDto, StatusCheck}; @@ -9,6 +10,7 @@ use crate::app::components::{IconButton, Sidebar, DashboardView, PlaylistView, P use crate::app::context::{ConfigContext, StatusContext}; use crate::model::ViewType; use crate::hooks::use_service_context; +use crate::services::WsMessage; #[function_component] pub fn Home() -> Html { @@ -16,6 +18,7 @@ pub fn Home() -> Html { let app_title = services.config.ui_config.app_title.as_ref().map_or("tuliprox", |v| v.as_str()); let config = use_state(|| None::>); let status = use_state(|| None::>); + let status_holder = use_state(|| Rc::new(RefCell::new(None::>))); let view_visible = use_state(|| ViewType::Users); @@ -53,32 +56,65 @@ pub fn Home() -> Html { { let services_ctx = services.clone(); let status_signal = status.clone(); + let status_holder_signal = status_holder.clone(); use_effect_with((), move |_| { - let fetch_status = { - let status = status_signal.clone(); - let services_ctx = services_ctx.clone(); - move || { - let status = status.clone(); - let services_ctx = services_ctx.clone(); - spawn_local(async move { - status.set(services_ctx.status.get_server_status().await.ok()); - }); + let subid = services_ctx.websocket.subscribe(move |msg| { + match msg { + WsMessage::ServerStatus(server_status) => { + *status_holder_signal.borrow_mut() = Some(Rc::clone(&server_status)); + status_signal.set(Some(server_status)); + } + WsMessage::ActiveUser(user_count, connections) => { + let mut server_status = { + if let Some(old_status) = status_holder_signal.borrow().as_ref() { + (**old_status).clone() + } else { + StatusCheck::default() + } + }; + server_status.active_users = user_count; + server_status.active_user_connections = connections; + let new_status = Rc::new(server_status); + *status_holder_signal.borrow_mut() = Some(Rc::clone(&new_status)); + status_signal.set(Some(new_status)); + } + WsMessage::ActiveProvider(provider, connections) => { + let mut server_status = { + if let Some(old_status) = status_holder_signal.borrow().as_ref() { + (**old_status).clone() + } else { + StatusCheck::default() + } + }; + if let Some(treemap) = server_status.active_provider_connections.as_mut() { + if connections == 0 { + treemap.remove(&provider); + } else { + treemap.insert(provider, connections); + } + } else { + if connections > 0 { + let mut treemap = BTreeMap::new(); + treemap.insert(provider, connections); + server_status.active_provider_connections = Some(treemap); + } + } + let new_status = Rc::new(server_status); + *status_holder_signal.borrow_mut() = Some(Rc::clone(&new_status)); + status_signal.set(Some(new_status)); + } } - }; - - fetch_status(); - // all 5 seconds - let interval = Interval::new(5000, move || { - fetch_status(); }); - - // Cleanup function - || drop(interval) + let services_clone = services_ctx.clone(); + spawn_local(async move { + services_clone.websocket.get_server_status().await; + }); + let services_clone = services_ctx.clone(); + move || services_clone.websocket.unsubscribe(subid) }); } - let config_context = ConfigContext { config: (*config).clone(), }; diff --git a/webui/src/app/components/playlist/input_table.rs b/webui/src/app/components/playlist/input_table.rs index a2aaa28d5..c7b927a40 100644 --- a/webui/src/app/components/playlist/input_table.rs +++ b/webui/src/app/components/playlist/input_table.rs @@ -160,7 +160,7 @@ pub fn InputTable(props: &InputTableProps) -> Html { let popup_is_open_state = popup_is_open.clone(); let confirm = dialog.clone(); let translate = translate.clone(); - let selected_dto = selected_dto.clone(); + // let selected_dto = selected_dto.clone(); Callback::from(move |name:String| { if let Ok(action) = TableAction::from_str(&name) { match action { diff --git a/webui/src/app/components/userlist/userlist_view.rs b/webui/src/app/components/userlist/userlist_view.rs index 8e49d0fb7..2e1cebf51 100644 --- a/webui/src/app/components/userlist/userlist_view.rs +++ b/webui/src/app/components/userlist/userlist_view.rs @@ -66,9 +66,9 @@ pub fn UserlistView() -> Html { }); }; - let handle_create = { - Callback::from(move |cmd: String| {}) - }; + // let handle_create = { + // Callback::from(move |cmd: String| {}) + // }; html! { context={userlist_context}> diff --git a/webui/src/hooks/use_service_context.rs b/webui/src/hooks/use_service_context.rs index 2fc90fbb5..0ed950ec1 100644 --- a/webui/src/hooks/use_service_context.rs +++ b/webui/src/hooks/use_service_context.rs @@ -1,13 +1,14 @@ use std::rc::Rc; use yew::prelude::*; use crate::model::WebConfig; -use crate::services::{AuthService, ConfigService, PlaylistService, StatusService}; +use crate::services::{AuthService, ConfigService, PlaylistService, StatusService, WebSocketService}; pub struct Services { pub auth: Rc, pub config: Rc, pub status: Rc, pub playlist: Rc, + pub websocket: Rc, } impl Services { @@ -16,11 +17,13 @@ impl Services { let config = Rc::new(ConfigService::new(config)); let status = Rc::new(StatusService::new()); let playlist = Rc::new(PlaylistService::new()); + let websocket = Rc::new(WebSocketService::new(Rc::clone(&status))); Self { auth, config, status, - playlist + playlist, + websocket } } } diff --git a/webui/src/services/mod.rs b/webui/src/services/mod.rs index 16899c4c3..59786e1a6 100644 --- a/webui/src/services/mod.rs +++ b/webui/src/services/mod.rs @@ -4,6 +4,7 @@ mod requests; mod status_service; mod dialog_service; mod playlist_service; +mod websocket_service; pub use self::auth_service::*; pub use self::config_service::*; @@ -11,3 +12,4 @@ pub use self::requests::*; pub use self::status_service::*; pub use self::dialog_service::*; pub use self::playlist_service::*; +pub use self::websocket_service::*; diff --git a/webui/src/services/status_service.rs b/webui/src/services/status_service.rs index a8586f075..eaf6b342a 100644 --- a/webui/src/services/status_service.rs +++ b/webui/src/services/status_service.rs @@ -14,9 +14,7 @@ impl Default for StatusService { } impl StatusService { - pub fn new() -> Self { - Self {} - } + pub fn new() -> Self { Self {} } pub async fn get_server_status(&self) -> Result, crate::error::Error> { request_get::>(STATUS_PATH).await diff --git a/webui/src/services/websocket_service.rs b/webui/src/services/websocket_service.rs new file mode 100644 index 000000000..677074fe8 --- /dev/null +++ b/webui/src/services/websocket_service.rs @@ -0,0 +1,190 @@ +use wasm_bindgen::prelude::*; +use wasm_bindgen::JsCast; +use web_sys::{WebSocket, MessageEvent, Event, ErrorEvent, CloseEvent}; +use std::cell::RefCell; +use std::collections::{HashMap}; +use std::rc::Rc; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use web_sys::js_sys::{Uint8Array, ArrayBuffer}; +use log::{debug, error, trace}; +use shared::model::{ProtocolMessage, StatusCheck, PROTOCOL_VERSION}; +use crate::services::{get_token, StatusService}; + +#[derive(Clone)] +pub enum WsMessage { + ServerStatus(Rc), + ActiveUser(usize, usize), + ActiveProvider(String, usize), +} + +const WS_PATH: &str = "/ws"; + +pub struct WebSocketService { + connected: Rc, + ws: Rc>>, + status_service: Rc, + subscriber_id: Rc, + subscribers: Rc>>>, +} + +impl WebSocketService { + pub fn new(status_service: Rc) -> Self { + Self { + connected: Rc::new(AtomicBool::new(false)), + ws: Rc::new(RefCell::new(None)), + status_service, + subscriber_id: Rc::new(AtomicUsize::new(0)), + subscribers: Rc::new(RefCell::new(HashMap::new())), + } + } + + pub fn subscribe(&self, callback: F) -> usize { + let sub_id = self.subscriber_id.fetch_add(1, Ordering::SeqCst); + self.subscribers.borrow_mut().insert(sub_id, Box::new(callback)); + sub_id + } + + pub fn unsubscribe(&self, sub_id: usize) { + self.subscribers.borrow_mut().remove(&sub_id); + } + + pub fn broadcast(&self, msg: WsMessage) { + for (_, cb) in self.subscribers.borrow().iter() { + cb(msg.clone()); + } + } + + pub fn connect_ws(&self) { + if self.connected.load(Ordering::SeqCst) { + return; + } + match WebSocket::new(WS_PATH) { + Err(err) => error!("Failed to open websocket connection: {err:?}"), + Ok(socket) => { + socket.set_binary_type(web_sys::BinaryType::Arraybuffer); + let ws_clone = self.ws.clone(); + *ws_clone.borrow_mut() = Some(socket.clone()); + let subscribers_clone = self.subscribers.clone(); + let broadcast = move |msg: WsMessage| { + for (_, cb) in subscribers_clone.borrow().iter() { + cb(msg.clone()); + } + }; + + // onmessage + let onmessage_callback = Closure::::wrap(Box::new(move |event: MessageEvent| { + trace!("WebSocket received message: {:?}", event); + if let Ok(buf) = event.data().dyn_into::() { + let array = Uint8Array::new(&buf); + let bytes = bytes::Bytes::from(array.to_vec()); + match ProtocolMessage::from_bytes(bytes) { + Ok(message) => { + match message { + ProtocolMessage::ActiveUserResponse(user_count, connections) => { + broadcast(WsMessage::ActiveUser(user_count, connections)); + }, + ProtocolMessage::ActiveProviderResponse(user_count, connections) => { + broadcast(WsMessage::ActiveProvider(user_count, connections)); + }, + ProtocolMessage::StatusResponse(status) => { + let data = Rc::new(status); + broadcast(WsMessage::ServerStatus(data)); + } + ProtocolMessage::Version(_) + | ProtocolMessage::StatusRequest(_) => {} + } + } + Err(err) => error!("Failed to decode websocket message: {err}") + } + } + })); + socket.set_onmessage(Some(onmessage_callback.as_ref().unchecked_ref())); + onmessage_callback.forget(); // Important: leak the closure to keep it alive + + let ws_open_clone = Rc::clone(&ws_clone); + let connected_clone = self.connected.clone(); + // onopen + let onopen_callback = Closure::::wrap(Box::new(move |_event: Event| { + trace!("WebSocket connection opened."); + connected_clone.store(true, Ordering::SeqCst); + match ProtocolMessage::Version(PROTOCOL_VERSION).to_bytes() { + Ok(bytes) => Self::try_send_message(ws_open_clone.borrow().as_ref(), bytes), + Err(err) => error!("Failed to create WebSocket protocol version message: {err}"), + } + })); + socket.set_onopen(Some(onopen_callback.as_ref().unchecked_ref())); + onopen_callback.forget(); + + let ws_close_rc = self.ws.clone(); + let connected_clone = self.connected.clone(); + let onclose_callback = Closure::::wrap(Box::new(move |e: CloseEvent| { + debug!("Websocket closed, Code: {}, Reason: {}, WasClean: {}", e.code(), e.reason(), e.was_clean()); + *ws_close_rc.borrow_mut() = None; + connected_clone.store(false, Ordering::SeqCst); + })); + socket.set_onclose(Some(onclose_callback.as_ref().unchecked_ref())); + onclose_callback.forget(); + + let connected_clone = self.connected.clone(); + // onerror + let onerror_callback = Closure::::wrap(Box::new(move |e: ErrorEvent| { + error!("WebSocket error"); + connected_clone.store(false, Ordering::SeqCst); + web_sys::console::error_1(&e); + })); + socket.set_onerror(Some(onerror_callback.as_ref().unchecked_ref())); + onerror_callback.forget(); + } + } + } + + fn try_send_message(ws_opt: Option<&WebSocket>, data: bytes::Bytes) { + if let Some(ws) = ws_opt { + if let Err(err) = ws.send_with_u8_array(data.as_ref()) { + error!("Failed to send a websocket message: {err:?}"); + } + } + } + + pub fn send_message(&self, data: bytes::Bytes) { + Self::try_send_message(self.ws.borrow().as_ref(), data); + } + + pub async fn get_server_status(&self) { + if self.connected.load(Ordering::SeqCst) { + self.send_message(ProtocolMessage::StatusRequest(get_token().unwrap_or_default()).to_bytes().unwrap()); + } else { + match self.status_service.get_server_status().await { + Ok(status) => { + self.broadcast(WsMessage::ServerStatus(status)); + } + Err(err) => {error!("Failed to get server status: {err:?}");} + } + } + + + // TODO + // on no wesocket connection + + // let fetch_status = { + // let status = status_signal.clone(); + // let services_ctx = services_ctx.clone(); + // move || { + // let status = status.clone(); + // let services_ctx = services_ctx.clone(); + // spawn_local(async move { + // status.set(services_ctx.status.get_server_status().await.ok()); + // }); + // } + // }; + // + // fetch_status(); + // // all 5 seconds + // let interval = Interval::new(5000, move || { + // fetch_status(); + // }); + // + // // Cleanup function + // || drop(interval) + } +} \ No newline at end of file