diff --git a/CHANGELOG.md b/CHANGELOG.md index cf637d4c5..b4acfe020 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -93,6 +93,16 @@ happens in `main`, making the path easier to test and reuse. - **Admission Failure Deduplication**: Repeated admission failure response logic in `hls_api`, `m3u_api`, and `xtream_api` has been centralized into shared helpers. +- **API User Network Access Restrictions**: API proxy users can now be restricted by source CIDR ranges and/or GeoIP country + codes via `network_access`. + - Matching any configured CIDR or any configured country is sufficient for access. + - Network access checks are centralized in the API user request context so denied requests stop before endpoint handling + or upstream forwarding. + - Country-based checks require the GeoIP database. If GeoIP is unavailable, the secure default is to deny requests that + did not match a configured CIDR. + - Operators can explicitly opt into allowing this GeoIP-unavailable country-rule case with + `reverse_proxy.geoip.unavailable_policy: allow`. + - CIDR-only misses, unknown countries, and country mismatches still deny. - **QoS Aggregation Efficiency**: - QoS snapshot listing does not rely on a full unbounded materialization path for filtered UI/API reads. - Current-day QoS rebuilds are skipped when the history day is unchanged. @@ -214,6 +224,16 @@ - Added `qos_aggregation` (optional) with: - `enabled` (`bool`) - `interval_secs` (`u64`) +- **config.yml (`reverse_proxy.geoip`)**: + - Added `unavailable_policy` (`deny` | `allow`, default `deny`). + - `deny` keeps country-based `network_access` restrictions closed when GeoIP is disabled, missing, or not loaded. + - `allow` is an explicit risk acceptance that allows country-based `network_access` restrictions only when GeoIP is + unavailable. CIDR-only misses, unknown countries, and country mismatches still deny. +- **api-proxy.yml (`user.credentials[].network_access`)**: + - Added optional per-user network restrictions: + - `allowed_networks`: CIDR ranges such as `192.168.0.0/16` or `10.0.0.1/32`. + - `allowed_countries`: ISO-style country codes resolved through GeoIP. + - The rules use OR semantics: any matching CIDR or country allows the request. ## 3.3.0 (2026-04-02) diff --git a/Cargo.lock b/Cargo.lock index 4af74de62..e93c3d800 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4650,6 +4650,7 @@ dependencies = [ "hyper-util", "iana-time-zone", "indexmap", + "ipnet", "jsonwebtoken", "libc", "log", diff --git a/README.md b/README.md index a69669170..cea3f8087 100644 --- a/README.md +++ b/README.md @@ -162,6 +162,8 @@ Generate all four formats simultaneously from the same source — one setup, eve - **Download & Recording Manager**: Provider-aware VOD downloads and live recordings with retries, fairness, and RBAC-controlled actions - **Config Editor**: Direct editing of config.yml, source.yml, mapping.yml in the browser - **User Management**: API users with category selection, priority, soft-priority, normal/soft connection limits, auto-generated credentials +- **Network Access Policy UI**: API-user network restrictions can be configured with CIDR and GeoIP country rules, including + the global GeoIP-unavailable `deny`/`allow` policy. - **RBAC Admin Panel**: Tabbed user/group management, permission checkbox grid, write-without-read warnings - **Stream Table**: Real-time stream monitoring with copy-to-clipboard, bandwidth metrics, episode titles - **EPG View**: Timeline with channels, now-line, program details @@ -231,6 +233,12 @@ Generate all four formats simultaneously from the same source — one setup, eve - **JWT authentication**: Compact bitmask encoding with password-version tracking for automatic token invalidation - **Rate limiting**: Per-IP rate limiting with configurable burst and period - **Content Security Policy**: Configurable CSP headers +- **API User Network Access Restrictions**: Per-user `network_access` rules can restrict API proxy accounts by IPv4/IPv6 + CIDR ranges and/or GeoIP country codes. Matching any configured CIDR or country allows the request; otherwise the + request is denied before it is forwarded upstream. +- **GeoIP-Unavailable Policy**: Country-based network restrictions default to `deny` when GeoIP is disabled, missing, or not + loaded. Operators can explicitly accept that risk with `reverse_proxy.geoip.unavailable_policy: allow`; CIDR-only misses, + unknown countries, and country mismatches still deny. - **SSL/TLS support**: Configurable including `accept_insecure_ssl_certificates` option - **Proxy support**: HTTP, HTTPS, SOCKS5 proxies for all outgoing requests - **Header stripping**: Configurable removal of referer, Cloudflare, and X-headers diff --git a/backend/Cargo.toml b/backend/Cargo.toml index c41682b03..3a191787a 100644 --- a/backend/Cargo.toml +++ b/backend/Cargo.toml @@ -80,6 +80,7 @@ crc32fast = "1.5.0" memchr = "2.8.0" http-body-util = "0.1.3" lru = "0.12" +ipnet = "2" [target.'cfg(unix)'.dependencies] libc = "0.2" diff --git a/backend/src/api/api_utils.rs b/backend/src/api/api_utils.rs index 9a5b8e938..0bb120740 100644 --- a/backend/src/api/api_utils.rs +++ b/backend/src/api/api_utils.rs @@ -465,7 +465,7 @@ use crate::api::panel_api::{can_provision_on_exhausted, create_panel_api_provisi use crate::utils::LRUResourceCache; pub use internal_server_error; use shared::error::TuliproxError; -use shared::model::{AdmissionStrategy, ConnectFailureReason, FailureStage}; +use shared::model::{AdmissionStrategy, ConnectFailureReason, FailureStage, GeoIpUnavailablePolicy}; use shared::utils::{default_catchup_session_ttl_secs, default_hls_session_ttl_secs}; pub use try_option_bad_request; pub use try_option_forbidden; @@ -551,7 +551,7 @@ pub async fn serve_file(file_path: &Path, mime_type: String, cache_control: Opti pub fn get_user_target_by_username( username: &str, app_state: &Arc, -) -> Option<(ProxyUserCredentials, Arc)> { +) -> Option<(Arc, Arc)> { if !username.is_empty() { return app_state.app_config.get_target_for_username(username); } @@ -563,7 +563,7 @@ pub fn get_user_target_by_credentials<'a>( password: &str, api_req: &'a UserApiRequest, app_state: &'a AppState, -) -> Option<(ProxyUserCredentials, Arc)> { +) -> Option<(Arc, Arc)> { if !username.is_empty() && !password.is_empty() { app_state.app_config.get_target_for_user(username, password) } else { @@ -579,12 +579,124 @@ pub fn get_user_target_by_credentials<'a>( pub fn get_user_target<'a>( api_req: &'a UserApiRequest, app_state: &'a AppState, -) -> Option<(ProxyUserCredentials, Arc)> { +) -> Option<(Arc, Arc)> { let username = api_req.username.as_str().trim(); let password = api_req.password.as_str().trim(); get_user_target_by_credentials(username, password, api_req, app_state) } +/// Result of a policy-aware network access evaluation. +#[derive(Debug, Copy, Clone, Eq, PartialEq)] +pub enum NetworkAccessDecision { + /// Request is allowed (matched CIDR or country with `GeoIP` available). + Allowed, + /// Request is allowed because `GeoIP` is unavailable and policy is Allow. + AllowedGeoIpUnavailable, + /// Request is denied with a typed reason. + Denied(NetworkAccessDenyReason), +} + +#[derive(Debug, Copy, Clone, Eq, PartialEq)] +pub enum NetworkAccessDenyReason { + NoCidrMatch, + NoCountryMatch, + GeoIpUnavailable, + CountryUnknown, + MalformedClientIp, +} + +impl NetworkAccessDenyReason { + pub const fn as_str(self) -> &'static str { + match self { + Self::NoCidrMatch => "no_cidr_match", + Self::NoCountryMatch => "no_country_match", + Self::GeoIpUnavailable => "geoip_unavailable", + Self::CountryUnknown => "country_unknown", + Self::MalformedClientIp => "malformed_client_ip", + } + } +} + +/// Logs a network access denial with structured context for operator debugging. +/// Do NOT log passwords or secrets. +#[allow(clippy::uninlined_format_args)] +pub fn log_network_access_denied(username: &str, client_ip: &str, reason: &str) { + let sanitized_username = sanitize_sensitive_info(username); + let sanitized_client_ip = sanitize_sensitive_info(client_ip); + warn!( + target: "network_access", + "Network access denied: user=\"{}\" client_ip=\"{}\" reason={}", + sanitized_username, + sanitized_client_ip, + reason + ); +} + +/// Logs a network access allowed-without-GeoIP event for explicit-risk observability. +#[allow(clippy::uninlined_format_args)] +pub fn log_network_access_allowed_geoip_unavailable(username: &str, client_ip: &str) { + warn!( + target: "network_access", + "Network access allowed because GeoIP is unavailable and reverse_proxy.geoip.unavailable_policy=allow; user=\"{}\" client_ip=\"{}\"", + sanitize_sensitive_info(username), + sanitize_sensitive_info(client_ip) + ); +} + +/// Evaluates network access with the configured GeoIP-unavailable policy. +/// Returns a structured decision for logging and HTTP response mapping. +pub fn evaluate_network_access( + user: &ProxyUserCredentials, + client_ip: &str, + geoip: &Arc>, + geoip_unavailable_policy: GeoIpUnavailablePolicy, +) -> NetworkAccessDecision { + let Some(access) = user.network_access.as_ref() else { + return NetworkAccessDecision::Allowed; + }; + if access.is_empty() { + return NetworkAccessDecision::Allowed; + } + + let Ok(parsed_ip) = client_ip.parse::() else { + return NetworkAccessDecision::Denied(NetworkAccessDenyReason::MalformedClientIp); + }; + + // CIDR check + for net in &access.allowed_networks { + if net.contains(&parsed_ip) { + return NetworkAccessDecision::Allowed; + } + } + + // No CIDR match — check country rules + if access.allowed_countries.is_empty() { + return NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCidrMatch); + } + + // Country rules exist — check if GeoIP is available + let geoip_guard = geoip.load(); + let Some(geoip_db) = geoip_guard.as_ref() else { + // GeoIP unavailable — apply policy + return match geoip_unavailable_policy { + GeoIpUnavailablePolicy::Allow => NetworkAccessDecision::AllowedGeoIpUnavailable, + GeoIpUnavailablePolicy::Deny => NetworkAccessDecision::Denied(NetworkAccessDenyReason::GeoIpUnavailable), + }; + }; + + // GeoIP is loaded — do country lookup + match geoip_db.lookup(client_ip) { + Some(country) => { + if access.allowed_countries.iter().any(|c| c == &country) { + NetworkAccessDecision::Allowed + } else { + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCountryMatch) + } + } + None => NetworkAccessDecision::Denied(NetworkAccessDenyReason::CountryUnknown), + } +} + pub struct StreamOptions { pub stream_retry: bool, pub buffer_enabled: bool, @@ -3257,6 +3369,7 @@ pub fn create_api_proxy_user(app_state: &Arc) -> ProxyUserCredentials soft_connections: 0, soft_priority: 0, t_is_api_user: true, + network_access: None, } } @@ -3285,8 +3398,8 @@ mod tests { }, auth::Fingerprint, model::{ - AppConfig, Config, ConfigInput, ConfigInputAlias, ConfigTarget, MediaToolCapabilities, ProcessTargets, - ProxyUserCredentials, SourcesConfig, + AppConfig, Config, ConfigInput, ConfigInputAlias, ConfigTarget, + MediaToolCapabilities, NetworkAccess, ProcessTargets, ProxyUserCredentials, SourcesConfig, }, utils::{FileLockManager, GeoIp}, }; @@ -3298,7 +3411,7 @@ mod tests { foundation::Filter, model::{ ClusterFlags, ConfigPaths, ConfigTargetOptions, InputFetchMethod, InputType, PlaylistItemType, - ProcessingOrder, StreamChannel, XtreamCluster, + ProcessingOrder, ProxyType, StreamChannel, XtreamCluster, }, utils::{default_catchup_session_ttl_secs, default_hls_session_ttl_secs, Internable}, }; @@ -7290,4 +7403,345 @@ mod tests { assert_eq!(extension, ""); assert_eq!(query_path, "1"); } + + // ========================================================================================= + // evaluate_network_access tests + // ========================================================================================= + + /// Helper to create a test user with specific network access + fn user_with_network_access(network_access: Option) -> ProxyUserCredentials { + ProxyUserCredentials { + username: "test".to_string(), + password: "test".to_string(), + token: None, + proxy: ProxyType::default(), + server: None, + epg_timeshift: None, + epg_request_timeshift: None, + created_at: None, + exp_date: None, + max_connections: 0, + status: None, + output_clusters: ClusterFlags::all(), + ui_enabled: true, + comment: None, + priority: 0, + soft_connections: 0, + soft_priority: 0, + t_is_api_user: false, + network_access, + } + } + + #[test] + fn no_config_allows_all() { + let user = user_with_network_access(None); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!(evaluate_network_access(&user, "192.168.1.1", &geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + } + + #[test] + fn empty_config_allows_all() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec![], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!(evaluate_network_access(&user, "192.168.1.1", &geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + } + + #[test] + fn cidr_match_allows() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!(evaluate_network_access(&user, "192.168.1.42", &geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + } + + #[test] + fn cidr_miss_denies() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!( + evaluate_network_access(&user, "10.0.0.1", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCidrMatch) + ); + } + + #[test] + fn country_match_allows() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let mock_geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::test_new("DE"))))); + assert_eq!(evaluate_network_access(&user, "8.8.8.8", &mock_geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + } + + #[test] + fn country_miss_denies() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let mock_geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::test_new("US"))))); + assert_eq!( + evaluate_network_access(&user, "8.8.8.8", &mock_geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCountryMatch) + ); + } + + #[test] + fn no_geoip_denies_on_country_restriction() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!( + evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::GeoIpUnavailable) + ); + } + + #[test] + fn ipv4_vs_ipv6_denies_gracefully() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec!["2001:db8::/32".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!( + evaluate_network_access(&user, "192.168.1.1", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCidrMatch) + ); + } + + #[test] + fn ipv6_vs_ipv4_denies_gracefully() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!( + evaluate_network_access(&user, "2001:db8::1", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCidrMatch) + ); + } + + #[test] + fn either_cidr_or_country_match_allows() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["US".to_string()], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], + })); + let mock_geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::test_new("DE"))))); + assert_eq!(evaluate_network_access(&user, "192.168.1.42", &mock_geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + } + + #[test] + fn single_ip_cidr() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec!["192.168.1.1/32".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!(evaluate_network_access(&user, "192.168.1.1", &geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + assert_eq!( + evaluate_network_access(&user, "192.168.1.2", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCidrMatch) + ); + } + + // ========================================================================================= + // network denied reason tests + // ========================================================================================= + + #[test] + fn network_denied_reason_cidr_no_match() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!( + evaluate_network_access(&user, "10.0.0.1", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCidrMatch) + ); + } + + #[test] + fn network_denied_reason_country_no_match() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let mock_geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::test_new("US"))))); + assert_eq!( + evaluate_network_access(&user, "8.8.8.8", &mock_geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCountryMatch) + ); + } + + #[test] + fn network_denied_reason_geoip_unavailable() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!( + evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::GeoIpUnavailable) + ); + } + + #[test] + fn network_denied_reason_country_unknown_when_geoip_loaded_but_unknown_ip() { + // GeoIP is loaded (not None), but lookup returns None for this IP (private/unknown). + // We need a GeoIP that only covers a private range, so public IPs get None. + // Use a CIDR-only restriction (no country rules) so we can verify + // that when countries ARE checked, lookup None gives "country_unknown". + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], // miss CIDR first + })); + // Use the real GeoIp::new() which only seeds private ranges. + // For 8.8.8.8 (public), lookup returns None. + let geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::new())))); + // CIDR miss -> country check -> geoip loaded but lookup returns None + assert_eq!( + evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::CountryUnknown) + ); + } + + #[test] + fn network_denied_reason_none_when_allowed() { + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let mock_geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::test_new("DE"))))); + assert_eq!(evaluate_network_access(&user, "8.8.8.8", &mock_geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + } + + #[test] + fn network_denied_reason_none_when_no_config() { + let user = user_with_network_access(None); + let geoip = Arc::new(ArcSwapOption::::default()); + assert_eq!(evaluate_network_access(&user, "192.168.1.1", &geoip, GeoIpUnavailablePolicy::Deny), NetworkAccessDecision::Allowed); + } + + // ========================================================================================= + // GeoIP unavailable policy tests + // ========================================================================================= + + #[test] + fn geoip_unavailable_default_deny_denies() { + // Country rule exists but GeoIP is unavailable — default policy is Deny + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + let decision = evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Deny); + assert_eq!(decision, NetworkAccessDecision::Denied(NetworkAccessDenyReason::GeoIpUnavailable)); + } + + #[test] + fn geoip_unavailable_explicit_allow_allows() { + // Country rule exists, GeoIP unavailable, but policy is Allow — allows + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + let decision = evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Allow); + assert_eq!(decision, NetworkAccessDecision::AllowedGeoIpUnavailable); + } + + #[test] + fn geoip_unavailable_cidr_only_still_denies() { + // CIDR only rules, no match — should deny even with Allow policy + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec![], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + let decision = evaluate_network_access(&user, "10.0.0.1", &geoip, GeoIpUnavailablePolicy::Allow); + assert_eq!(decision, NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCidrMatch)); + } + + #[test] + fn geoip_unavailable_cidr_match_allows() { + // CIDR match always allows, regardless of policy + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec!["10.0.0.0/8".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + let decision = evaluate_network_access(&user, "10.0.0.1", &geoip, GeoIpUnavailablePolicy::Deny); + assert_eq!(decision, NetworkAccessDecision::Allowed); + } + + #[test] + fn geoip_loaded_country_mismatch_still_denies() { + // Loaded GeoIP but country doesn't match — should deny under Allow policy + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let mock_geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::test_new("US"))))); + let decision = evaluate_network_access(&user, "8.8.8.8", &mock_geoip, GeoIpUnavailablePolicy::Allow); + assert_eq!(decision, NetworkAccessDecision::Denied(NetworkAccessDenyReason::NoCountryMatch)); + } + + #[test] + fn geoip_loaded_unknown_country_still_denies() { + // Loaded GeoIP but lookup returns None — should deny under Allow policy + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec!["192.168.1.0/24".parse().unwrap()], + })); + let geoip = Arc::new(ArcSwapOption::from(Some(Arc::new(GeoIp::new())))); + let decision = evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Allow); + assert_eq!(decision, NetworkAccessDecision::Denied(NetworkAccessDenyReason::CountryUnknown)); + } + + #[test] + fn malformed_ip_denies_even_when_geoip_unavailable_policy_is_allow() { + let user = user_with_network_access(Some(NetworkAccess::from(&shared::model::NetworkAccessDto { + allowed_countries: Some(vec!["DE".to_string()]), + allowed_networks: None, + }))); + let geoip = Arc::new(ArcSwapOption::::default()); + + let decision = evaluate_network_access(&user, "not-an-ip", &geoip, GeoIpUnavailablePolicy::Allow); + + assert_eq!(decision, NetworkAccessDecision::Denied(NetworkAccessDenyReason::MalformedClientIp)); + } + + #[test] + fn evaluate_network_access_respects_allow_policy() { + // verify evaluate_network_access returns AllowedGeoIpUnavailable with Allow policy + let user = user_with_network_access(Some(NetworkAccess { + allowed_countries: vec!["DE".to_string()], + allowed_networks: vec![], + })); + let geoip = Arc::new(ArcSwapOption::::default()); + // Default deny policy should return Denied + assert_eq!( + evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Deny), + NetworkAccessDecision::Denied(NetworkAccessDenyReason::GeoIpUnavailable) + ); + // Allow policy should return AllowedGeoIpUnavailable + assert_eq!(evaluate_network_access(&user, "8.8.8.8", &geoip, GeoIpUnavailablePolicy::Allow), NetworkAccessDecision::AllowedGeoIpUnavailable); + } } diff --git a/backend/src/api/endpoints/custom_video_stream_api.rs b/backend/src/api/endpoints/custom_video_stream_api.rs index 2cb5df10f..e8515cce7 100644 --- a/backend/src/api/endpoints/custom_video_stream_api.rs +++ b/backend/src/api/endpoints/custom_video_stream_api.rs @@ -4,6 +4,7 @@ use crate::{ }; use axum::response::IntoResponse; use std::{str::FromStr, sync::Arc}; +use crate::auth::resolve_api_user_context; async fn cvs_api( fingerprint: Fingerprint, @@ -16,12 +17,12 @@ async fn cvs_api( return axum::http::StatusCode::NOT_FOUND.into_response(); }; - let Some((user, _target)) = app_state.app_config.get_target_for_user(&username, &password) else { + let Some((user, target)) = app_state.app_config.get_target_for_user(&username, &password) else { return app_state.app_config.get_auth_error_status().into_response(); }; - if user.permission_denied(&app_state) { - return axum::http::StatusCode::FORBIDDEN.into_response(); + if let Err(e) = resolve_api_user_context(user.clone(), target.clone(), fingerprint.clone(), &app_state) { + return e.into_player_response(app_state.app_config.get_auth_error_status()); } create_custom_video_stream_response(&app_state, &fingerprint.addr, custom_video_type).into_response() diff --git a/backend/src/api/endpoints/hdhomerun_api.rs b/backend/src/api/endpoints/hdhomerun_api.rs index 23499f1cf..5177b2d96 100644 --- a/backend/src/api/endpoints/hdhomerun_api.rs +++ b/backend/src/api/endpoints/hdhomerun_api.rs @@ -3,7 +3,7 @@ use crate::{ api_utils::{internal_server_error, try_unwrap_body}, model::HdHomerunAppState, }, - auth::AuthBasic, + auth::{try_check_network_access_only, AuthBasic, Fingerprint}, model::{AppConfig, ConfigInputFlags, ConfigTarget, ProxyUserCredentials}, processing::parser::xtream::get_xtream_url, repository::{iter_raw_m3u_target_playlist, M3uPlaylistIterator, XtreamPlaylistIterator}, @@ -243,13 +243,19 @@ async fn discover_json( } async fn lineup_status( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse { + let cfg = Arc::clone(&app_state.app_state.app_config); + let network_access_denied = + cfg.get_target_for_username(&app_state.device.t_username).is_none_or(|(credentials, _)| { + try_check_network_access_only(&credentials, &fingerprint, &app_state.app_state).is_err() + }); let current_state = app_state.hd_scan_state.load(std::sync::atomic::Ordering::Acquire); - if current_state < 0 { + if network_access_denied || current_state < 0 { axum::Json(json!({ "ScanInProgress": 0, - "ScanPossible": 1, + "ScanPossible": i32::from(!network_access_denied), "Source": "Cable", "SourceList": ["Cable"], })) @@ -259,47 +265,48 @@ async fn lineup_status( let final_state = if new_state > 100 { 100 } else { new_state }; let cfg = Arc::clone(&app_state.app_state.app_config); - let num_of_channels = if let Some((user, target)) = cfg.get_target_for_username(&app_state.device.t_username) { - if target.has_output(TargetType::M3u) { - let credentials = Arc::new(user); - if let Some(iter) = iter_raw_m3u_target_playlist(&cfg, &target, None).await { - iter.filter_map(move |res| { - let credentials = Arc::clone(&credentials); - async move { - let item = res.ok()?; - credentials.allows_item_type(item.item_type).then_some(item) + let num_of_channels = + if let Some((credentials, target)) = cfg.get_target_for_username(&app_state.device.t_username) { + if target.has_output(TargetType::M3u) { + if let Some(iter) = iter_raw_m3u_target_playlist(&cfg, &target, None).await { + let cred = Arc::clone(&credentials); + iter.filter_map(move |res| { + let cred = Arc::clone(&cred); + async move { + let item = res.ok()?; + cred.allows_item_type(item.item_type).then_some(item) + } + }) + .count() + .await + } else { + 0 + } + } else if target.has_output(TargetType::Xtream) { + let cred = Arc::clone(&credentials); + let live = if cred.allows_cluster(XtreamCluster::Live) { + match XtreamPlaylistIterator::new(XtreamCluster::Live, &cfg, &target, None, &cred).await { + Ok(stream) => stream.count().await, + Err(_) => 0, } - }) - .count() - .await + } else { + 0 + }; + let vod = if cred.allows_cluster(XtreamCluster::Video) { + match XtreamPlaylistIterator::new(XtreamCluster::Video, &cfg, &target, None, &cred).await { + Ok(stream) => stream.count().await, + Err(_) => 0, + } + } else { + 0 + }; + live + vod } else { 0 } - } else if target.has_output(TargetType::Xtream) { - let credentials = Arc::new(user); - let live = if credentials.allows_cluster(XtreamCluster::Live) { - match XtreamPlaylistIterator::new(XtreamCluster::Live, &cfg, &target, None, &credentials).await { - Ok(stream) => stream.count().await, - Err(_) => 0, - } - } else { - 0 - }; - let vod = if credentials.allows_cluster(XtreamCluster::Video) { - match XtreamPlaylistIterator::new(XtreamCluster::Video, &cfg, &target, None, &credentials).await { - Ok(stream) => stream.count().await, - Err(_) => 0, - } - } else { - 0 - }; - live + vod } else { 0 - } - } else { - 0 - }; + }; if final_state >= 100 { app_state.hd_scan_state.store(-1, std::sync::atomic::Ordering::Release); @@ -322,9 +329,17 @@ struct LineupPostQuery { } async fn lineup_post( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, axum::extract::Query(query): axum::extract::Query, ) -> impl IntoResponse { + let cfg = Arc::clone(&app_state.app_state.app_config); + let allowed = cfg.get_target_for_username(&app_state.device.t_username).is_some_and(|(credentials, _)| { + try_check_network_access_only(&credentials, &fingerprint, &app_state.app_state).is_ok() + }); + if !allowed { + return axum::http::StatusCode::FORBIDDEN.into_response(); + } match query.scan.as_str() { "start" => { app_state.hd_scan_state.store(0, std::sync::atomic::Ordering::Release); @@ -429,6 +444,7 @@ async fn lineup( } async fn auth_lineup_json( + fingerprint: Fingerprint, AuthBasic((username, password)): AuthBasic, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse { @@ -437,18 +453,25 @@ async fn auth_lineup_json( if !username.eq(&credentials.username) || !password.eq(&credentials.password) { return axum::http::StatusCode::UNAUTHORIZED.into_response(); } - let user_credentials = Arc::new(credentials); + if let Err(e) = try_check_network_access_only(&credentials, &fingerprint, &app_state.app_state) { + return e.into_player_response(app_state.app_state.app_config.get_auth_error_status()); + } + let user_credentials = Arc::clone(&credentials); return lineup(&app_state, &cfg, &user_credentials, &target).await.into_response(); } axum::http::StatusCode::NOT_FOUND.into_response() } async fn lineup_json( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse { let cfg = Arc::clone(&app_state.app_state.app_config); if let Some((credentials, target)) = cfg.get_target_for_username(&app_state.device.t_username) { - let user_credentials = Arc::new(credentials); + if let Err(e) = try_check_network_access_only(&credentials, &fingerprint, &app_state.app_state) { + return e.into_player_response(app_state.app_state.app_config.get_auth_error_status()); + } + let user_credentials = Arc::clone(&credentials); return lineup(&app_state, &cfg, &user_credentials, &target).await.into_response(); } axum::http::StatusCode::NOT_FOUND.into_response() diff --git a/backend/src/api/endpoints/hls_api.rs b/backend/src/api/endpoints/hls_api.rs index aaaccf34f..47e12784f 100644 --- a/backend/src/api/endpoints/hls_api.rs +++ b/backend/src/api/endpoints/hls_api.rs @@ -1,10 +1,11 @@ use crate::{ api::{ api_utils::{ - connection_priority_for_kind, create_session_fingerprint, force_provider_stream_response, get_headers_from_request, - get_hls_session_ttl_secs, - admission_failure_response, get_stream_alternative_url, is_seek_request, local_stream_response, try_option_bad_request, - try_unwrap_body, HeaderFilter, + admission_failure_response, connection_priority_for_kind, + create_session_fingerprint, force_provider_stream_response, get_headers_from_request, + get_hls_session_ttl_secs, get_stream_alternative_url, is_seek_request, local_stream_response, + try_option_bad_request, try_unwrap_body, + HeaderFilter, }, model::{ AppState, CustomVideoStreamType, ProviderAllocation, UserSession, @@ -26,6 +27,7 @@ use shared::{ use std::sync::Arc; use url::Url; use shared::model::ConnectFailureReason; +use crate::auth::check_network_access_only; const PLAYLIST_TEMPLATE: &str = r"#EXTM3U #EXT-X-VERSION:3 @@ -377,11 +379,13 @@ async fn hls_api_stream( axum::extract::Path(params): axum::extract::Path, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse + Send { - let (user, target) = try_option_bad_request!( - app_state.app_config.get_target_for_user(¶ms.username, ¶ms.password), - false, - format!("Could not find any user for hls stream {}", params.username) - ); + let Some((user, target)) = app_state.app_config.get_target_for_user(¶ms.username, ¶ms.password) else { + return axum::http::StatusCode::BAD_REQUEST.into_response(); + }; + // Network access check only - permission check is done later with full stream info + if let Err(e) = check_network_access_only(&user, &fingerprint, &app_state) { + return e.into_player_response(app_state.app_config.get_auth_error_status()); + } let target_name = &target.name; let virtual_id = params.stream_id; let input = try_option_bad_request!( @@ -391,19 +395,12 @@ async fn hls_api_stream( ); if user.permission_denied(&app_state) { - let denied_channel = resolve_stream_channel( - &app_state, - &target, - &input, - virtual_id, - &Arc::from(String::new()), - ) - .await; + let stream_channel = resolve_stream_channel(&app_state, &target, &input, virtual_id, "").await; return admission_failure_response( &app_state, &fingerprint, &user, - denied_channel, + stream_channel, input.name.clone(), &req_headers, ConnectFailureReason::UserAccountExpired, diff --git a/backend/src/api/endpoints/m3u_api.rs b/backend/src/api/endpoints/m3u_api.rs index f031464f7..cf47c91d6 100644 --- a/backend/src/api/endpoints/m3u_api.rs +++ b/backend/src/api/endpoints/m3u_api.rs @@ -1,12 +1,14 @@ use crate::{ api::{ api_utils::{ - admission_failure_response, create_catchup_session_key, create_session_fingerprint, - force_provider_stream_response, get_session_reservation_ttl_secs, get_user_target, - get_user_target_by_credentials, is_seek_request, is_session_based_playback, is_stream_share_enabled, - local_stream_response, redirect, redirect_response, resource_response, separate_number_and_remainder, - should_allow_exhausted_shared_reconnect, stream_response, try_option_bad_request, try_option_forbidden, - try_result_bad_request, try_result_not_found, try_unwrap_body, RedirectParams, + admission_failure_response, create_catchup_session_key, + create_session_fingerprint, force_provider_stream_response, get_session_reservation_ttl_secs, + get_user_target, get_user_target_by_credentials, is_seek_request, is_session_based_playback, + is_stream_share_enabled, local_stream_response, + redirect, redirect_response, resource_response, + separate_number_and_remainder, should_allow_exhausted_shared_reconnect, stream_response, + try_option_bad_request, try_result_bad_request, try_result_not_found, + try_unwrap_body, RedirectParams, }, endpoints::{ hls_api::handle_hls_stream_request, @@ -15,6 +17,7 @@ use crate::{ model::{AppState, UserApiRequest, UserApiRequestQueryOrBody}, }, auth::Fingerprint, + model::{ConfigTarget, ProxyUserCredentials}, repository::{m3u_get_item_for_stream_id, m3u_load_rewrite_playlist, storage_const}, utils::debug_if_enabled, }; @@ -28,20 +31,18 @@ use shared::{ utils::{concat_path, extract_extension_from_url, sanitize_sensitive_info}, }; use std::sync::Arc; +use crate::auth::{check_network_access_only, resolve_api_user_context, ApiUserAuthError}; -async fn m3u_api(api_req: &UserApiRequest, app_state: &AppState) -> impl IntoResponse + Send { - api_req.log_sanitized("m3u_api"); - let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target(api_req, app_state), - auth_status, - false, - format!("Could not find any user for m3u api {}", api_req.username) - ); +async fn m3u_api( + user: Arc, + target: Arc, + app_state: &AppState, + content_type: &str, +) -> impl IntoResponse + Send { + let _guard = app_state.app_config.file_locks.write_lock_str(&user.username).await; match m3u_load_rewrite_playlist(&app_state.app_config, &target, &user).await { Ok(m3u_iter) => { - // Convert the stream into a stream of `Bytes` let content_stream = m3u_iter.map(|mut line| { line.push('\n'); Ok::(Bytes::from(line)) @@ -50,8 +51,8 @@ async fn m3u_api(api_req: &UserApiRequest, app_state: &AppState) -> impl IntoRes let mut builder = axum::response::Response::builder() .status(axum::http::StatusCode::OK) .header(axum::http::header::CONTENT_TYPE, mime::TEXT_PLAIN_UTF_8.to_string()); - if api_req.content_type == "m3u_plus" { - builder = builder.header("Content-Disposition", "attachment; filename=\"playlist.m3u\""); + if content_type == "m3u_plus" { + builder = builder.header(axum::http::header::CONTENT_DISPOSITION, "attachment; filename=\"playlist.m3u\""); } try_unwrap_body!(builder.body(axum::body::Body::from_stream(content_stream))) } @@ -62,37 +63,66 @@ async fn m3u_api(api_req: &UserApiRequest, app_state: &AppState) -> impl IntoRes } } +fn m3u_api_with_auth( + fingerprint: &Fingerprint, + app_state: &Arc, + api_req: &UserApiRequest, +) -> Result<(Arc, Arc), ApiUserAuthError> { + let (user, target) = get_user_target(api_req, app_state) + .ok_or(ApiUserAuthError::AuthFailed)?; + check_network_access_only(&user, fingerprint, app_state)?; + Ok((user, target)) +} + +/// Network-only auth for stream endpoints. Permission check is done later by the stream +/// handler with full stream info for `admission_failure_response`. +fn m3u_api_stream_network_auth( + fingerprint: &Fingerprint, + app_state: &Arc, + api_req: &UserApiRequest, + stream_req: &ApiStreamRequest<'_>, +) -> Result<(Arc, Arc), ApiUserAuthError> { + let (user, target) = get_user_target_by_credentials(stream_req.username, stream_req.password, api_req, app_state) + .ok_or(ApiUserAuthError::AuthFailed)?; + check_network_access_only(&user, fingerprint, app_state)?; + Ok((user, target)) +} + async fn m3u_api_get( + fingerprint: Fingerprint, axum::extract::Query(api_req): axum::extract::Query, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse + Send { - m3u_api(&api_req, &app_state).await + let auth_status = app_state.app_config.get_auth_error_status(); + let (user, target) = match m3u_api_with_auth(&fingerprint, &app_state, &api_req) { + Ok(ctx) => ctx, + Err(e) => return e.into_player_response(auth_status), + }; + m3u_api(user, target, &app_state, &api_req.content_type).await.into_response() } async fn m3u_api_post( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, UserApiRequestQueryOrBody(api_req): UserApiRequestQueryOrBody, ) -> impl IntoResponse + Send { - m3u_api(&api_req, &app_state).await.into_response() + let auth_status = app_state.app_config.get_auth_error_status(); + let (user, target) = match m3u_api_with_auth(&fingerprint, &app_state, &api_req) { + Ok(ctx) => ctx, + Err(e) => return e.into_player_response(auth_status), + }; + m3u_api(user, target, &app_state, &api_req.content_type).await.into_response() } #[allow(clippy::too_many_lines)] async fn m3u_api_stream( + user: Arc, + target: Arc, fingerprint: &Fingerprint, req_headers: &axum::http::HeaderMap, app_state: &Arc, - api_req: &UserApiRequest, stream_req: ApiStreamRequest<'_>, - // _addr: &std::net::SocketAddr, ) -> impl IntoResponse + Send { - let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target_by_credentials(stream_req.username, stream_req.password, api_req, app_state), - auth_status, - false, - format!("Could not find any user for m3u stream {}", stream_req.username) - ); - let _guard = app_state.app_config.file_locks.write_lock_str(&user.username).await; let target_name = &target.name; @@ -347,7 +377,21 @@ async fn m3u_api_stream( .into_response() } +fn m3u_api_resource_auth( + fingerprint: &Fingerprint, + app_state: &Arc, + api_req: &UserApiRequest, + username: &str, + password: &str, +) -> Result<(Arc, Arc), ApiUserAuthError> { + let (user, target) = get_user_target_by_credentials(username, password, api_req, app_state) + .ok_or(ApiUserAuthError::AuthFailed)?; + resolve_api_user_context(user.clone(), target.clone(), fingerprint.clone(), app_state)?; + Ok((user, target)) +} + async fn m3u_api_resource( + fingerprint: Fingerprint, req_headers: axum::http::HeaderMap, axum::extract::Query(api_req): axum::extract::Query, axum::extract::Path((username, password, stream_id, resource)): axum::extract::Path<( @@ -362,15 +406,10 @@ async fn m3u_api_resource( return axum::http::StatusCode::BAD_REQUEST.into_response(); }; let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target_by_credentials(&username, &password, &api_req, &app_state), - auth_status, - false, - format!("Could not find any user for m3u resource {username}") - ); - if user.permission_denied(&app_state) { - return axum::http::StatusCode::FORBIDDEN.into_response(); - } + let (user, target) = match m3u_api_resource_auth(&fingerprint, &app_state, &api_req, &username, &password) { + Ok(ctx) => ctx, + Err(e) => return e.into_player_response(auth_status), + }; let target_name = &target.name; if !target.has_output(TargetType::M3u) { @@ -411,17 +450,24 @@ macro_rules! create_m3u_api_stream { axum::extract::Query(api_req): axum::extract::Query, axum::extract::Path((username, password, stream_id)): axum::extract::Path<(String, String, String)>, axum::extract::State(app_state): axum::extract::State>, - // axum::extract::ConnectInfo(addr): axum::extract::ConnectInfo, ) -> impl IntoResponse + Send { - m3u_api_stream( - &fingerprint, - &req_headers, - &app_state, - &api_req, - ApiStreamRequest::from($context, &username, &password, &stream_id, ""), - ) - .await - .into_response() + let stream_req = ApiStreamRequest::from($context, &username, &password, &stream_id, ""); + let auth_status = app_state.app_config.get_auth_error_status(); + match m3u_api_stream_network_auth(&fingerprint, &app_state, &api_req, &stream_req) { + Ok((user, target)) => { + m3u_api_stream( + user, + target, + &fingerprint, + &req_headers, + &app_state, + stream_req, + ) + .await + .into_response() + } + Err(e) => e.into_player_response(auth_status), + } } }; } diff --git a/backend/src/api/endpoints/user_api.rs b/backend/src/api/endpoints/user_api.rs index 4541bc48e..d174a9f89 100644 --- a/backend/src/api/endpoints/user_api.rs +++ b/backend/src/api/endpoints/user_api.rs @@ -1,6 +1,8 @@ use crate::{ api::{ - api_utils::{get_user_target_by_username, get_username_from_auth_header, try_unwrap_body}, + api_utils::{ + get_user_target_by_username, get_username_from_auth_header, try_unwrap_body, + }, model::AppState, }, auth::{validator_api_user, AuthBearer}, diff --git a/backend/src/api/endpoints/v1_api_user.rs b/backend/src/api/endpoints/v1_api_user.rs index d6dd457b5..535305133 100644 --- a/backend/src/api/endpoints/v1_api_user.rs +++ b/backend/src/api/endpoints/v1_api_user.rs @@ -23,7 +23,9 @@ async fn save_config_api_proxy_user( }; let _lock = app_state.app_config.file_locks.write_lock(Path::new(&api_proxy_file_path)).await; - credential.prepare(); + if let Err(err) = credential.prepare() { + return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": err.to_string()}))).into_response(); + } if let Err(err) = credential.validate() { return (axum::http::StatusCode::BAD_REQUEST, axum::Json(json!({"error": err.to_string()}))).into_response(); } @@ -97,11 +99,11 @@ async fn save_config_api_proxy_user( if user_target_idx == target_idx { // Update - api_proxy.user[user_target_idx].credentials[user_idx] = ProxyUserCredentials::from(&credential); + api_proxy.user[user_target_idx].credentials[user_idx] = Arc::new(ProxyUserCredentials::from(&credential)); } else { // Move: remove from old target and insert into new target api_proxy.user[user_target_idx].credentials.remove(user_idx); - api_proxy.user[target_idx].credentials.push(ProxyUserCredentials::from(&credential)); + api_proxy.user[target_idx].credentials.push(Arc::new(ProxyUserCredentials::from(&credential))); remove_empty_target = api_proxy.user[user_target_idx].credentials.is_empty(); } @@ -110,7 +112,7 @@ async fn save_config_api_proxy_user( } } else { // new user - api_proxy.user[target_idx].credentials.push(ProxyUserCredentials::from(&credential)); + api_proxy.user[target_idx].credentials.push(Arc::new(ProxyUserCredentials::from(&credential))); } let new_api_proxy = Arc::new(api_proxy); diff --git a/backend/src/api/endpoints/xmltv_api.rs b/backend/src/api/endpoints/xmltv_api.rs index 8709e43af..abe7fb6e8 100644 --- a/backend/src/api/endpoints/xmltv_api.rs +++ b/backend/src/api/endpoints/xmltv_api.rs @@ -1,16 +1,18 @@ use crate::{ api::{ api_utils::{ - create_api_proxy_user, empty_json_response_as_array, get_user_target, get_user_target_by_credentials, - internal_server_error, resource_response, stream_json_or_bin_response_stream, - try_option_forbidden, try_unwrap_body, + create_api_proxy_user, empty_json_response_as_array, get_user_target, + get_user_target_by_credentials, internal_server_error, + resource_response, + stream_json_or_bin_response_stream, try_unwrap_body, }, - model::{AppState, UserApiRequestQueryOrBody, UserApiRequest}, + model::{AppState, UserApiRequest, UserApiRequestQueryOrBody}, }, + auth::Fingerprint, model::{Config, ConfigTarget, ProxyUserCredentials, TargetOutput, EPG_ATTRIB_ID, EPG_TAG_CHANNEL}, repository::{ - get_target_storage_path, m3u_get_epg_file_path_for_target, storage_const, xtream_get_epg_file_path_for_target, - xtream_get_storage_path, BPlusTreeQuery, epg_query_channels, LockedReceiverStream, XML_PREAMBLE, + epg_query_channels, get_target_storage_path, m3u_get_epg_file_path_for_target, storage_const, + xtream_get_epg_file_path_for_target, xtream_get_storage_path, BPlusTreeQuery, LockedReceiverStream, XML_PREAMBLE, }, utils, utils::{ @@ -20,7 +22,7 @@ use crate::{ }; use axum::response::IntoResponse; use chrono::{DateTime, TimeZone}; -use log::{debug, error, trace}; +use log::{error, trace}; use quick_xml::events::{BytesEnd, BytesStart, BytesText, Event}; use shared::{ concat_string, @@ -38,6 +40,7 @@ use std::{ use tokio::{io::AsyncWriteExt, sync::mpsc, task}; use tokio_stream::StreamExt; use tokio_util::io::ReaderStream; +use crate::auth::resolve_api_user_context; use crate::model::ApiProxyServerInfo; pub fn get_empty_epg_response() -> axum::response::Response { @@ -667,24 +670,18 @@ pub(crate) async fn stream_epg_api( /// let router = xmltv_api_register(); /// // A GET request to /xmltv.php with valid query parameters will invoke this handler. /// ``` -async fn xmltv_api(api_req: UserApiRequest, app_state: &Arc) -> impl IntoResponse + Send { +async fn xmltv_api(fingerprint: &Fingerprint, api_req: UserApiRequest, app_state: &Arc) -> impl IntoResponse + Send { api_req.log_sanitized("xmltv_api"); let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target(&api_req, app_state), - auth_status, - false, - format!("Could not find any user for xmltv api {}", api_req.username) - ); - - if user.permission_denied(app_state) { - return axum::http::StatusCode::FORBIDDEN.into_response(); + let Some((user, target)) = get_user_target(&api_req, app_state) else { + return auth_status.into_response(); + }; + if let Err(e) = resolve_api_user_context(user.clone(), target.clone(), fingerprint.clone(), app_state) { + return e.into_player_response(auth_status); } let config = &app_state.app_config.config.load(); let Some(epg_path) = get_epg_path_for_target(config, &target) else { - // No epg configured, No processing or timeshift, epg can't be mapped to the channels. - // we do not deliver epg return get_empty_epg_response(); }; @@ -692,34 +689,34 @@ async fn xmltv_api(api_req: UserApiRequest, app_state: &Arc) -> impl I } async fn xmltv_api_get( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, axum::extract::Query(api_req): axum::extract::Query, ) -> impl IntoResponse + Send { - xmltv_api(api_req, &app_state).await + xmltv_api(&fingerprint, api_req, &app_state).await } async fn xmltv_api_post( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, UserApiRequestQueryOrBody(api_req): UserApiRequestQueryOrBody, ) -> impl IntoResponse + Send { - xmltv_api(api_req, &app_state).await + xmltv_api(&fingerprint, api_req, &app_state).await } async fn epg_api_resource( + fingerprint: Fingerprint, req_headers: axum::http::HeaderMap, axum::extract::Query(api_req): axum::extract::Query, axum::extract::Path((username, password, resource)): axum::extract::Path<(String, String, String)>, axum::extract::State(app_state): axum::extract::State>, ) -> impl IntoResponse + Send { let auth_status = app_state.app_config.get_auth_error_status(); - let (user, _target) = try_option_forbidden!( - get_user_target_by_credentials(&username, &password, &api_req, &app_state), - auth_status, - false, - format!("Could not find any user for epg resource {username}") - ); - if user.permission_denied(&app_state) { - return axum::http::StatusCode::FORBIDDEN.into_response(); + let Some((user, target)) = get_user_target_by_credentials(&username, &password, &api_req, &app_state) else { + return auth_status.into_response(); + }; + if let Err(e) = resolve_api_user_context(user.clone(), target.clone(), fingerprint.clone(), &app_state) { + return e.into_player_response(auth_status); } let encrypt_secret = app_state.get_encrypt_secret(); diff --git a/backend/src/api/endpoints/xtream_api.rs b/backend/src/api/endpoints/xtream_api.rs index b6bb5c548..a29b4d3f4 100644 --- a/backend/src/api/endpoints/xtream_api.rs +++ b/backend/src/api/endpoints/xtream_api.rs @@ -4,13 +4,15 @@ use crate::{ api::{ api_utils, api_utils::{ - admission_failure_response, create_api_proxy_user, create_catchup_session_key, create_session_fingerprint, - empty_json_response_as_array, empty_json_response_as_object, force_provider_stream_response, - get_session_reservation_ttl_secs, get_user_target, get_user_target_by_credentials, internal_server_error, - is_seek_request, is_session_based_playback, is_stream_share_enabled, local_stream_response, redirect, - redirect_response, resource_response, separate_number_and_remainder, - should_allow_exhausted_shared_reconnect, stream_response, try_option_bad_request, try_option_forbidden, - try_result_bad_request, try_result_not_found, try_unwrap_body, RedirectParams, + admission_failure_response, create_api_proxy_user, create_catchup_session_key, + create_session_fingerprint, empty_json_response_as_array, empty_json_response_as_object, + force_provider_stream_response, get_session_reservation_ttl_secs, get_user_target, + get_user_target_by_credentials, internal_server_error, is_seek_request, is_session_based_playback, + is_stream_share_enabled, local_stream_response, + redirect, redirect_response, resource_response, + separate_number_and_remainder, should_allow_exhausted_shared_reconnect, stream_response, + try_option_bad_request, try_result_bad_request, try_result_not_found, + try_unwrap_body, RedirectParams, }, endpoints::{ hls_api::handle_hls_stream_request, @@ -62,6 +64,9 @@ use std::{ str::FromStr, sync::Arc, }; +use crate::auth::{check_network_access_only, check_permission_and_network_access_only}; +// https://github.com/tellytv/go.xtream-codes/blob/master/structs.go +// Xtream api -> https://9tzx6f0ozj.apidog.io/ #[derive(Serialize, Deserialize, Debug, Copy, Clone, Eq, PartialEq)] pub enum ApiStreamContext { @@ -207,7 +212,7 @@ async fn xtream_player_api_stream( app_state: &Arc, api_req: &UserApiRequest, stream_req: ApiStreamRequest<'_>, - user_target: Option<(ProxyUserCredentials, Arc)>, + user_target: Option<(Arc, Arc)>, ) -> impl IntoResponse + Send { // if log::log_enabled!(log::Level::Debug) { // debug!( @@ -225,14 +230,18 @@ async fn xtream_player_api_stream( let auth_status = app_state.app_config.get_auth_error_status(); let (user, target) = match user_target { - None => try_option_forbidden!( - get_user_target_by_credentials(stream_req.username, stream_req.password, api_req, app_state), - auth_status, - false, - format!("Could not find any user for xc stream {}", stream_req.username) - ), + None => { + let Some((user, target)) = get_user_target_by_credentials(stream_req.username, stream_req.password, api_req, app_state) else { + return auth_status.into_response(); + }; + (user, target) + } Some((user, target)) => (user, target), }; + // Network access check only - permission check is done later with full stream info + if let Err(e) = check_network_access_only(&user, fingerprint, app_state) { + return e.into_player_response(auth_status); + } let _guard = app_state.app_config.file_locks.write_lock_str(&user.username).await; @@ -713,20 +722,18 @@ async fn xtream_player_api_stream_with_token( } async fn xtream_player_api_resource( + fingerprint: &Fingerprint, req_headers: &HeaderMap, api_req: &UserApiRequest, app_state: &Arc, resource_req: ApiStreamRequest<'_>, ) -> impl IntoResponse { let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target_by_credentials(resource_req.username, resource_req.password, api_req, app_state), - auth_status, - false, - format!("Could not find any user xc resource {}", resource_req.username) - ); - if user.permission_denied(app_state) { - return axum::http::StatusCode::FORBIDDEN.into_response(); + let Some((user, target)) = get_user_target_by_credentials(resource_req.username, resource_req.password, api_req, app_state) else { + return auth_status.into_response(); + }; + if let Err(e) = check_permission_and_network_access_only(&user, fingerprint, app_state) { + return e.into_player_response(auth_status); } let target_name = &target.name; if !target.has_output(TargetType::Xtream) { @@ -787,6 +794,7 @@ macro_rules! create_xtream_player_api_stream { macro_rules! create_xtream_player_api_resource { ($fn_name:ident, $context:expr) => { async fn $fn_name( + fingerprint: Fingerprint, axum::extract::Path((username, password, stream_id, resource)): axum::extract::Path<( String, String, @@ -798,6 +806,7 @@ macro_rules! create_xtream_player_api_resource { req_headers: HeaderMap, ) -> impl IntoResponse { xtream_player_api_resource( + &fingerprint, &req_headers, &api_req, &app_state, @@ -845,12 +854,9 @@ async fn xtream_player_api_timeshift_stream( let api_req = UserApiRequest::merge_prefer_primary(&path_req, &query_req); let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target_by_credentials(&api_req.username, &api_req.password, &api_req, &app_state), - auth_status, - false, - format!("Could not find any user {}", api_req.username) - ); + let Some((user, target)) = get_user_target_by_credentials(&api_req.username, &api_req.password, &api_req, &app_state) else { + return auth_status.into_response(); + }; let epg_timeshift = parse_timeshift(user.epg_request_timeshift.as_deref()); let start = apply_timeshift(&api_req.start, &epg_timeshift); @@ -894,12 +900,9 @@ async fn xtream_player_api_timeshift_query_stream( } let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target_by_credentials(&api_req.username, &api_req.password, &api_req, &app_state), - auth_status, - false, - format!("Could not find any user {}", api_req.username) - ); + let Some((user, target)) = get_user_target_by_credentials(&api_req.username, &api_req.password, &api_req, &app_state) else { + return auth_status.into_response(); + }; let epg_timeshift = parse_timeshift(user.epg_request_timeshift.as_deref()); let start = apply_timeshift(&api_req.start, &epg_timeshift); @@ -1292,15 +1295,15 @@ macro_rules! skip_flag_optional { } #[allow(clippy::too_many_lines)] -async fn xtream_player_api(api_req: UserApiRequest, app_state: &Arc) -> impl IntoResponse + Send { +async fn xtream_player_api(fingerprint: &Fingerprint, api_req: UserApiRequest, app_state: &Arc) -> impl IntoResponse + Send { api_req.log_sanitized("xtream_player_api"); let auth_status = app_state.app_config.get_auth_error_status(); - let (user, target) = try_option_forbidden!( - get_user_target(&api_req, app_state), - auth_status, - false, - format!("Could not find any user for xc player api {}", api_req.username) - ); + let Some((user, target)) = get_user_target(&api_req, app_state) else { + return auth_status.into_response(); + }; + if let Err(e) = check_network_access_only(&user, fingerprint, app_state) { + return e.into_player_response(auth_status); + } if !target.has_output(TargetType::Xtream) { return get_user_info(&user, app_state).await.map_or_else( || axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response(), @@ -1460,17 +1463,19 @@ where } async fn xtream_player_api_get( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, axum::extract::Query(api_req): axum::extract::Query, ) -> impl IntoResponse + Send { - xtream_player_api(api_req, &app_state).await + xtream_player_api(&fingerprint, api_req, &app_state).await } async fn xtream_player_api_post( + fingerprint: Fingerprint, axum::extract::State(app_state): axum::extract::State>, UserApiRequestQueryOrBody(api_req): UserApiRequestQueryOrBody, ) -> impl IntoResponse + Send { - xtream_player_api(api_req, &app_state).await + xtream_player_api(&fingerprint, api_req, &app_state).await } macro_rules! register_xtream_api { diff --git a/backend/src/api/model/active_user_manager.rs b/backend/src/api/model/active_user_manager.rs index f91ca4e57..16f6de2b0 100644 --- a/backend/src/api/model/active_user_manager.rs +++ b/backend/src/api/model/active_user_manager.rs @@ -3992,6 +3992,7 @@ mod tests { soft_connections, soft_priority: 0, t_is_api_user: false, + network_access: None, } } diff --git a/backend/src/api/model/streams/provider_stream.rs b/backend/src/api/model/streams/provider_stream.rs index 31093b438..e4e2fff02 100644 --- a/backend/src/api/model/streams/provider_stream.rs +++ b/backend/src/api/model/streams/provider_stream.rs @@ -16,6 +16,7 @@ use serde::{Deserialize, Deserializer, Serialize, Serializer}; use shared::model::PlaylistItemType; use std::{fmt, net::SocketAddr, str::FromStr, sync::Arc}; use tokio_util::sync::CancellationToken; +use shared::error::TuliproxError; #[derive(Debug, Copy, Clone, PartialEq, PartialOrd, Eq, Ord, Hash)] pub enum CustomVideoStreamType { @@ -42,7 +43,7 @@ impl fmt::Display for CustomVideoStreamType { } impl FromStr for CustomVideoStreamType { - type Err = String; + type Err = TuliproxError; fn from_str(s: &str) -> Result { match s.to_lowercase().as_str() { @@ -52,7 +53,7 @@ impl FromStr for CustomVideoStreamType { "low_priority_preempted" => Ok(Self::LowPriorityPreempted), "user_account_expired" => Ok(Self::UserAccountExpired), "provisioning" => Ok(Self::Provisioning), - _ => Err(format!("Unknown stream type: {s}")), + _ => Err(TuliproxError::Config(format!("Unknown stream type: {s}"))), } } } diff --git a/backend/src/auth/api_user_context.rs b/backend/src/auth/api_user_context.rs new file mode 100644 index 000000000..8d3d0cc37 --- /dev/null +++ b/backend/src/auth/api_user_context.rs @@ -0,0 +1,175 @@ +use crate::{ + api::{ + api_utils::{ + evaluate_network_access, log_network_access_allowed_geoip_unavailable, log_network_access_denied, + NetworkAccessDecision, NetworkAccessDenyReason, + }, + model::AppState, + }, + model::ProxyUserPermissionDenyReason, +}; +use axum::response::IntoResponse; +use log::debug; +use shared::utils::sanitize_sensitive_info; +use std::sync::Arc; + +#[derive(Debug, Clone)] +pub enum PermissionDenyReason { + Expired, + Disabled, + Banned, + Inactive, +} + +impl From for PermissionDenyReason { + fn from(reason: ProxyUserPermissionDenyReason) -> Self { + match reason { + ProxyUserPermissionDenyReason::Expired | ProxyUserPermissionDenyReason::ExpiredStatus => { + PermissionDenyReason::Expired + } + ProxyUserPermissionDenyReason::Disabled => PermissionDenyReason::Disabled, + ProxyUserPermissionDenyReason::Banned => PermissionDenyReason::Banned, + ProxyUserPermissionDenyReason::Inactive => PermissionDenyReason::Inactive, + } + } +} + +#[derive(Debug)] +pub enum ApiUserAuthError { + AuthFailed, + PermissionDenied(PermissionDenyReason), + NetworkDenied(NetworkAccessDenyReason), +} + +impl std::fmt::Display for ApiUserAuthError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + ApiUserAuthError::AuthFailed => write!(f, "Authentication failed"), + ApiUserAuthError::PermissionDenied(reason) => match reason { + PermissionDenyReason::Expired => write!(f, "User access denied, expired"), + PermissionDenyReason::Disabled => write!(f, "User access denied, status disabled"), + PermissionDenyReason::Banned => write!(f, "User access denied, status banned"), + PermissionDenyReason::Inactive => write!(f, "User access denied, status inactive"), + }, + ApiUserAuthError::NetworkDenied(reason) => match reason { + NetworkAccessDenyReason::NoCidrMatch => write!(f, "Network access denied, no CIDR match"), + NetworkAccessDenyReason::NoCountryMatch => write!(f, "Network access denied, no country match"), + NetworkAccessDenyReason::GeoIpUnavailable => write!(f, "Network access denied, GeoIP unavailable"), + NetworkAccessDenyReason::CountryUnknown => write!(f, "Network access denied, country unknown"), + NetworkAccessDenyReason::MalformedClientIp => write!(f, "Network access denied, malformed client IP"), + }, + } + } +} + +impl ApiUserAuthError { + /// Returns a player-endpoint compatible response with configured auth error status and empty body. + /// This is used by proxy player endpoints (m3u, xtream, xmltv, etc.) where external players + /// expect empty responses with HTTP status codes only. + pub fn into_player_response(self, auth_error_status: axum::http::StatusCode) -> axum::response::Response { + let status = match &self { + ApiUserAuthError::AuthFailed => auth_error_status, + ApiUserAuthError::PermissionDenied(_) | ApiUserAuthError::NetworkDenied(_) => { + axum::http::StatusCode::FORBIDDEN + } + }; + status.into_response() + } +} + +#[derive(Debug)] +pub struct ApiUserContext { + pub user: Arc, + pub target: Arc, + pub fingerprint: crate::auth::Fingerprint, +} + +/// Checks only network access (no permission check). Used by stream endpoints +/// where permission check must happen later with full stream info for `admission_failure_response`. +pub fn check_network_access_only( + user: &Arc, + fingerprint: &crate::auth::Fingerprint, + app_state: &Arc, +) -> Result<(), ApiUserAuthError> { + let geoip_unavailable_policy = app_state.app_config.get_geoip_unavailable_policy(); + match evaluate_network_access(user, &fingerprint.client_ip, &app_state.geoip, geoip_unavailable_policy) { + NetworkAccessDecision::Allowed => Ok(()), + NetworkAccessDecision::AllowedGeoIpUnavailable => { + log_network_access_allowed_geoip_unavailable(&user.username, &fingerprint.client_ip); + Ok(()) + } + NetworkAccessDecision::Denied(reason) => { + log_network_access_denied(&user.username, &fingerprint.client_ip, reason.as_str()); + Err(ApiUserAuthError::NetworkDenied(reason)) + } + } +} + +/// Checks network access without logging. Used for high-frequency polling endpoints +/// (e.g., `HDHomeRun` `lineup_status`) where repeated log output would be noisy. +pub fn try_check_network_access_only( + user: &Arc, + fingerprint: &crate::auth::Fingerprint, + app_state: &Arc, +) -> Result<(), ApiUserAuthError> { + let geoip_unavailable_policy = app_state.app_config.get_geoip_unavailable_policy(); + match evaluate_network_access(user, &fingerprint.client_ip, &app_state.geoip, geoip_unavailable_policy) { + NetworkAccessDecision::Allowed | NetworkAccessDecision::AllowedGeoIpUnavailable => Ok(()), + NetworkAccessDecision::Denied(reason) => Err(ApiUserAuthError::NetworkDenied(reason)), + } +} + +pub fn resolve_api_user_context( + user: Arc, + target: Arc, + fingerprint: crate::auth::Fingerprint, + app_state: &Arc, +) -> Result { + // Permission check + if let Some(reason) = user.permission_denied_reason(app_state) { + debug!("User access denied for {}: {:?}", sanitize_sensitive_info(&user.username), reason); + return Err(ApiUserAuthError::PermissionDenied(reason.into())); + } + + // Network access check with policy + let geoip_unavailable_policy = app_state.app_config.get_geoip_unavailable_policy(); + match evaluate_network_access(&user, &fingerprint.client_ip, &app_state.geoip, geoip_unavailable_policy) { + NetworkAccessDecision::Allowed => Ok(ApiUserContext { user, target, fingerprint }), + NetworkAccessDecision::AllowedGeoIpUnavailable => { + log_network_access_allowed_geoip_unavailable(&user.username, &fingerprint.client_ip); + Ok(ApiUserContext { user, target, fingerprint }) + } + NetworkAccessDecision::Denied(reason) => { + log_network_access_denied(&user.username, &fingerprint.client_ip, reason.as_str()); + Err(ApiUserAuthError::NetworkDenied(reason)) + } + } +} + +/// Checks permission and network access without taking ownership. Used when +/// user/target are needed after the auth check (avoids unnecessary Arc clones). +pub fn check_permission_and_network_access_only( + user: &Arc, + fingerprint: &crate::auth::Fingerprint, + app_state: &Arc, +) -> Result<(), ApiUserAuthError> { + // Permission check + if let Some(reason) = user.permission_denied_reason(app_state) { + debug!("User access denied for {}: {:?}", sanitize_sensitive_info(&user.username), reason); + return Err(ApiUserAuthError::PermissionDenied(reason.into())); + } + + // Network access check + let geoip_unavailable_policy = app_state.app_config.get_geoip_unavailable_policy(); + match evaluate_network_access(user, &fingerprint.client_ip, &app_state.geoip, geoip_unavailable_policy) { + NetworkAccessDecision::Allowed => Ok(()), + NetworkAccessDecision::AllowedGeoIpUnavailable => { + log_network_access_allowed_geoip_unavailable(&user.username, &fingerprint.client_ip); + Ok(()) + } + NetworkAccessDecision::Denied(reason) => { + log_network_access_denied(&user.username, &fingerprint.client_ip, reason.as_str()); + Err(ApiUserAuthError::NetworkDenied(reason)) + } + } +} diff --git a/backend/src/auth/mod.rs b/backend/src/auth/mod.rs index 1cc8e3899..0b99523c0 100644 --- a/backend/src/auth/mod.rs +++ b/backend/src/auth/mod.rs @@ -6,6 +6,8 @@ mod auth_bearer; mod auth_basic; mod access_token; mod fingerprint; +mod api_user_context; + type Rejection = (StatusCode, &'static str); #[macro_export] @@ -27,3 +29,4 @@ pub use self::password::*; pub use self::fingerprint::*; pub use self::auth_basic::*; pub use self::auth_bearer::*; +pub use self::api_user_context::*; \ No newline at end of file diff --git a/backend/src/model/config/api_proxy.rs b/backend/src/model/config/api_proxy.rs index caf253268..0b050f83d 100644 --- a/backend/src/model/config/api_proxy.rs +++ b/backend/src/model/config/api_proxy.rs @@ -218,10 +218,10 @@ impl ApiProxyConfig { } } - pub fn get_target_name(&self, username: &str, password: &str) -> Option<(ProxyUserCredentials, String)> { + pub fn get_target_name(&self, username: &str, password: &str) -> Option<(Arc, String)> { for target_user in &self.user { if let Some((credentials, target_name)) = target_user.get_target_name(username, password) { - return Some((credentials.clone(), target_name.to_string())); + return Some((Arc::clone(&credentials), target_name.to_string())); } } if log::log_enabled!(log::Level::Debug) && !username.eq(API_USER) { @@ -230,22 +230,22 @@ impl ApiProxyConfig { None } - pub fn get_target_name_by_token(&self, token: &str) -> Option<(ProxyUserCredentials, String)> { + pub fn get_target_name_by_token(&self, token: &str) -> Option<(Arc, String)> { for target_user in &self.user { if let Some((credentials, target_name)) = target_user.get_target_name_by_token(token) { - return Some((credentials.clone(), target_name.to_string())); + return Some((Arc::clone(&credentials), target_name.to_string())); } } None } - pub fn get_user_credentials(&self, username: &str) -> Option { + pub fn get_user_credentials(&self, username: &str) -> Option> { let result = self .user .iter() - .flat_map(|target_user| &target_user.credentials) + .flat_map(|target_user| target_user.credentials.iter()) .find(|credential| credential.username == username) - .cloned(); + .map(Arc::clone); if result.is_none() && (username != TEST_USER && username != API_USER) { debug!("Could not find any user credentials for: {username}"); } diff --git a/backend/src/model/config/api_user.rs b/backend/src/model/config/api_user.rs index f90b84ddd..c9e89dfbd 100644 --- a/backend/src/model/config/api_user.rs +++ b/backend/src/model/config/api_user.rs @@ -3,14 +3,95 @@ use crate::model::{macros, Config}; use arc_swap::access::Access; use arc_swap::ArcSwap; use chrono::Local; -use log::debug; +use log::{debug, warn}; use shared::model::{ - ClusterFlags, ProxyType, ProxyUserCredentialsDto, ProxyUserStatus, TargetUserDto, UserConnectionPermission, - XtreamCluster, + ClusterFlags, NetworkAccessDto, ProxyType, ProxyUserCredentialsDto, ProxyUserStatus, TargetUserDto, + UserConnectionPermission, XtreamCluster, }; use std::sync::Arc; use zeroize::Zeroize; +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ProxyUserPermissionDenyReason { + Expired, + Disabled, + Banned, + ExpiredStatus, + Inactive, +} + +#[derive(Debug, Clone, Default)] +pub struct NetworkAccess { + pub allowed_countries: Vec, + pub allowed_networks: Vec, +} + +impl NetworkAccess { + pub fn is_empty(&self) -> bool { + self.allowed_countries.is_empty() && self.allowed_networks.is_empty() + } +} + +impl From<&NetworkAccessDto> for NetworkAccess { + fn from(dto: &NetworkAccessDto) -> Self { + let mut seen_countries = std::collections::HashSet::new(); + let allowed_countries: Vec = dto + .allowed_countries + .as_ref() + .map(|countries| { + countries + .iter() + .filter_map(|c| { + let upper = c.trim().to_uppercase(); + if upper.is_empty() { + None + } else if seen_countries.insert(upper.clone()) { + Some(upper) + } else { + None + } + }) + .collect() + }) + .unwrap_or_default(); + + let allowed_networks: Vec = dto + .allowed_networks + .as_ref() + .map(|networks| { + networks + .iter() + .filter_map(|n| match n.trim().parse::() { + Ok(net) => Some(net), + Err(err) => { + warn!("Skipping invalid CIDR '{n}': {err}"); + None + } + }) + .collect() + }) + .unwrap_or_default(); + + Self { allowed_countries, allowed_networks } + } +} + +impl From<&NetworkAccess> for NetworkAccessDto { + fn from(instance: &NetworkAccess) -> Self { + let allowed_countries = if instance.allowed_countries.is_empty() { + None + } else { + Some(instance.allowed_countries.clone()) + }; + let allowed_networks = if instance.allowed_networks.is_empty() { + None + } else { + Some(instance.allowed_networks.iter().map(std::string::ToString::to_string).collect()) + }; + Self { allowed_countries, allowed_networks } + } +} + #[derive(Debug, Clone, Default)] pub struct ProxyUserCredentials { pub username: String, @@ -31,6 +112,7 @@ pub struct ProxyUserCredentials { pub soft_connections: u16, pub soft_priority: i8, pub t_is_api_user: bool, + pub network_access: Option, } macros::from_impl!(ProxyUserCredentials); @@ -55,6 +137,11 @@ impl From<&ProxyUserCredentialsDto> for ProxyUserCredentials { soft_connections: dto.soft_connections, soft_priority: dto.soft_priority, t_is_api_user: false, + network_access: dto + .network_access + .as_ref() + .map(NetworkAccess::from) + .filter(|network_access| !network_access.is_empty()), } } } @@ -79,6 +166,7 @@ impl From<&ProxyUserCredentials> for ProxyUserCredentialsDto { priority: instance.priority, soft_connections: instance.soft_connections, soft_priority: instance.soft_priority, + network_access: instance.network_access.as_ref().map(NetworkAccessDto::from), } } } @@ -95,30 +183,50 @@ impl ProxyUserCredentials { self.username.eq(username) && self.password.eq(password) } + #[inline] pub fn has_permissions(&self, app_state: &AppState) -> bool { + self.permission_denied_reason(app_state).is_none() + } + + #[inline] + pub fn permission_denied(&self, app_state: &AppState) -> bool { + !self.has_permissions(app_state) + } + + pub fn permission_denied_reason(&self, app_state: &AppState) -> Option { let config = > as Access>::load(&app_state.app_config.config); if config.user_access_control { if let Some(exp_date) = self.exp_date.as_ref() { let now = Local::now(); if (exp_date - now.timestamp()) < 0 { debug!("User access denied, expired: {}", self.username); - return false; + return Some(ProxyUserPermissionDenyReason::Expired); } } if let Some(status) = &self.status { - if !matches!(status, ProxyUserStatus::Active | ProxyUserStatus::Trial) { - debug!("User access denied, status invalid: {status} for user: {}", self.username); - return false; + match status { + ProxyUserStatus::Disabled => { + debug!("User access denied, status disabled: {}", self.username); + return Some(ProxyUserPermissionDenyReason::Disabled); + } + ProxyUserStatus::Banned => { + debug!("User access denied, status banned: {}", self.username); + return Some(ProxyUserPermissionDenyReason::Banned); + } + ProxyUserStatus::Expired => { + debug!("User access denied, status expired: {}", self.username); + return Some(ProxyUserPermissionDenyReason::ExpiredStatus); + } + ProxyUserStatus::Active | ProxyUserStatus::Trial => {} + ProxyUserStatus::Pending => { + debug!("User access denied, status pending: {}", self.username); + return Some(ProxyUserPermissionDenyReason::Inactive); + } } - } // NO STATUS SET, ok admins fault, we take this as a valid status + } } - true - } - - #[inline] - pub fn permission_denied(&self, app_state: &AppState) -> bool { - !self.has_permissions(app_state) + None } pub fn allows_cluster(&self, cluster: XtreamCluster) -> bool { @@ -152,30 +260,39 @@ impl Drop for ProxyUserCredentials { #[derive(Debug, Clone)] pub struct TargetUser { pub target: String, - pub credentials: Vec, + pub credentials: Vec>, } macros::from_impl!(TargetUser); impl From<&TargetUserDto> for TargetUser { fn from(dto: &TargetUserDto) -> Self { - Self { target: dto.target.clone(), credentials: dto.credentials.iter().map(Into::into).collect() } + Self { + target: dto.target.clone(), + credentials: dto.credentials.iter().map(|c| Arc::new(c.into())).collect(), + } } } impl From<&TargetUser> for TargetUserDto { fn from(instance: &TargetUser) -> Self { - Self { target: instance.target.clone(), credentials: instance.credentials.iter().map(Into::into).collect() } + Self { + target: instance.target.clone(), + credentials: instance.credentials.iter().map(|c| c.as_ref().into()).collect(), + } } } impl TargetUser { - pub fn get_target_name(&self, username: &str, password: &str) -> Option<(&ProxyUserCredentials, &str)> { + pub fn get_target_name(&self, username: &str, password: &str) -> Option<(Arc, &str)> { self.credentials .iter() .find(|c| c.matches(username, password)) - .map(|credentials| (credentials, self.target.as_str())) + .map(|credentials| (Arc::clone(credentials), self.target.as_str())) } - pub fn get_target_name_by_token(&self, token: &str) -> Option<(&ProxyUserCredentials, &str)> { - self.credentials.iter().find(|c| c.matches_token(token)).map(|credentials| (credentials, self.target.as_str())) + pub fn get_target_name_by_token(&self, token: &str) -> Option<(Arc, &str)> { + self.credentials + .iter() + .find(|c| c.matches_token(token)) + .map(|credentials| (Arc::clone(credentials), self.target.as_str())) } } diff --git a/backend/src/model/config/app.rs b/backend/src/model/config/app.rs index 227a72821..ef38976e8 100644 --- a/backend/src/model/config/app.rs +++ b/backend/src/model/config/app.rs @@ -1,15 +1,15 @@ use crate::api::model::TransportStreamBuffer; use crate::model::{ ApiProxyConfig, ApiProxyServerInfo, Config, ConfigInput, ConfigInputOptions, ConfigTarget, CustomStreamResponse, - GracePeriodOptions, HdHomeRunConfig, HdHomeRunFlags, Mappings, MediaToolCapabilities, ProxyUserCredentials, - ReverseProxyDisabledHeaderConfig, SourcesConfig, TargetOutput, + GracePeriodOptions, HdHomeRunConfig, HdHomeRunFlags, Mappings, MediaToolCapabilities, + ProxyUserCredentials, ReverseProxyDisabledHeaderConfig, SourcesConfig, TargetOutput, }; use crate::utils; use arc_swap::{ArcSwap, ArcSwapOption}; use log::{error, warn}; use rand::Rng; use shared::error::TuliproxError; -use shared::model::ConfigPaths; +use shared::model::{ConfigPaths, GeoIpUnavailablePolicy}; use shared::utils::{ CHANNEL_UNAVAILABLE, LOW_PRIORITY_PREEMPTED, PANEL_API_PROVISIONING, PROVIDER_CONNECTIONS_EXHAUSTED, USER_ACCOUNT_EXPIRED, USER_CONNECTIONS_EXHAUSTED, @@ -153,14 +153,14 @@ impl AppConfig { config.reverse_proxy.as_ref().map(|r| r.rewrite_secret) } - fn intern_get_target_for_user(&self, user_target: Option<(ProxyUserCredentials, String)>) -> Option<(ProxyUserCredentials, Arc)> { + fn intern_get_target_for_user(&self, user_target: Option<(Arc, String)>) -> Option<(Arc, Arc)> { match user_target { Some((user, target_name)) => { let sources = self.sources.load(); for source in &sources.sources { for target in &source.targets { if target_name.eq_ignore_ascii_case(&target.name) { - return Some((user, Arc::clone(target))); + return Some((Arc::clone(&user), Arc::clone(target))); } } } @@ -186,7 +186,7 @@ impl AppConfig { None } - pub fn get_target_for_username(&self, username: &str) -> Option<(ProxyUserCredentials, Arc)> { + pub fn get_target_for_username(&self, username: &str) -> Option<(Arc, Arc)> { if let Some(credentials) = self.get_user_credentials(username) { return self.api_proxy.load().as_ref() .and_then(|api_proxy| self.intern_get_target_for_user(api_proxy.get_target_name(&credentials.username, &credentials.password))); @@ -194,15 +194,15 @@ impl AppConfig { None } - pub fn get_target_for_user(&self, username: &str, password: &str) -> Option<(ProxyUserCredentials, Arc)> { + pub fn get_target_for_user(&self, username: &str, password: &str) -> Option<(Arc, Arc)> { self.api_proxy.load().as_ref().and_then(|api_proxy| self.intern_get_target_for_user(api_proxy.get_target_name(username, password))) } - pub fn get_target_for_user_by_token(&self, token: &str) -> Option<(ProxyUserCredentials, Arc)> { - self.api_proxy.load().as_ref().as_ref().and_then(|api_proxy| self.intern_get_target_for_user(api_proxy.get_target_name_by_token(token))) + pub fn get_target_for_user_by_token(&self, token: &str) -> Option<(Arc, Arc)> { + self.api_proxy.load().as_ref().and_then(|api_proxy| self.intern_get_target_for_user(api_proxy.get_target_name_by_token(token))) } - pub fn get_user_credentials(&self, username: &str) -> Option { + pub fn get_user_credentials(&self, username: &str) -> Option> { self.api_proxy.load().as_ref().as_ref().and_then(|api_proxy| api_proxy.get_user_credentials(username)) } @@ -466,6 +466,10 @@ impl AppConfig { self.config.load().get_grace_options() } + pub fn get_geoip_unavailable_policy(&self) -> GeoIpUnavailablePolicy { + self.config.load().get_geoip_unavailable_policy() + } + pub async fn is_ffprobe_enabled(&self) -> bool { let ffprobe_enabled_in_config = { let config = self.config.load(); diff --git a/backend/src/model/config/base.rs b/backend/src/model/config/base.rs index 676603c42..d50f01b4a 100644 --- a/backend/src/model/config/base.rs +++ b/backend/src/model/config/base.rs @@ -1,13 +1,13 @@ use crate::model::{ - macros, ConfigApi, HdHomeRunConfig, HdHomeRunFlags, IpCheckConfig, LibraryConfig, LogConfig, MetadataUpdateConfig, - MessagingConfig, ProxyConfig, ReverseProxyConfig, ReverseProxyDisabledHeaderConfig, ScheduleConfig, VideoConfig, - WebUiConfig, + macros, ConfigApi, HdHomeRunConfig, HdHomeRunFlags, IpCheckConfig, LibraryConfig, + LogConfig, MetadataUpdateConfig, MessagingConfig, ProxyConfig, ReverseProxyConfig, + ReverseProxyDisabledHeaderConfig, ScheduleConfig, VideoConfig, WebUiConfig, }; use crate::utils; use log::{error, info}; use path_clean::PathClean; use shared::error::TuliproxError; -use shared::model::{ConfigDto, HdHomeRunDeviceOverview}; +use shared::model::{ConfigDto, GeoIpUnavailablePolicy, HdHomeRunDeviceOverview}; use shared::utils::{default_grace_period_millis, default_grace_period_timeout_secs, set_sanitize_sensitive_info, DEFAULT_BACKUP_DIR, DEFAULT_CACHE_DIR, DEFAULT_DOWNLOAD_DIR, DEFAULT_STORAGE_DIR, DEFAULT_STORAGE_TEMP_DIR, DEFAULT_USER_CONFIG_DIR}; use std::borrow::Cow; use std::path::{Path, PathBuf}; @@ -262,6 +262,13 @@ impl Config { self.reverse_proxy.as_ref().is_some_and(|r| r.geoip.as_ref().is_some_and(|g| g.enabled)) } + pub fn get_geoip_unavailable_policy(&self) -> GeoIpUnavailablePolicy { + self.reverse_proxy + .as_ref() + .and_then(|r| r.geoip.as_ref()) + .map_or(GeoIpUnavailablePolicy::Deny, |g| g.unavailable_policy) + } + pub fn get_disabled_headers(&self) -> Option { self.reverse_proxy .as_ref() @@ -321,7 +328,7 @@ impl From<&ConfigDto> for Config { #[cfg(test)] mod tests { use super::Config; - use shared::model::ConfigDto; + use shared::model::{ConfigDto, GeoIpUnavailablePolicy}; use tempfile::tempdir; #[test] @@ -402,4 +409,55 @@ mod tests { assert_eq!(stream_history.stream_history_retention_days, 14); assert_eq!(stream_history.stream_history_directory, "/var/lib/tuliprox/history"); } + + #[test] + fn get_geoip_unavailable_policy_no_reverse_proxy_returns_deny() { + let dto = ConfigDto::default(); + let config = Config::from(&dto); + assert_eq!(config.get_geoip_unavailable_policy(), GeoIpUnavailablePolicy::Deny); + } + + #[test] + fn get_geoip_unavailable_policy_no_geoip_returns_deny() { + let dto = ConfigDto { + reverse_proxy: Some(shared::model::ReverseProxyConfigDto::default()), + ..Default::default() + }; + let config = Config::from(&dto); + assert_eq!(config.get_geoip_unavailable_policy(), GeoIpUnavailablePolicy::Deny); + } + + #[test] + fn get_geoip_unavailable_policy_missing_field_returns_deny() { + let dto = ConfigDto { + reverse_proxy: Some(shared::model::ReverseProxyConfigDto { + geoip: Some(shared::model::GeoIpConfigDto { + enabled: true, + url: "https://example.com/db.csv".to_string(), + ..shared::model::GeoIpConfigDto::default() + }), + ..Default::default() + }), + ..Default::default() + }; + let config = Config::from(&dto); + assert_eq!(config.get_geoip_unavailable_policy(), GeoIpUnavailablePolicy::Deny); + } + + #[test] + fn get_geoip_unavailable_policy_allow_returns_allow() { + let dto = ConfigDto { + reverse_proxy: Some(shared::model::ReverseProxyConfigDto { + geoip: Some(shared::model::GeoIpConfigDto { + enabled: true, + url: "https://example.com/db.csv".to_string(), + unavailable_policy: shared::model::GeoIpUnavailablePolicy::Allow, + }), + ..Default::default() + }), + ..Default::default() + }; + let config = Config::from(&dto); + assert_eq!(config.get_geoip_unavailable_policy(), GeoIpUnavailablePolicy::Allow); + } } diff --git a/backend/src/model/config/geoip.rs b/backend/src/model/config/geoip.rs index 88d44c611..23f8ddb92 100644 --- a/backend/src/model/config/geoip.rs +++ b/backend/src/model/config/geoip.rs @@ -1,10 +1,11 @@ -use shared::model::GeoIpConfigDto; +use shared::model::{GeoIpConfigDto, GeoIpUnavailablePolicy}; use crate::model::macros; #[derive(Debug, Clone)] pub struct GeoIpConfig { pub(crate) enabled: bool, pub(crate) url: String, + pub(crate) unavailable_policy: GeoIpUnavailablePolicy, } macros::from_impl!(GeoIpConfig); @@ -14,6 +15,7 @@ impl From<&GeoIpConfigDto> for GeoIpConfig { Self { enabled: dto.enabled, url: dto.url.clone(), + unavailable_policy: dto.unavailable_policy, } } } @@ -23,6 +25,45 @@ impl From<&GeoIpConfig> for GeoIpConfigDto { Self { enabled: instance.enabled, url: instance.url.clone(), + unavailable_policy: instance.unavailable_policy, } } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn geoip_unavailable_policy_dto_default_converts_to_deny() { + let dto = GeoIpConfigDto { + enabled: false, + url: String::new(), + unavailable_policy: GeoIpUnavailablePolicy::Deny, + }; + let config: GeoIpConfig = (&dto).into(); + assert_eq!(config.unavailable_policy, GeoIpUnavailablePolicy::Deny); + } + + #[test] + fn geoip_unavailable_policy_dto_allow_converts_to_allow() { + let dto = GeoIpConfigDto { + enabled: true, + url: "https://example.com/db.csv".to_string(), + unavailable_policy: GeoIpUnavailablePolicy::Allow, + }; + let config: GeoIpConfig = (&dto).into(); + assert_eq!(config.unavailable_policy, GeoIpUnavailablePolicy::Allow); + } + + #[test] + fn geoip_unavailable_policy_domain_allow_converts_to_dto_allow() { + let config = GeoIpConfig { + enabled: true, + url: "https://example.com/db.csv".to_string(), + unavailable_policy: GeoIpUnavailablePolicy::Allow, + }; + let dto: GeoIpConfigDto = (&config).into(); + assert_eq!(dto.unavailable_policy, GeoIpUnavailablePolicy::Allow); + } +} diff --git a/backend/src/repository/bplustree_migration.rs b/backend/src/repository/bplustree_migration.rs index 7b62e9106..aa5453551 100644 --- a/backend/src/repository/bplustree_migration.rs +++ b/backend/src/repository/bplustree_migration.rs @@ -1,8 +1,10 @@ -use super::bplustree::{BPlusTree, MAGIC, STORAGE_VERSION}; -use super::storage_const; +use super::{ + bplustree::{BPlusTree, MAGIC, STORAGE_VERSION}, + storage_const, +}; use fs2::FileExt as _; use log::{info, trace, warn}; -use shared::model::{ClusterFlags, ConfigPaths, ProxyType, ProxyUserStatus}; +use shared::model::{ClusterFlags, ConfigPaths, NetworkAccessDto, ProxyType, ProxyUserStatus}; use std::{ collections::{HashSet, VecDeque}, ffi::OsStr, @@ -19,7 +21,7 @@ const HEADER_FLAG_HAS_TOMBSTONES: u32 = 1 << 30; const HEADER_METADATA_LEN_MASK: u32 = !(HEADER_FLAG_HAS_METADATA_FLAGS | HEADER_FLAG_HAS_TOMBSTONES); const MARKER_FILE_GUARD_PREFIX: &str = ".db_mergeto_v"; const MARKER_FILE_GUARD_PREFIX_LEGACY_ALT: &str = ".db_mergedto"; -const MARKER_FILE_API_USER_GUARD: &str = ".userdb_mergeto_v5"; +const MARKER_FILE_API_USER_GUARD: &str = ".userdb_mergeto_v6"; const MARKER_VERSION_KEY: &str = "migrated_to"; const MARKER_ROOTS_FINGERPRINT_KEY: &str = "roots_fingerprint"; @@ -38,9 +40,7 @@ struct BPlusTreeStartupMigrator { } impl BPlusTreeStartupMigrator { - pub fn new(roots: Vec) -> Self { - Self { roots, migration_marker_path: None } - } + pub fn new(roots: Vec) -> Self { Self { roots, migration_marker_path: None } } pub fn new_with_marker(roots: Vec, migration_marker_path: PathBuf) -> Self { Self { roots, migration_marker_path: Some(migration_marker_path) } @@ -380,9 +380,7 @@ pub fn migrate_bplustree_databases(roots: &[PathBuf]) -> io::Result PathBuf { - marker_dir.join(marker_file_name()) -} +pub fn bplustree_migration_marker_path(marker_dir: &Path) -> PathBuf { marker_dir.join(marker_file_name()) } pub fn migrate_bplustree_databases_with_marker( roots: &[PathBuf], @@ -392,9 +390,7 @@ pub fn migrate_bplustree_databases_with_marker( BPlusTreeStartupMigrator::new_with_marker(roots.to_vec(), marker_path).run() } -fn marker_file_name() -> String { - format!("{MARKER_FILE_GUARD_PREFIX}{STORAGE_VERSION}") -} +fn marker_file_name() -> String { format!("{MARKER_FILE_GUARD_PREFIX}{STORAGE_VERSION}") } // // The user database has gone through six serialization schemas (MessagePack, @@ -412,6 +408,7 @@ fn marker_file_name() -> String { // overwrite the freshly migrated data. #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] struct StoredApiUserV1 { pub target: String, pub username: String, @@ -429,6 +426,7 @@ struct StoredApiUserV1 { } #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] struct StoredApiUserV2 { pub target: String, pub username: String, @@ -449,6 +447,7 @@ struct StoredApiUserV2 { // V3 mirror — same layout as user_repository::StoredProxyUserCredentials. // Defined here so the migration has no dependency on user_repository internals. #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] struct StoredApiUserV3 { pub target: String, pub username: String, @@ -512,6 +511,7 @@ impl StoredApiUserV3 { // V4 mirror — same layout as user_repository::StoredProxyUserCredentials. // Defined here so the migration has no dependency on user_repository internals. #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] struct StoredApiUserV4 { pub target: String, pub username: String, @@ -555,18 +555,15 @@ impl StoredApiUserV4 { } } - fn from_v2(v2: &StoredApiUserV2) -> Self { - Self::from_v3(&StoredApiUserV3::from_v2(v2)) - } + fn from_v2(v2: &StoredApiUserV2) -> Self { Self::from_v3(&StoredApiUserV3::from_v2(v2)) } - fn from_v1(v1: &StoredApiUserV1) -> Self { - Self::from_v3(&StoredApiUserV3::from_v1(v1)) - } + fn from_v1(v1: &StoredApiUserV1) -> Self { Self::from_v3(&StoredApiUserV3::from_v1(v1)) } } -// V5 mirror — same layout as user_repository::StoredProxyUserCredentials. +// V5 mirror — same layout as the previous user_repository::StoredProxyUserCredentials. // Defined here so the migration has no dependency on user_repository internals. #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] struct StoredApiUserV5 { pub target: String, pub username: String, @@ -612,17 +609,71 @@ impl StoredApiUserV5 { } } - fn from_v3(v3: &StoredApiUserV3) -> Self { - Self::from_v4(&StoredApiUserV4::from_v3(v3)) + fn from_v3(v3: &StoredApiUserV3) -> Self { Self::from_v4(&StoredApiUserV4::from_v3(v3)) } + + fn from_v2(v2: &StoredApiUserV2) -> Self { Self::from_v4(&StoredApiUserV4::from_v2(v2)) } + + fn from_v1(v1: &StoredApiUserV1) -> Self { Self::from_v4(&StoredApiUserV4::from_v1(v1)) } +} + +// V6 mirror — same layout as user_repository::StoredProxyUserCredentials. +// Defined here so the migration has no dependency on user_repository internals. +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] +struct StoredApiUserV6 { + pub target: String, + pub username: String, + pub password: String, + pub token: Option, + pub proxy: ProxyType, + pub server: Option, + pub epg_timeshift: Option, + pub epg_request_timeshift: Option, + pub created_at: Option, + pub exp_date: Option, + pub max_connections: Option, + pub status: Option, + pub output_clusters: ClusterFlags, + pub ui_enabled: bool, + pub comment: Option, + pub priority: Option, + pub soft_connections: Option, + pub soft_priority: Option, + pub network_access: Option, +} + +impl StoredApiUserV6 { + fn from_v5(v5: &StoredApiUserV5) -> Self { + Self { + target: v5.target.clone(), + username: v5.username.clone(), + password: v5.password.clone(), + token: v5.token.clone(), + proxy: v5.proxy, + server: v5.server.clone(), + epg_timeshift: v5.epg_timeshift.clone(), + epg_request_timeshift: v5.epg_request_timeshift.clone(), + created_at: v5.created_at, + exp_date: v5.exp_date, + max_connections: v5.max_connections, + status: v5.status, + output_clusters: v5.output_clusters, + ui_enabled: v5.ui_enabled, + comment: v5.comment.clone(), + priority: v5.priority, + soft_connections: v5.soft_connections, + soft_priority: v5.soft_priority, + network_access: None, + } } - fn from_v2(v2: &StoredApiUserV2) -> Self { - Self::from_v4(&StoredApiUserV4::from_v2(v2)) - } + fn from_v4(v4: &StoredApiUserV4) -> Self { Self::from_v5(&StoredApiUserV5::from_v4(v4)) } - fn from_v1(v1: &StoredApiUserV1) -> Self { - Self::from_v4(&StoredApiUserV4::from_v1(v1)) - } + fn from_v3(v3: &StoredApiUserV3) -> Self { Self::from_v5(&StoredApiUserV5::from_v3(v3)) } + + fn from_v2(v2: &StoredApiUserV2) -> Self { Self::from_v5(&StoredApiUserV5::from_v2(v2)) } + + fn from_v1(v1: &StoredApiUserV1) -> Self { Self::from_v5(&StoredApiUserV5::from_v1(v1)) } } fn create_user_db_merge_guard(merge_guard_path: &Path) -> io::Result<()> { @@ -632,68 +683,76 @@ fn create_user_db_merge_guard(merge_guard_path: &Path) -> io::Result<()> { Ok(()) } -pub(crate) fn user_db_merge_guard_path(config_dir: &Path) -> PathBuf { - config_dir.join(MARKER_FILE_API_USER_GUARD) -} +pub(crate) fn user_db_merge_guard_path(config_dir: &Path) -> PathBuf { config_dir.join(MARKER_FILE_API_USER_GUARD) } -/// Migrates the user database file from V1, V2, V3, or V4 schema to V5 (current) in +/// Migrates the user database file from V1-V5 schema to V6 (current) in /// place and creates a merge-guard file so config-driven merges are skipped /// until the operator explicitly removes it. /// /// Returns `true` when a migration was performed, `false` when the file was -/// already in V5 format or did not exist. +/// already in V6 format or did not exist. fn migrate_user_db_schema(db_path: &Path, merge_guard_path: &Path) -> io::Result { if !db_path.exists() { return Ok(false); } - if BPlusTree::::load(db_path).is_ok() { + if let Ok(tree) = BPlusTree::::load(db_path) { + let mut v6_tree: BPlusTree = BPlusTree::new(); + for (key, v5) in &tree { + v6_tree.insert(key.clone(), StoredApiUserV6::from_v5(v5)); + } + create_user_db_merge_guard(merge_guard_path)?; + v6_tree.store(db_path)?; + return Ok(true); + } + + if BPlusTree::::load(db_path).is_ok() { return Ok(false); } if let Ok(tree) = BPlusTree::::load(db_path) { - let mut v5_tree: BPlusTree = BPlusTree::new(); + let mut v6_tree: BPlusTree = BPlusTree::new(); for (key, v4) in &tree { - v5_tree.insert(key.clone(), StoredApiUserV5::from_v4(v4)); + v6_tree.insert(key.clone(), StoredApiUserV6::from_v4(v4)); } create_user_db_merge_guard(merge_guard_path)?; - v5_tree.store(db_path)?; + v6_tree.store(db_path)?; return Ok(true); } if let Ok(tree) = BPlusTree::::load(db_path) { - let mut v5_tree: BPlusTree = BPlusTree::new(); + let mut v6_tree: BPlusTree = BPlusTree::new(); for (key, v3) in &tree { - v5_tree.insert(key.clone(), StoredApiUserV5::from_v3(v3)); + v6_tree.insert(key.clone(), StoredApiUserV6::from_v3(v3)); } create_user_db_merge_guard(merge_guard_path)?; - v5_tree.store(db_path)?; + v6_tree.store(db_path)?; return Ok(true); } if let Ok(tree) = BPlusTree::::load(db_path) { - let mut v5_tree: BPlusTree = BPlusTree::new(); + let mut v6_tree: BPlusTree = BPlusTree::new(); for (key, v2) in &tree { - v5_tree.insert(key.clone(), StoredApiUserV5::from_v2(v2)); + v6_tree.insert(key.clone(), StoredApiUserV6::from_v2(v2)); } create_user_db_merge_guard(merge_guard_path)?; - v5_tree.store(db_path)?; + v6_tree.store(db_path)?; return Ok(true); } if let Ok(tree) = BPlusTree::::load(db_path) { - let mut v5_tree: BPlusTree = BPlusTree::new(); + let mut v6_tree: BPlusTree = BPlusTree::new(); for (key, v1) in &tree { - v5_tree.insert(key.clone(), StoredApiUserV5::from_v1(v1)); + v6_tree.insert(key.clone(), StoredApiUserV6::from_v1(v1)); } create_user_db_merge_guard(merge_guard_path)?; - v5_tree.store(db_path)?; + v6_tree.store(db_path)?; return Ok(true); } Err(io::Error::new( io::ErrorKind::InvalidData, - format!("User DB at '{}' exists but could not be read as V1, V2, V3, V4, or V5 format", db_path.display()), + format!("User DB at '{}' exists but could not be read as V1, V2, V3, V4, V5, or V6 format", db_path.display()), )) } @@ -705,7 +764,7 @@ pub struct AllStartupMigrationStats { /// Runs all startup migrations in sequence: /// 1. B+Tree storage-format migration (V1 -> current binary format) -/// 2. User DB schema migration (V1/V2/V3/V4 -> V5 `MessagePack` layout) +/// 2. User DB schema migration (V1-V5 -> V6 `MessagePack` layout) /// /// `config_dir` is the directory that contains `api_user.db` and the merge-guard /// marker. `storage_dir` is used for the B+Tree migration marker. @@ -754,7 +813,7 @@ pub fn run_startup_migrations(config_paths: &ConfigPaths) { ); } if stats.user_db_migrated { - info!("User DB schema migrated to V5"); + info!("User DB schema migrated to V6"); } } Err(err) => { @@ -948,7 +1007,7 @@ mod tests { } #[test] - fn user_db_schema_migration_v2_to_v5_creates_merge_guard() -> io::Result<()> { + fn user_db_schema_migration_v2_to_v6_creates_merge_guard() -> io::Result<()> { let temp = tempdir()?; let db_path = temp.path().join(storage_const::API_USER_DB_FILE); let merge_guard_path = user_db_merge_guard_path(temp.path()); @@ -980,8 +1039,8 @@ mod tests { assert!(migrated); assert!(merge_guard_path.exists()); - let v5_tree = BPlusTree::::load(&db_path)?; - let user = v5_tree + let v6_tree = BPlusTree::::load(&db_path)?; + let user = v6_tree .query(&"alice".to_string()) .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "alice missing after migration"))?; assert_eq!(user.username, "alice"); @@ -990,12 +1049,13 @@ mod tests { assert_eq!(user.priority, None); assert_eq!(user.soft_connections, None); assert_eq!(user.soft_priority, None); + assert_eq!(user.network_access, None); Ok(()) } #[test] - fn user_db_schema_migration_v3_to_v5_creates_merge_guard() -> io::Result<()> { + fn user_db_schema_migration_v3_to_v6_creates_merge_guard() -> io::Result<()> { let temp = tempdir()?; let db_path = temp.path().join(storage_const::API_USER_DB_FILE); let merge_guard_path = user_db_merge_guard_path(temp.path()); @@ -1028,20 +1088,21 @@ mod tests { assert!(migrated); assert!(merge_guard_path.exists()); - let v5_tree = BPlusTree::::load(&db_path)?; - let user = v5_tree + let v6_tree = BPlusTree::::load(&db_path)?; + let user = v6_tree .query(&"bob".to_string()) .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "bob missing after migration"))?; assert_eq!(user.output_clusters, ClusterFlags::all()); assert_eq!(user.priority, Some(5)); assert_eq!(user.soft_connections, None); assert_eq!(user.soft_priority, None); + assert_eq!(user.network_access, None); Ok(()) } #[test] - fn user_db_schema_migration_v4_to_v5_creates_merge_guard() -> io::Result<()> { + fn user_db_schema_migration_v4_to_v6_creates_merge_guard() -> io::Result<()> { let temp = tempdir()?; let db_path = temp.path().join(storage_const::API_USER_DB_FILE); let merge_guard_path = user_db_merge_guard_path(temp.path()); @@ -1076,19 +1137,20 @@ mod tests { assert!(migrated); assert!(merge_guard_path.exists()); - let v5_tree = BPlusTree::::load(&db_path)?; - let user = v5_tree + let v6_tree = BPlusTree::::load(&db_path)?; + let user = v6_tree .query(&"carol".to_string()) .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "carol missing after migration"))?; assert_eq!(user.output_clusters, ClusterFlags::all()); assert_eq!(user.soft_connections, Some(2)); assert_eq!(user.soft_priority, Some(-4)); + assert_eq!(user.network_access, None); Ok(()) } #[test] - fn user_db_schema_v5_is_detected_without_writing_merge_guard() -> io::Result<()> { + fn user_db_schema_migration_v5_to_v6_creates_merge_guard() -> io::Result<()> { let temp = tempdir()?; let db_path = temp.path().join(storage_const::API_USER_DB_FILE); let merge_guard_path = user_db_merge_guard_path(temp.path()); @@ -1121,17 +1183,75 @@ mod tests { assert!(!merge_guard_path.exists()); let migrated = migrate_user_db_schema(&db_path, &merge_guard_path)?; - assert!(!migrated); - assert!(!merge_guard_path.exists()); + assert!(migrated); + assert!(merge_guard_path.exists()); - let v5_tree = BPlusTree::::load(&db_path)?; - let user = v5_tree + let v6_tree = BPlusTree::::load(&db_path)?; + let user = v6_tree .query(&"dave".to_string()) - .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "dave missing after v5 detection"))?; + .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "dave missing after v6 migration"))?; assert_eq!(user.output_clusters, ClusterFlags::Live | ClusterFlags::Vod); assert_eq!(user.priority, Some(5)); assert_eq!(user.soft_connections, Some(2)); assert_eq!(user.soft_priority, Some(-4)); + assert_eq!(user.network_access, None); + + Ok(()) + } + + #[test] + fn user_db_schema_v6_is_detected_without_writing_merge_guard() -> io::Result<()> { + let temp = tempdir()?; + let db_path = temp.path().join(storage_const::API_USER_DB_FILE); + let merge_guard_path = user_db_merge_guard_path(temp.path()); + + let mut v6_tree: BPlusTree = BPlusTree::new(); + v6_tree.insert( + "erin".to_string(), + StoredApiUserV6 { + target: "channels".to_string(), + username: "erin".to_string(), + password: "secret".to_string(), + token: None, + proxy: ProxyType::Reverse(None), + server: None, + epg_timeshift: None, + epg_request_timeshift: None, + created_at: None, + exp_date: None, + max_connections: Some(1), + status: Some(ProxyUserStatus::Active), + output_clusters: ClusterFlags::Live | ClusterFlags::Vod, + ui_enabled: true, + comment: None, + priority: Some(5), + soft_connections: Some(2), + soft_priority: Some(-4), + network_access: Some(NetworkAccessDto { + allowed_countries: Some(vec!["DE".to_string()]), + allowed_networks: Some(vec!["192.168.0.0/16".to_string()]), + }), + }, + ); + let _ = v6_tree.store(&db_path)?; + assert!(!merge_guard_path.exists()); + + let migrated = migrate_user_db_schema(&db_path, &merge_guard_path)?; + assert!(!migrated); + assert!(!merge_guard_path.exists()); + + let v6_tree = BPlusTree::::load(&db_path)?; + let user = v6_tree + .query(&"erin".to_string()) + .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "erin missing after v6 detection"))?; + assert_eq!(user.output_clusters, ClusterFlags::Live | ClusterFlags::Vod); + assert_eq!(user.priority, Some(5)); + assert_eq!(user.soft_connections, Some(2)); + assert_eq!(user.soft_priority, Some(-4)); + assert_eq!( + user.network_access.as_ref().and_then(|value| value.allowed_countries.as_ref()), + Some(&vec!["DE".to_string()]) + ); Ok(()) } diff --git a/backend/src/repository/strm_repository.rs b/backend/src/repository/strm_repository.rs index 53689e447..92fabc67f 100644 --- a/backend/src/repository/strm_repository.rs +++ b/backend/src/repository/strm_repository.rs @@ -1,24 +1,37 @@ -use crate::model::{ApiProxyServerInfo, AppConfig, ProxyUserCredentials}; -use crate::model::{ConfigTarget, StrmTargetFlags, StrmTargetOutput}; -use crate::repository::storage::ensure_target_storage_path; -use crate::repository::storage_const; -use crate::utils::{async_file_reader, async_file_writer, normalize_string_path, truncate_filename, - IO_BUFFER_SIZE}; +use crate::{ + model::{ + ApiProxyServerInfo, AppConfig, ConfigInput, ConfigTarget, ProxyUserCredentials, StrmTargetFlags, + StrmTargetOutput, + }, + repository::{storage::ensure_target_storage_path, storage_const}, + utils::{async_file_reader, async_file_writer, normalize_string_path, truncate_filename, IO_BUFFER_SIZE}, +}; use chrono::Datelike; use filetime::{set_file_times, FileTime}; use log::{error, trace}; use serde::Serialize; -use shared::error::TuliproxError; -use shared::model::{ClusterFlags, MediaQuality, PlaylistGroup, PlaylistItem, PlaylistItemType, StreamProperties, StrmExportStyle}; -use shared::utils::{arc_str_option_serde, arc_str_serde, clean_playlist_title, extract_extension_from_url, hash_bytes, - hash_string_as_hex, is_blank_optional_arc_str, truncate_string, ExportStyleConfig, CONSTANTS}; -use std::collections::{HashMap, HashSet, VecDeque}; -use std::path::{Path, PathBuf}; -use std::sync::Arc; -use tokio::fs::{create_dir_all, remove_dir, remove_file, File}; -use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt}; -use shared::model::UUIDType; -use std::borrow::Cow; +use shared::{ + error::TuliproxError, + model::{ + ClusterFlags, MediaQuality, PlaylistGroup, PlaylistItem, PlaylistItemType, StreamProperties, StrmExportStyle, + UUIDType, + }, + utils::{ + arc_str_option_serde, arc_str_serde, clean_playlist_title, extract_extension_from_url, hash_bytes, + hash_string_as_hex, is_blank_optional_arc_str, sanitize_sensitive_info, truncate_string, ExportStyleConfig, + CONSTANTS, PROVIDER_SCHEME_PREFIX, + }, +}; +use std::{ + borrow::Cow, + collections::{HashMap, HashSet, VecDeque}, + path::{Path, PathBuf}, + sync::Arc, +}; +use tokio::{ + fs::{create_dir_all, remove_dir, remove_file, File}, + io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt}, +}; /// Sanitizes a string to be safe for use as a file or directory name by /// following a strict "allow-list" approach and discarding invalid characters. @@ -36,7 +49,8 @@ fn sanitize_for_filename(text: &str, underscore_whitespace: bool) -> String { // Decide which characters to keep or transform. if c.is_alphanumeric() { Some(c) - } else if "+=,._-@#()[]".contains(c) { // <-- Allow list of safe punctuation, added [ and ] for quality tags. + } else if "+=,._-@#()[]".contains(c) { + // <-- Allow list of safe punctuation, added [ and ] for quality tags. Some(c) } else if c.is_whitespace() { if underscore_whitespace { @@ -82,7 +96,7 @@ fn style_rename_year<'a>( }); let cur_year = u32::try_from(chrono::Utc::now().year()).unwrap_or(0); - + // Check if we need to clean the title (remove year if present) or extract year from title // We iterate matches to either find the year (if meta_year is None) or remove it (if meta_year is Some) let mut new_name = String::with_capacity(name.len()); @@ -92,12 +106,13 @@ fn style_rename_year<'a>( for caps in style.year.captures_iter(name) { if let Some(year_match) = caps.get(1) { if let Ok(year) = year_match.as_str().parse::() { - if (1900..=cur_year + 5).contains(&year) { // Allow slightly future years + if (1900..=cur_year + 5).contains(&year) { + // Allow slightly future years // Found a valid year in title if extracted_year.is_none() { extracted_year = Some(year); } - + // We remove the year from the title in two cases: // A) We have a metadata year (clean up title to avoid "Movie (2000) (2000)") // B) We don't have metadata year (we extract it and remove it from title to re-append consistently later) @@ -111,21 +126,21 @@ fn style_rename_year<'a>( } } } - + new_name.push_str(&name[last_index..]); // Use metadata year if available, otherwise the one extracted from title let final_year = meta_year.or(extracted_year); - + // If we modified the string, trim it and return Owned if last_index > 0 { // Clean up potential double spaces or trailing punctuation left by removal // Remove trailing " -", ".", or "_" which might have been separators before the year let cleaned = new_name.trim().trim_end_matches(|c| " -_.".contains(c)).trim().to_string(); - + // Ensure we didn't make the name empty if cleaned.is_empty() { - return (Cow::Borrowed(name), final_year); + return (Cow::Borrowed(name), final_year); } (Cow::Owned(cleaned), final_year) } else { @@ -134,7 +149,11 @@ fn style_rename_year<'a>( } pub fn strm_get_file_paths(file_prefix: &str, target_path: &Path) -> PathBuf { - target_path.join(PathBuf::from(format!("{file_prefix}_{}.{}", storage_const::FILE_STRM, storage_const::FILE_SUFFIX_DB))) + target_path.join(PathBuf::from(format!( + "{file_prefix}_{}.{}", + storage_const::FILE_STRM, + storage_const::FILE_SUFFIX_DB + ))) } #[derive(Serialize)] @@ -169,9 +188,7 @@ struct StrmItemInfo { } impl StrmItemInfo { - pub(crate) fn get_file_ts(&self) -> Option { - self.added - } + pub(crate) fn get_file_ts(&self) -> Option { self.added } } fn extract_item_info(pli: &mut PlaylistItem) -> StrmItemInfo { @@ -183,61 +200,77 @@ fn extract_item_info(pli: &mut PlaylistItem) -> StrmItemInfo { let virtual_id = header.virtual_id; let input_name = header.input_name.clone(); let url = header.url.clone(); - + // Extract properties based on type // We prioritize name/title from additional_properties if available (e.g. from TMDB) - let (title, series_name, release_date, series_release_date, added, tmdb_id, season, episode) = match header.item_type { - PlaylistItemType::Series - | PlaylistItemType::LocalSeries => { - let (prop_name, release_date, series_release_date, added, tmdb_id, season, episode) = match header.additional_properties.as_ref() { - None => (None, None, None, None, None, None, None), - Some(props) => ( - // If series props are available, check if we have a valid name there - if let StreamProperties::Series(s) = props { - if s.name.is_empty() { None } else { Some(s.name.clone()) } - } else { None }, - - props.get_release_date(), - // Extract series-level release date from Episode properties - if let StreamProperties::Episode(ep) = props { ep.series_release_date.clone() } else { None }, - props.get_added(), - props.get_tmdb_id().filter(|&id| id != 0), - props.get_season(), - props.get_episode(), - ) - }; - + let (title, series_name, release_date, series_release_date, added, tmdb_id, season, episode) = match header + .item_type + { + PlaylistItemType::Series | PlaylistItemType::LocalSeries => { + let (prop_name, release_date, series_release_date, added, tmdb_id, season, episode) = + match header.additional_properties.as_ref() { + None => (None, None, None, None, None, None, None), + Some(props) => ( + // If series props are available, check if we have a valid name there + if let StreamProperties::Series(s) = props { + if s.name.is_empty() { + None + } else { + Some(s.name.clone()) + } + } else { + None + }, + props.get_release_date(), + // Extract series-level release date from Episode properties + if let StreamProperties::Episode(ep) = props { ep.series_release_date.clone() } else { None }, + props.get_added(), + props.get_tmdb_id().filter(|&id| id != 0), + props.get_season(), + props.get_episode(), + ), + }; + // For series title, we prefer the one from metadata (prop_name), then header.name, then header.title let final_series_name = prop_name.unwrap_or_else(|| { - if header.name.is_empty() { header.title.clone() } else { header.name.clone() } + if header.name.is_empty() { + header.title.clone() + } else { + header.name.clone() + } }); - + // Episode title relies on header.title unless we want to look deeper into props let ep_title = header.title.clone(); - + (ep_title, Some(final_series_name), release_date, series_release_date, added, tmdb_id, season, episode) } - PlaylistItemType::Video - | PlaylistItemType::LocalVideo => { + PlaylistItemType::Video | PlaylistItemType::LocalVideo => { let (prop_name, release_date, added, tmdb_id) = match header.additional_properties.as_ref() { None => (None, None, None, None), Some(props) => ( if let StreamProperties::Video(v) = props { - if v.name.is_empty() { None } else { Some(v.name.clone()) } - } else { None }, + if v.name.is_empty() { + None + } else { + Some(v.name.clone()) + } + } else { + None + }, props.get_release_date(), props.get_added(), props.get_tmdb_id().filter(|&id| id != 0), - ) + ), }; - + let final_title = prop_name.unwrap_or_else(|| header.title.clone()); - + (final_title, None, release_date, None, added, tmdb_id, None, None) } _ => (header.title.clone(), None, None, None, None, None, None, None), }; - + StrmItemInfo { group, title, @@ -294,9 +327,7 @@ async fn cleanup_strm_output_directory( processed: &HashSet, ) -> Result<(), String> { if !(root_path.exists() && root_path.is_dir()) { - return Err(format!( - "Error: STRM directory does not exist: {}", root_path.display() - )); + return Err(format!("Error: STRM directory does not exist: {}", root_path.display())); } let to_remove: HashSet = if cleanup { @@ -304,10 +335,7 @@ async fn cleanup_strm_output_directory( let mut found_files = HashSet::new(); let files = read_files_non_recursive(root_path).await.map_err(|err| err.to_string())?; for file_path in files { - if let Some(file_name) = file_path - .strip_prefix(root_path) - .ok() - .and_then(|p| p.to_str()) { + if let Some(file_name) = file_path.strip_prefix(root_path).ok().and_then(|p| p.to_str()) { found_files.insert(file_name.to_string()); } } @@ -331,16 +359,20 @@ async fn cleanup_strm_output_directory( fn filter_strm_item(pli: &PlaylistItem) -> bool { let item_type = pli.header.item_type; - matches!(item_type, PlaylistItemType::Live | PlaylistItemType::Video | PlaylistItemType::LocalVideo | PlaylistItemType::Series | PlaylistItemType::LocalSeries) + matches!( + item_type, + PlaylistItemType::Live + | PlaylistItemType::Video + | PlaylistItemType::LocalVideo + | PlaylistItemType::Series + | PlaylistItemType::LocalSeries + ) } fn get_relative_path_str(full_path: &Path, root_path: &Path) -> String { full_path .strip_prefix(root_path) - .map_or_else( - |_| full_path.to_string_lossy(), - |relative| relative.to_string_lossy(), - ) + .map_or_else(|_| full_path.to_string_lossy(), |relative| relative.to_string_lossy()) .to_string() } @@ -365,22 +397,15 @@ fn prepare_filename_parts( separator: &str, id_format: &str, // e.g. "{tmdb={}}" or "[tmdbid={}]" ) -> FilenameParts { - let id_string = if tmdb_id > 0 { - id_format.replace("{}", &tmdb_id.to_string()) - } else { - String::new() - }; - + let id_string = if tmdb_id > 0 { id_format.replace("{}", &tmdb_id.to_string()) } else { String::new() }; + // Determine source name and date based on type let (raw_name, raw_date) = match strm_item_info.item_type { PlaylistItemType::Series | PlaylistItemType::LocalSeries => ( - strm_item_info.series_name.as_ref().unwrap_or(&strm_item_info.title), - strm_item_info.series_release_date.as_ref() + strm_item_info.series_name.as_ref().unwrap_or(&strm_item_info.title), + strm_item_info.series_release_date.as_ref(), ), - _ => ( - &strm_item_info.title, - strm_item_info.release_date.as_ref() - ) + _ => (&strm_item_info.title, strm_item_info.release_date.as_ref()), }; // Use clean_playlist_title to remove IPTV garbage BEFORE parsing years @@ -422,8 +447,8 @@ fn format_for_kodi( if flat { if tmdb_id > 0 { - if let Some(path) = flat_dedup_paths.get(&tmdb_id) { - dir_path.clone_from(path); + if let Some(path) = flat_dedup_paths.get(&tmdb_id) { + dir_path.clone_from(path); } else { dir_path.push(&folder_name); flat_dedup_paths.insert(tmdb_id, dir_path.clone()); @@ -469,7 +494,7 @@ fn format_for_plex( flat: bool, flat_dedup_paths: &mut HashMap, ) -> (PathBuf, String) { - // Plex ID format: {tmdb-12345} + // Plex ID format: {tmdb-12345} let parts = prepare_filename_parts(strm_item_info, tmdb_id, separator, &format!("{separator}{{tmdb-{{}}}}")); let mut dir_path = PathBuf::new(); @@ -479,7 +504,7 @@ fn format_for_plex( let final_filename = parts.base_name; if flat { - if tmdb_id > 0 { + if tmdb_id > 0 { if let Some(path) = flat_dedup_paths.get(&tmdb_id) { dir_path.clone_from(path); } else { @@ -539,7 +564,7 @@ fn format_for_emby( let final_filename = format!("{}{}", parts.base_name, parts.id_string); if flat { - if tmdb_id > 0 { + if tmdb_id > 0 { if let Some(path) = flat_dedup_paths.get(&tmdb_id) { dir_path.clone_from(path); } else { @@ -556,7 +581,7 @@ fn format_for_emby( (dir_path, final_filename) } PlaylistItemType::Series | PlaylistItemType::LocalSeries => { - // For series, the ID goes in the folder name. + // For series, the ID goes in the folder name. let series_folder_name = format!("{}{}", parts.base_name, parts.id_string); let season_num = strm_item_info.season.unwrap_or(1); let episode_num = strm_item_info.episode.unwrap_or(1); @@ -589,7 +614,7 @@ fn format_for_jellyfin( flat: bool, flat_dedup_paths: &mut HashMap, ) -> (PathBuf, String) { - // Jellyfin ID format: [tmdbid-12345] + // Jellyfin ID format: [tmdbid-12345] let parts = prepare_filename_parts(strm_item_info, tmdb_id, separator, &format!("{separator}[tmdbid-{{}}]")); let mut dir_path = PathBuf::new(); @@ -600,7 +625,7 @@ fn format_for_jellyfin( let final_filename = folder_name.clone(); if flat { - if tmdb_id > 0 { + if tmdb_id > 0 { if let Some(path) = flat_dedup_paths.get(&tmdb_id) { dir_path.clone_from(path); } else { @@ -650,7 +675,6 @@ fn style_based_rename( ) -> (PathBuf, String) { let separator = if underscore_whitespace { "_" } else { " " }; - let tmdb_id = tmdb.or(strm_item_info.tmdb_id).unwrap_or(0); // Dispatch the call to the responsible function based on the style. @@ -662,14 +686,8 @@ fn style_based_rename( } } -fn prepare_strm_files( - new_playlist: &mut [PlaylistGroup], - strm_target_output: &StrmTargetOutput, -) -> Vec { - let channel_count = new_playlist - .iter() - .map(|g| g.filter_count(filter_strm_item)) - .sum(); +fn prepare_strm_files(new_playlist: &mut [PlaylistGroup], strm_target_output: &StrmTargetOutput) -> Vec { + let channel_count = new_playlist.iter().map(|g| g.filter_count(filter_strm_item)).sum(); // contains all paths (dir + filename) to detect collisions let mut all_filenames: HashSet = HashSet::with_capacity(channel_count); // contains only collision filenames (PathBuf) @@ -681,31 +699,22 @@ fn prepare_strm_files( // first we create the names to identify name collisions for pg in new_playlist.iter_mut() { for pli in pg.channels.iter_mut().filter(|c| filter_strm_item(c)) { - let strm_item_info = extract_item_info(pli); let (dir_path, strm_file_name) = style_based_rename( &strm_item_info, pli.get_tmdb_id(), strm_target_output.style, - strm_target_output - .flags - .contains(StrmTargetFlags::UnderscoreWhitespace), + strm_target_output.flags.contains(StrmTargetFlags::UnderscoreWhitespace), strm_target_output.flags.contains(StrmTargetFlags::Flat), &mut flat_dedup_paths, ); // Conditionally generate the quality string based on the new config flag - let separator = if strm_target_output - .flags - .contains(StrmTargetFlags::UnderscoreWhitespace) - { - "_" - } else { - " " - }; + let separator = + if strm_target_output.flags.contains(StrmTargetFlags::UnderscoreWhitespace) { "_" } else { " " }; let quality_string = get_quality(strm_target_output, pli, separator); - + // Add category suffix for flat movie structure to avoid collisions let category_suffix = if strm_target_output.flags.contains(StrmTargetFlags::Flat) && pli.get_tmdb_id().is_some() @@ -713,13 +722,13 @@ fn prepare_strm_files( { let cat = sanitize_for_filename(&strm_item_info.group, false); format!("{separator}[{cat}]") - } else { - String::new() + } else { + String::new() }; let final_filename = format!("{strm_file_name}{quality_string}{category_suffix}"); let filename = Arc::new(final_filename); - + // Construct the full relative path for collision checking let full_relative_path = dir_path.join(filename.as_str()); @@ -727,25 +736,15 @@ fn prepare_strm_files( collisions.insert(full_relative_path.clone()); } all_filenames.insert(full_relative_path); - result.push(StrmFile { - file_name: filename, - dir_path, - strm_info: strm_item_info, - }); + result.push(StrmFile { file_name: filename, dir_path, strm_info: strm_item_info }); } } if !collisions.is_empty() { // This separator is specifically for the multi-version naming convention. let version_separator = " "; - let separator = if strm_target_output - .flags - .contains(StrmTargetFlags::UnderscoreWhitespace) - { - "_" - } else { - " " - }; + let separator = + if strm_target_output.flags.contains(StrmTargetFlags::UnderscoreWhitespace) { "_" } else { " " }; for s in &mut result { { @@ -769,29 +768,23 @@ fn prepare_strm_files( } fn get_quality(strm_target_output: &StrmTargetOutput, pli: &PlaylistItem, separator: &str) -> String { - if strm_target_output - .flags - .contains(StrmTargetFlags::AddQualityToFilename) - { + if strm_target_output.flags.contains(StrmTargetFlags::AddQualityToFilename) { // Use `additional_properties` which are populated by metadata_update_manager/probe let (audio, video) = match pli.header.additional_properties.as_ref() { None => (None, None), - Some(props) => { - match props { - StreamProperties::Live(_) - | StreamProperties::Series(_) => (None, None), - StreamProperties::Video(video) => - video.details.as_ref().map_or_else(|| (None, None), |d| (d.audio.as_deref(), d.video.as_deref())), - StreamProperties::Episode(episode) => - (episode.audio.as_deref(), episode.video.as_deref()) + Some(props) => match props { + StreamProperties::Live(_) | StreamProperties::Series(_) => (None, None), + StreamProperties::Video(video) => { + video.details.as_ref().map_or_else(|| (None, None), |d| (d.audio.as_deref(), d.video.as_deref())) } - } + StreamProperties::Episode(episode) => (episode.audio.as_deref(), episode.video.as_deref()), + }, }; if let Some(media_quality) = MediaQuality::from_ffprobe_info(audio, video) { let formatted = media_quality.format_for_filename(separator); if !formatted.is_empty() { // Hard-coded separator for filename clarity. - return format!(" - [{formatted}]") + return format!(" - [{formatted}]"); } } } @@ -814,16 +807,9 @@ fn strm_contains_tmdb_marker(s: &str) -> bool { /// returning the equivalent no-tmdb path used for identity matching. fn strip_tmdb_markers(s: &str) -> String { let mut result = s.to_string(); - for marker_prefix in &[ - " {tmdb=", - " {tmdb-", - " [tmdbid=", - " [tmdbid-", - "_{tmdb=", - "_{tmdb-", - "_[tmdbid=", - "_[tmdbid-", - ] { + for marker_prefix in + &[" {tmdb=", " {tmdb-", " [tmdbid=", " [tmdbid-", "_{tmdb=", "_{tmdb-", "_[tmdbid=", "_[tmdbid-"] + { while let Some(start) = result.find(marker_prefix) { let close_char = if marker_prefix.contains('{') { '}' } else { ']' }; let search_from = start + marker_prefix.len(); @@ -849,11 +835,10 @@ pub async fn write_strm_playlist( } let config = app_config.config.load(); - let Some(root_path) = crate::utils::get_file_path( - &config.storage_dir, - Some(std::path::PathBuf::from(&target_output.directory)), - ) else { - return Err(TuliproxError::Config(format!("Failed to get file path for {}",target_output.directory))); + let Some(root_path) = + crate::utils::get_file_path(&config.storage_dir, Some(std::path::PathBuf::from(&target_output.directory))) + else { + return Err(TuliproxError::Config(format!("Failed to get file path for {}", target_output.directory))); }; let user_and_server_info = get_credentials_and_server_info(app_config, target_output.username.as_deref()); @@ -862,12 +847,8 @@ pub async fn write_strm_playlist( let strm_index_path = strm_get_file_paths(&strm_file_prefix, &ensure_target_storage_path(&config, target.name.as_str()).await?); let existing_strm = { - let _file_lock = app_config - .file_locks - .read_lock(&strm_index_path).await; - read_strm_file_index(&strm_index_path) - .await - .unwrap_or_else(|_| HashSet::with_capacity(4096)) + let _file_lock = app_config.file_locks.read_lock(&strm_index_path).await; + read_strm_file_index(&strm_index_path).await.unwrap_or_else(|_| HashSet::with_capacity(4096)) }; let mut processed_strm: HashSet = HashSet::with_capacity(existing_strm.len()); @@ -885,12 +866,10 @@ pub async fn write_strm_playlist( prepare_strm_output_directory(&root_path).await?; let target_force_redirect = target.options.as_ref().and_then(|o| o.force_redirect.as_ref()); + let mut input_by_name: HashMap, Option>> = HashMap::new(); + + let strm_files = prepare_strm_files(new_playlist, target_output); - let strm_files = prepare_strm_files( - new_playlist, - target_output, - ); - for strm_file in strm_files { // file paths let output_path = truncate_filename(&root_path.join(&strm_file.dir_path), 255); @@ -898,24 +877,30 @@ pub async fn write_strm_playlist( let relative_file_path = get_relative_path_str(&file_path, &root_path); - let (target_relative_file_path, target_file_path) = if strm_file.strm_info.tmdb_id.is_none() { - if let Some(enriched_path) = enriched_strm.get(&relative_file_path) { - (enriched_path.clone(), root_path.join(enriched_path)) - } else { - (relative_file_path.clone(), file_path) - } - } else { - (relative_file_path.clone(), file_path) - }; + let (target_relative_file_path, target_file_path) = get_target_strm_file_path( + &root_path, + &enriched_strm, + &relative_file_path, + file_path, + strm_file.strm_info.tmdb_id, + ); let file_exists = target_file_path.exists(); // create content - let url = get_strm_url(target_force_redirect, user_and_server_info.as_ref(), &strm_file.strm_info); - let mut content = target_output.strm_props.as_ref().map_or_else(Vec::new, std::clone::Clone::clone); - content.push(url.to_string()); - let content_text = content.join("\r\n"); - let content_as_bytes = content_text.as_bytes(); - let content_hash = hash_bytes(content_as_bytes); + let url = match resolve_strm_file_url( + app_config, + &mut input_by_name, + target_force_redirect, + user_and_server_info.as_ref(), + &strm_file.strm_info, + ) { + Ok(url) => url, + Err(err) => { + failed.push(err); + continue; + } + }; + let (content_as_bytes, content_hash) = build_strm_content(target_output, &url); // check if file exists and has same hash if file_exists && has_strm_file_same_hash(&target_file_path, content_hash).await { @@ -923,39 +908,31 @@ pub async fn write_strm_playlist( continue; // skip creation } - // if we can't create the directory skip this entry - let target_output_path = target_file_path.parent().map_or_else(|| output_path.clone(), std::path::Path::to_path_buf); - if !ensure_strm_file_directory(&mut failed, &target_output_path).await { + if !write_strm_output_file( + &mut failed, + &target_file_path, + &output_path, + &content_as_bytes, + strm_file.strm_info.get_file_ts(), + ) + .await + { continue; } - - match write_strm_file( - &target_file_path, - content_as_bytes, - strm_file.strm_info.get_file_ts(), - ).await - { - Ok(()) => { - processed_strm.insert(target_relative_file_path); - } - Err(err) => { - failed.push(err); - } - }; + processed_strm.insert(target_relative_file_path); } if let Err(err) = write_strm_index_file(app_config, &processed_strm, &strm_index_path).await { failed.push(err); } - if let Err(err) = - cleanup_strm_output_directory( - target_output.flags.contains(StrmTargetFlags::Cleanup), - &root_path, - &existing_strm, - &processed_strm, - ) - .await + if let Err(err) = cleanup_strm_output_directory( + target_output.flags.contains(StrmTargetFlags::Cleanup), + &root_path, + &existing_strm, + &processed_strm, + ) + .await { failed.push(err); } @@ -971,9 +948,7 @@ async fn write_strm_index_file( entries: &HashSet, index_file_path: &PathBuf, ) -> Result<(), String> { - let _file_lock = cfg - .file_locks - .write_lock(index_file_path).await; + let _file_lock = cfg.file_locks.write_lock(index_file_path).await; let file = File::create(index_file_path) .await .map_err(|err| format!("Failed to create strm index file: {} {err}", index_file_path.display()))?; @@ -984,35 +959,22 @@ async fn write_strm_index_file( for entry in entries { let bytes = entry.as_bytes(); write_counter += bytes.len() + 1; - writer - .write_all(bytes) - .await - .map_err(|err| format!("Failed to write strm index entry: {err}"))?; - writer - .write_all(new_line) - .await - .map_err(|err| format!("Failed to write strm index entry: {err}"))?; + writer.write_all(bytes).await.map_err(|err| format!("Failed to write strm index entry: {err}"))?; + writer.write_all(new_line).await.map_err(|err| format!("Failed to write strm index entry: {err}"))?; if write_counter >= IO_BUFFER_SIZE { write_counter = 0; writer.flush().await.map_err(|err| format!("Failed to flush: {err}"))?; } } - writer - .flush() - .await - .map_err(|err| format!("failed to write strm index entry: {err}"))?; - writer - .shutdown() - .await - .map_err(|err| format!("failed to write strm index entry: {err}"))?; + writer.flush().await.map_err(|err| format!("failed to write strm index entry: {err}"))?; + writer.shutdown().await.map_err(|err| format!("failed to write strm index entry: {err}"))?; Ok(()) } async fn ensure_strm_file_directory(failed: &mut Vec, output_path: &Path) -> bool { if !output_path.exists() { if let Err(e) = create_dir_all(output_path).await { - let err_msg = - format!("Failed to create directory for strm playlist: {} {e}", output_path.display()); + let err_msg = format!("Failed to create directory for strm playlist: {} {e}", output_path.display()); error!("{err_msg}"); failed.push(err_msg); return false; // skip creation, could not create directory @@ -1021,11 +983,29 @@ async fn ensure_strm_file_directory(failed: &mut Vec, output_path: &Path true } -async fn write_strm_file( - file_path: &Path, +async fn write_strm_output_file( + failed: &mut Vec, + target_file_path: &Path, + output_path: &Path, content_as_bytes: &[u8], timestamp: Option, -) -> Result<(), String> { +) -> bool { + let target_output_path = + target_file_path.parent().map_or_else(|| output_path.to_path_buf(), std::path::Path::to_path_buf); + if !ensure_strm_file_directory(failed, &target_output_path).await { + return false; + } + + match write_strm_file(target_file_path, content_as_bytes, timestamp).await { + Ok(()) => true, + Err(err) => { + failed.push(err); + false + } + } +} + +async fn write_strm_file(file_path: &Path, content_as_bytes: &[u8], timestamp: Option) -> Result<(), String> { File::create(file_path) .await .map_err(|err| format!("Failed to create strm file: {err}"))? @@ -1066,10 +1046,10 @@ async fn has_strm_file_same_hash(file_path: &PathBuf, content_hash: UUIDType) -> fn get_credentials_and_server_info( cfg: &AppConfig, username: Option<&str>, -) -> Option<(ProxyUserCredentials, ApiProxyServerInfo)> { +) -> Option<(Arc, ApiProxyServerInfo)> { let username = username?; let credentials = cfg.get_user_credentials(username)?; - let server_info = cfg.get_user_server_info(&credentials)?; + let server_info = cfg.get_user_server_info(credentials.as_ref())?; Some((credentials, server_info)) } @@ -1084,16 +1064,85 @@ async fn read_strm_file_index(strm_file_index_path: &Path) -> std::io::Result, Option>>, + target_force_redirect: Option<&ClusterFlags>, + user_and_server_info: Option<&(Arc, ApiProxyServerInfo)>, + str_item_info: &StrmItemInfo, +) -> Result, String> { + let input = if user_and_server_info.is_none() && str_item_info.url.starts_with(PROVIDER_SCHEME_PREFIX) { + input_by_name + .entry(Arc::clone(&str_item_info.input_name)) + .or_insert_with(|| app_config.get_input_by_name(&str_item_info.input_name)) + .clone() + } else { + None + }; + + get_strm_url(target_force_redirect, user_and_server_info, input.as_deref(), str_item_info) +} + +fn get_target_strm_file_path( + root_path: &Path, + enriched_strm: &HashMap, + relative_file_path: &str, + file_path: PathBuf, + tmdb_id: Option, +) -> (String, PathBuf) { + if tmdb_id.is_none() { + if let Some(enriched_path) = enriched_strm.get(relative_file_path) { + return (enriched_path.clone(), root_path.join(enriched_path)); + } + } + (relative_file_path.to_string(), file_path) +} + +fn build_strm_content(target_output: &StrmTargetOutput, url: &str) -> (Vec, UUIDType) { + let mut content = target_output.strm_props.as_ref().map_or_else(Vec::new, std::clone::Clone::clone); + content.push(url.to_string()); + let content_text = content.join("\r\n"); + let content_as_bytes = content_text.into_bytes(); + let content_hash = hash_bytes(&content_as_bytes); + (content_as_bytes, content_hash) +} + +fn resolve_strm_source_url(input: Option<&ConfigInput>, str_item_info: &StrmItemInfo) -> Result, String> { + if !str_item_info.url.starts_with(PROVIDER_SCHEME_PREFIX) { + return Ok(Arc::clone(&str_item_info.url)); + } + + let input = input.ok_or_else(|| { + format!( + "Failed to resolve STRM provider URL for input '{}' because the source input is missing: {}", + sanitize_sensitive_info(&str_item_info.input_name), + sanitize_sensitive_info(&str_item_info.url) + ) + })?; + + input.resolve_url(&str_item_info.url).map(|resolved| Arc::::from(resolved.into_owned())).map_err(|err| { + format!( + "Failed to resolve STRM provider URL for input '{}': {} ({err})", + sanitize_sensitive_info(&str_item_info.input_name), + sanitize_sensitive_info(&str_item_info.url) + ) + }) +} + fn get_strm_url( target_force_redirect: Option<&ClusterFlags>, - user_and_server_info: Option<&(ProxyUserCredentials, ApiProxyServerInfo)>, + user_and_server_info: Option<&(Arc, ApiProxyServerInfo)>, + input: Option<&ConfigInput>, str_item_info: &StrmItemInfo, -) -> Arc { - let Some((user, server_info)) = user_and_server_info else { return str_item_info.url.clone(); }; +) -> Result, String> { + let Some((user, server_info)) = user_and_server_info else { + return resolve_strm_source_url(input, str_item_info); + }; - let redirect = user.proxy.is_redirect(str_item_info.item_type) || target_force_redirect.is_some_and(|f| f.has_cluster(str_item_info.item_type)); + let redirect = user.proxy.is_redirect(str_item_info.item_type) + || target_force_redirect.is_some_and(|f| f.has_cluster(str_item_info.item_type)); if redirect { - return str_item_info.url.clone(); + return Ok(Arc::clone(&str_item_info.url)); } if let Some(stream_type) = match str_item_info.item_type { @@ -1102,21 +1151,21 @@ fn get_strm_url( | PlaylistItemType::SeriesInfo | PlaylistItemType::LocalSeries | PlaylistItemType::LocalSeriesInfo => Some("series"), - PlaylistItemType::Video - | PlaylistItemType::LocalVideo => Some("movie"), + PlaylistItemType::Video | PlaylistItemType::LocalVideo => Some("movie"), _ => None, } { let url = &str_item_info.url; let ext = extract_extension_from_url(url).unwrap_or_default(); - format!( + Ok(format!( "{}/{stream_type}/{}/{}/{}{ext}", server_info.get_base_url(), - user.username, - user.password, + user.as_ref().username, + user.as_ref().password, str_item_info.virtual_id - ).into() + ) + .into()) } else { - str_item_info.url.clone() + Ok(Arc::clone(&str_item_info.url)) } } @@ -1128,29 +1177,19 @@ fn get_strm_url( #[derive(Debug, Clone)] struct DirNode { path: PathBuf, - is_root: bool, // is root -> do not delete! + is_root: bool, // is root -> do not delete! has_files: bool, // has content -> do not delete! children: HashSet, parent: Option, } impl DirNode { - fn new(path: PathBuf, parent: Option) -> Self { - Self::new_with_flag(path, parent, false) - } + fn new(path: PathBuf, parent: Option) -> Self { Self::new_with_flag(path, parent, false) } - fn new_root(path: PathBuf) -> Self { - Self::new_with_flag(path, None, true) - } + fn new_root(path: PathBuf) -> Self { Self::new_with_flag(path, None, true) } fn new_with_flag(path: PathBuf, parent: Option, is_root: bool) -> Self { - Self { - path, - is_root, - has_files: false, - children: HashSet::new(), - parent, - } + Self { path, is_root, has_files: false, children: HashSet::new(), parent } } } @@ -1200,10 +1239,7 @@ async fn build_directory_tree(root_path: &Path) -> HashMap { // now we need to build an ordered flat list, // We walk from top to bottom. // (PS: you can only delete in reverse order, because delete first children, then the parents) -fn flatten_tree( - root_path: &Path, - mut tree_nodes: HashMap, -) -> Vec { +fn flatten_tree(root_path: &Path, mut tree_nodes: HashMap) -> Vec { let mut paths_to_process = Vec::new(); // List of paths to process { @@ -1222,10 +1258,7 @@ fn flatten_tree( } } - paths_to_process - .iter() - .filter_map(|path| tree_nodes.remove(path)) - .collect() + paths_to_process.iter().filter_map(|path| tree_nodes.remove(path)).collect() } async fn delete_empty_dirs_from_tree(root_path: &Path, tree_nodes: HashMap) { @@ -1243,3 +1276,67 @@ async fn remove_empty_dirs(root_path: PathBuf) { let tree_nodes = build_directory_tree(&root_path).await; delete_empty_dirs_from_tree(&root_path, tree_nodes).await; } + +#[cfg(test)] +mod tests { + use super::{resolve_strm_source_url, StrmItemInfo}; + use crate::model::{ConfigInput, ConfigProvider}; + use shared::model::{ConfigProviderDto, InputType, PlaylistItemType, ProviderUrlSelectionPolicy}; + use std::{collections::HashMap, sync::Arc}; + + fn make_strm_item(url: &str, input_name: &str) -> StrmItemInfo { + StrmItemInfo { + group: Arc::from("group"), + title: Arc::from("title"), + item_type: PlaylistItemType::Live, + provider_id: None, + virtual_id: 1, + input_name: Arc::from(input_name), + url: Arc::from(url), + series_name: None, + release_date: None, + series_release_date: None, + season: None, + episode: None, + added: None, + tmdb_id: None, + } + } + + fn make_input_with_provider() -> ConfigInput { + let provider = ConfigProvider::from(&ConfigProviderDto { + name: "myprovider".into(), + urls: vec!["http://provider.example.com".into()], + provider_url_selection_policy: ProviderUrlSelectionPolicy::ResumeLastWorking, + dns: None, + }); + + ConfigInput { + name: Arc::from("input-a"), + input_type: InputType::M3u, + headers: HashMap::new(), + url: "http://input.example.com".to_string(), + provider_configs: Some(vec![Arc::new(provider)]), + ..Default::default() + } + } + + #[test] + fn resolve_strm_source_url_resolves_provider_scheme_without_user_context() { + let input = make_input_with_provider(); + let strm_item = make_strm_item("provider://myprovider/live/1.ts", "input-a"); + + let resolved = resolve_strm_source_url(Some(&input), &strm_item); + + assert_eq!(resolved.as_deref(), Ok("http://provider.example.com/live/1.ts")); + } + + #[test] + fn resolve_strm_source_url_fails_when_provider_input_is_missing() { + let strm_item = make_strm_item("provider://myprovider/live/1.ts", "missing-input"); + + let err = resolve_strm_source_url(None, &strm_item).err(); + + assert!(err.is_some_and(|message| message.contains("source input is missing"))); + } +} diff --git a/backend/src/repository/user_repository.rs b/backend/src/repository/user_repository.rs index 3113f4bba..f2e4ee263 100644 --- a/backend/src/repository/user_repository.rs +++ b/backend/src/repository/user_repository.rs @@ -1,6 +1,6 @@ -use crate::model::Config; -use crate::model::PlaylistXtreamCategory; +use crate::model::{Config, NetworkAccess, PlaylistXtreamCategory}; use crate::model::{AppConfig, ProxyUserCredentials, TargetUser}; +use std::sync::Arc; use crate::repository::storage_const; use crate::repository::xtream_get_playlist_categories; use crate::repository::BPlusTree; @@ -16,7 +16,7 @@ use std::io::Error; use std::path::{Path, PathBuf}; use tokio::task; -// V5 (current): added output_clusters. V1-V4 are migrated to V5 at startup +// V6 (current): added network_access. V1-V5 are migrated to V6 at startup // by `bplustree_migration::run_all_startup_migrations`. #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] struct StoredProxyUserCredentials { @@ -38,6 +38,8 @@ struct StoredProxyUserCredentials { pub priority: Option, pub soft_connections: Option, pub soft_priority: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub network_access: Option, } impl StoredProxyUserCredentials { @@ -61,6 +63,7 @@ impl StoredProxyUserCredentials { priority: if proxy.priority != 0 { Some(proxy.priority) } else { None }, soft_connections: if proxy.soft_connections > 0 { Some(proxy.soft_connections) } else { None }, soft_priority: if proxy.soft_priority != 0 { Some(proxy.soft_priority) } else { None }, + network_access: proxy.network_access.as_ref().map(Into::into), } } @@ -84,6 +87,7 @@ impl StoredProxyUserCredentials { soft_connections: stored.soft_connections.unwrap_or(0), soft_priority: stored.soft_priority.unwrap_or(0), t_is_api_user: false, + network_access: stored.network_access.as_ref().map(NetworkAccess::from), } } } @@ -167,10 +171,10 @@ fn collect_target_users(user_tree: &BPlusTree { - entry.get_mut().credentials.push(proxy_user); + entry.get_mut().credentials.push(Arc::new(proxy_user)); } std::collections::hash_map::Entry::Vacant(entry) => { - entry.insert(TargetUser { target: stored_user.target.clone(), credentials: vec![proxy_user] }); + entry.insert(TargetUser { target: stored_user.target.clone(), credentials: vec![Arc::new(proxy_user)] }); } } } @@ -505,7 +509,7 @@ mod tests { let user = TargetUser { target: "test".to_string(), credentials: vec![ - ProxyUserCredentials { + Arc::new(ProxyUserCredentials { username: "Test".to_string(), password: "Test".to_string(), token: Some("Test".to_string()), @@ -524,8 +528,9 @@ mod tests { soft_connections: 0, soft_priority: 0, t_is_api_user: false, - }, - ProxyUserCredentials { + network_access: None, + }), + Arc::new(ProxyUserCredentials { username: "Test2".to_string(), password: "Test".to_string(), token: Some("Test".to_string()), @@ -544,8 +549,9 @@ mod tests { soft_connections: 0, soft_priority: 0, t_is_api_user: false, - }, - ProxyUserCredentials { + network_access: None, + }), + Arc::new(ProxyUserCredentials { username: "Test3".to_string(), password: "Test".to_string(), token: Some("Test".to_string()), @@ -564,8 +570,9 @@ mod tests { soft_connections: 0, soft_priority: 0, t_is_api_user: false, - }, - ProxyUserCredentials { + network_access: None, + }), + Arc::new(ProxyUserCredentials { username: "Test4".to_string(), password: "Test".to_string(), token: Some("Test".to_string()), @@ -584,7 +591,11 @@ mod tests { soft_connections: 2, soft_priority: -3, t_is_api_user: false, - }, + network_access: Some(crate::model::NetworkAccess { + allowed_countries: vec!["DE".to_string(), "AT".to_string()], + allowed_networks: vec!["10.0.0.0/8".parse().unwrap(), "192.168.1.0/24".parse().unwrap()], + }), + }), ], }; @@ -626,6 +637,15 @@ mod tests { assert_eq!(test4.soft_connections, 2); assert_eq!(test4.soft_priority, -3); assert_eq!(test4.output_clusters, ClusterFlags::Live | ClusterFlags::Vod); + // Verify network_access round-trip (allowed_countries + allowed_networks). + let test4_na = test4.network_access.as_ref().expect("Test4 network_access should be set"); + let mut countries = test4_na.allowed_countries.clone(); + countries.sort(); + assert_eq!(countries, vec!["AT", "DE"]); + let networks: Vec = test4_na.allowed_networks.iter().map(std::string::ToString::to_string).collect(); + assert_eq!(networks.len(), 2); + assert!(networks.iter().any(|n| n == "10.0.0.0/8")); + assert!(networks.iter().any(|n| n == "192.168.1.0/24")); } #[tokio::test] diff --git a/backend/src/utils/geoip.rs b/backend/src/utils/geoip.rs index 40f2bcdc3..411683002 100644 --- a/backend/src/utils/geoip.rs +++ b/backend/src/utils/geoip.rs @@ -12,7 +12,6 @@ pub struct GeoIp { tree: BPlusTree, } - impl GeoIp { fn seed_private_ranges(tree: &mut BPlusTree) { @@ -91,14 +90,26 @@ impl Default for GeoIp { } #[cfg(test)] -mod test { - // https://raw.githubusercontent.com/sapics/ip-location-db/refs/heads/main/asn-country/asn-country-ipv4.csv +impl GeoIp { + /// Creates a `GeoIp` instance for testing that returns the specified country for any IPv4 address. + #[cfg(test)] + pub fn test_new(country: &str) -> Self { + let mut tree = BPlusTree::new(); + tree.insert(0, (u32::MAX, country.to_string())); + Self { tree } + } +} - // use crate::utils::geoip::GeoIp; + +// #[cfg(test)] +// mod test { // use std::fs::File; // use std::path::PathBuf; // use crate::utils::file_reader; + // https://raw.githubusercontent.com/sapics/ip-location-db/refs/heads/main/asn-country/asn-country-ipv4.csv + + // #[test] // pub fn test_csv() { // let db_file = PathBuf::from("/projects/m3u-test/asn-country-ipv4.db"); @@ -115,4 +126,4 @@ mod test { // panic!("GeoIP lookup returned no result"); // } // } -} +// } diff --git a/backend/src/utils/mod.rs b/backend/src/utils/mod.rs index a265c53ee..26994b29c 100644 --- a/backend/src/utils/mod.rs +++ b/backend/src/utils/mod.rs @@ -9,7 +9,7 @@ mod trakt; mod json_utils; mod binary_utils; mod telegram; -mod geoip; +pub(crate) mod geoip; mod db_viewer; pub(crate) mod stream_history_viewer; mod epg_parser; diff --git a/config/api-proxy.yml b/config/api-proxy.yml index 506548dd8..2fdf2d02d 100644 --- a/config/api-proxy.yml +++ b/config/api-proxy.yml @@ -24,7 +24,46 @@ user: max_connections: 0 status: Active ui_enabled: true + # network_access: + # allowed_networks: + # - "192.168.0.0/16" + # - "10.0.0.0/8" + - username: vpn-only + password: vpnsecret + proxy: reverse + output_clusters: [live, vod, series] + server: external + max_connections: 2 + status: Active + # Network access restriction — uses OR logic: matching ANY allowed_networks + # OR ANY allowed_countries is sufficient for access. + # + # Private/VPN ranges must use CIDR notation (/16, /24, /32 for single IPv4s, + # /128 for a single IPv6, e.g. 2001:db8::/32). + # Country-based restrictions require GeoIP database to be configured. + # If GeoIP is unavailable and country restrictions exist, access is denied by default. + # To explicitly accept that risk globally, set this in config.yml: + # + # reverse_proxy: + # geoip: + # unavailable_policy: allow + # + # The policy is global, not per user. CIDR-only misses, unknown countries, + # and country mismatches still deny. + # + # Client IP is derived from X-Real-IP / X-Forwarded-For headers when behind + # a reverse proxy. Ensure your reverse proxy is configured to set these, + # otherwise network restrictions may not function correctly. + network_access: + allowed_networks: + - "10.200.0.0/16" # WireGuard VPN range + allowed_countries: + - DE # Germany + - AT # Austria + ui_enabled: false + + # No network restrictions — allow from any source - username: external password: externalsecret token: "77418" diff --git a/config/config.yml b/config/config.yml index 5e20952d0..7e4a36a45 100644 --- a/config/config.yml +++ b/config/config.yml @@ -83,6 +83,13 @@ reverse_proxy: max_attempts: 3 backoff_millis: 250 backoff_multiplier: 1.0 + # geoip: + # enabled: true + # # unavailable_policy: deny # default. If GeoIP is disabled, missing, or not loaded, + # # country-based network_access rules deny non-CIDR-matching requests. + # # allow: explicit risk acceptance. Country-based network_access rules + # # allow when GeoIP is unavailable; CIDR-only misses still deny. + # unavailable_policy: deny stream: shared_burst_buffer_mb: 12 # default 12 MB # stream_history: diff --git a/docs/src/configuration/api-proxy.md b/docs/src/configuration/api-proxy.md index cf2208611..56a16bbaa 100644 --- a/docs/src/configuration/api-proxy.md +++ b/docs/src/configuration/api-proxy.md @@ -127,6 +127,56 @@ in your `config.yml`. Without it, these fields are purely cosmetic! | `exp_date` | UnixTs | No | `None` | Locks the user out after this Unix timestamp. **Requires** `user_access_control: true` in `config.yml` to be enforced. | | `ui_enabled` | Bool | No | `true` | Allows this specific user to log into the Web UI to manage their own favorites/bouquets. | | `priority` | Int (i8) | No | `0` | Stream preemption priority. Priority range: `-128` to `127`, where `-128` has the highest priority. Negative numbers are explicitly allowed for top-tier access. (see [user priority](#user-priorities-priority) below) | +| `network_access` | Block | No | `None` | Per-user network/country access restrictions. Uses OR logic — matching ANY `allowed_networks` (CIDR) OR ANY `allowed_countries` grants access. Requires GeoIP for country checks. Client IP from `X-Real-IP` / `X-Forwarded-For`. See [Network Access Restrictions](#network-access-restrictions) below. | + +--- + +### Network Access Restrictions + +Tuliprox supports per-user network access restrictions to limit streaming based on the client's IP address or geographic location. + +**How it works:** + +- Uses OR logic — matching ANY `allowed_networks` (CIDR range) OR ANY `allowed_countries` grants access +- `allowed_networks`: CIDR notation for IPv4 and IPv6 + (e.g., `192.168.0.0/16`, `10.0.0.0/8`, `192.168.1.1/32` for a single IPv4, + `2001:db8::/32` or `2001:db8::1/128` for IPv6) +- `allowed_countries`: ISO 3166-1 alpha-2 country codes (e.g., `DE`, `US`). Requires GeoIP database +- If no restrictions are configured, all IPs are allowed +- If GeoIP is unavailable and country restrictions exist, access is denied by default +- The global `reverse_proxy.geoip.unavailable_policy` setting can explicitly change only the GeoIP-unavailable case: + - `deny` (default): country-based restrictions deny when GeoIP is disabled, missing, or not loaded + - `allow`: explicit risk acceptance; country-based restrictions allow when GeoIP is unavailable +- CIDR-only misses, unknown countries, and country mismatches still deny + +**Reverse proxy dependency:** Client IP is derived from `X-Real-IP` or `X-Forwarded-For` headers. If your reverse proxy +doesn't set these, network restrictions will not work correctly. + +**Example:** + +```yaml +network_access: + allowed_networks: + - "10.200.0.0/16" # VPN range + - "192.168.1.1/32" # Single IP + - "2001:db8::/32" # IPv6 range + allowed_countries: + - DE # Germany + - AT # Austria +``` + +Global GeoIP-unavailable policy example: + +```yaml +reverse_proxy: + geoip: + enabled: true + unavailable_policy: deny # default; use allow only as explicit risk acceptance +``` + +**Operator logging:** Denied requests emit structured logs: +`Network access denied: user="john" client_ip="203.0.113.5" reason=no_country_match`. Possible reasons: +`no_cidr_match`, `no_country_match`, `geoip_unavailable`, `country_unknown`, `malformed_client_ip`. --- diff --git a/frontend/public/assets/i18n/en.json b/frontend/public/assets/i18n/en.json index 03906b64e..9fe62f885 100644 --- a/frontend/public/assets/i18n/en.json +++ b/frontend/public/assets/i18n/en.json @@ -403,7 +403,8 @@ }, "GEO_IP_CONFIG": { "ENABLED": "If enabled, Geo-IP lookup is performed for incoming requests.", - "URL": "URL to download the Geo-IP database (e.g. MaxMind GeoLite2-City.mmdb)." + "URL": "URL to download the Geo-IP database (e.g. MaxMind GeoLite2-City.mmdb).", + "UNAVAILABLE_POLICY": "`deny` is the secure default. If Geo-IP is disabled, missing, or not loaded, country-based network access rules deny requests that did not match a configured CIDR.\n\n`allow` is an explicit risk acceptance. It allows country-based network access rules when Geo-IP is unavailable. CIDR-only misses, unknown countries, and country mismatches still deny." }, "HD_HOME_RUN_CONFIG": { "AUTH": "Authentication string or token required for HDHomeRun API access.", @@ -582,6 +583,8 @@ "SOFT_CONNECTIONS": "Additional provider slots above max_connections. Soft connections can be preempted by any normal connection or by a higher-priority soft connection.", "SOFT_PRIORITY": "Priority used while this user's connection is consuming a soft slot. Once promoted back to a normal slot, the regular priority applies again.", "PASSWORD": "Access password for this user's playlist and streams.", + "NETWORK_ACCESS_COUNTRIES": "Add one ISO 3166-1 alpha-2 country code per entry, for example `NL`, `FR`, or `IT`.\n\nIf set, requests from other countries are denied unless they also match an allowed network.", + "NETWORK_ACCESS_NETWORKS": "Add one allowed client network per entry in CIDR notation. Both IPv4 and IPv6 are supported, for example `192.168.1.5/32`, `192.168.0.0/16`, `10.0.0.0/8`, or `2001:db8::/32`.\n\nNetworks are checked before GeoIP country rules.", "PLAYLIST": "Playlist specifically assigned to the user proxy context.\n\nThe `L`, `V`, and `S` toggles control which output clusters this user should receive from that target:\n`L` = Live, `V` = VOD, `S` = Series.\n\nSelect at least one cluster to activate the filter. If no cluster is selected, the cluster filter is inactive and all clusters are delivered for the assigned target.", "PROXY": "Proxy access definitions or roles for the designated user.", "SERVER": "Target proxy server address or definition.", @@ -810,6 +813,8 @@ "ADD_EXTENSION": "Add Extension", "ADD_FORMAT": "Add Format", "ADD_HEADER": "Add header", + "ADD_COUNTRY": "Add country", + "ADD_NETWORK": "Add network", "ADD_MAPPING": "Add Mapping", "ADD_PATTERN": "Add Pattern", "ADD_PROPERTY": "Add Property", @@ -831,6 +836,8 @@ "ALIASES": "Aliases", "ALIAS_NAME": "Alias Name", "ALL": "All", + "ALLOWED_COUNTRIES": "Allowed Countries", + "ALLOWED_NETWORKS": "Allowed Networks", "API": "Api", "API_CONFIG": "API", "API_CONFIGURATION": "API Configuration", @@ -982,6 +989,9 @@ "FRIENDLY_NAME": "Friendly Name", "FUZZY_MATCHING": "Fuzzy Matching", "GEOIP": "Geo-IP", + "GEOIP_UNAVAILABLE_POLICY_DENY": "Deny", + "GEOIP_UNAVAILABLE_POLICY_ALLOW": "Allow", + "GEOIP_UNAVAILABLE_POLICY": "Geo-IP unavailable policy", "GITHUB": "GitHub", "GRACE_PERIOD": "Grace period ms", "GRACE_PERIOD_HOLD_STREAM": "Grace Period Hold Stream", @@ -1552,6 +1562,11 @@ }, "TARGET_NOT_EXISTS": "Target does not exist", "USER_DELETED": "User successfully deleted" + , + "VALIDATION": { + "NETWORK_ACCESS_COUNTRIES": "Enter one 2-letter ISO country code per entry, for example DE.", + "NETWORK_ACCESS_NETWORKS": "Enter one CIDR per entry, for example 192.168.0.0/16 or 2001:db8::/32." + } }, "SETUP": { "DESC": { diff --git a/frontend/src/app/components/config/admission_strategies.rs b/frontend/src/app/components/config/admission_strategies.rs new file mode 100644 index 000000000..eeac0b3fe --- /dev/null +++ b/frontend/src/app/components/config/admission_strategies.rs @@ -0,0 +1,299 @@ +use crate::{i18n::YewI18n, utils::t_safe}; +use shared::model::{AdmissionStrategy, StreamConfigDto}; +use std::str::FromStr; + +const LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_OLDEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_SAME_IP_OLDEST"; +const LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_LATEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_SAME_IP_LATEST"; +const LABEL_ADMISSION_STRATEGY_EVICT_USER_OLDEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_OLDEST"; +const LABEL_ADMISSION_STRATEGY_EVICT_USER_LATEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_LATEST"; +const LABEL_ADMISSION_STRATEGY_GRACE_INSTANT_STREAM: &str = "LABEL.ADMISSION_STRATEGY_GRACE_INSTANT_STREAM"; +const LABEL_ADMISSION_STRATEGY_GRACE_HOLD_STREAM: &str = "LABEL.ADMISSION_STRATEGY_GRACE_HOLD_STREAM"; + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct AdmissionStrategiesDto { + pub strategies: Option>, +} + +pub(crate) fn admission_strategy_label_key(strategy: AdmissionStrategy) -> &'static str { + match strategy { + AdmissionStrategy::EvictUserSameIpOldest => LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_OLDEST, + AdmissionStrategy::EvictUserSameIpLatest => LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_LATEST, + AdmissionStrategy::EvictUserOldest => LABEL_ADMISSION_STRATEGY_EVICT_USER_OLDEST, + AdmissionStrategy::EvictUserLatest => LABEL_ADMISSION_STRATEGY_EVICT_USER_LATEST, + AdmissionStrategy::GraceInstantStream => LABEL_ADMISSION_STRATEGY_GRACE_INSTANT_STREAM, + AdmissionStrategy::GraceHoldStream => LABEL_ADMISSION_STRATEGY_GRACE_HOLD_STREAM, + } +} + +pub(crate) fn admission_strategy_label(translate: &YewI18n, strategy: AdmissionStrategy) -> String { + t_safe(translate, admission_strategy_label_key(strategy)).unwrap_or_else(|| match strategy { + AdmissionStrategy::EvictUserSameIpOldest => "Evict same-IP oldest stream".to_string(), + AdmissionStrategy::EvictUserSameIpLatest => "Evict same-IP latest stream".to_string(), + AdmissionStrategy::EvictUserOldest => "Evict user oldest stream".to_string(), + AdmissionStrategy::EvictUserLatest => "Evict user latest stream".to_string(), + AdmissionStrategy::GraceInstantStream => "Grace instant stream".to_string(), + AdmissionStrategy::GraceHoldStream => "Grace hold stream".to_string(), + }) +} + +fn is_grace_strategy(strategy: AdmissionStrategy) -> bool { + matches!(strategy, AdmissionStrategy::GraceInstantStream | AdmissionStrategy::GraceHoldStream) +} + +pub(crate) fn is_grace_strategy_tag(tag: &str) -> bool { tag.trim().starts_with("grace_") } + +pub(crate) fn admission_strategy_tags(strategies: Option<&Vec>) -> Option> { + strategies.map(|entries| entries.iter().map(|entry| (*entry).to_string()).collect()) +} + +pub(crate) fn parse_admission_strategy_tags(tags: Option<&[String]>) -> Option> { + let tags = tags?; + let mut parsed = Vec::new(); + for tag in tags { + if let Ok(strategy) = AdmissionStrategy::from_str(tag) { + if !parsed.contains(&strategy) { + parsed.push(strategy); + } + } + } + Some(parsed) +} + +pub(crate) fn filter_disabled_grace_strategy_tags(tags: Vec, grace_period_millis: u64) -> Vec { + if grace_period_millis == 0 { + tags.into_iter().filter(|tag| !is_grace_strategy_tag(tag)).collect() + } else { + tags + } +} + +pub(crate) fn filter_disabled_grace_strategies( + strategies: Option>, + grace_period_millis: u64, +) -> Option> { + strategies.map(|entries| { + if grace_period_millis == 0 { + entries.into_iter().filter(|strategy| !is_grace_strategy(*strategy)).collect() + } else { + entries + } + }) +} + +pub(crate) fn admission_strategy_tag_label(translate: &YewI18n, tag: &str) -> String { + AdmissionStrategy::from_str(tag) + .ok() + .map(|strategy| admission_strategy_label(translate, strategy)) + .unwrap_or_else(|| tag.to_string()) +} + +pub(crate) fn legacy_admission_strategy_tags(stream: &StreamConfigDto) -> Vec { + if stream.grace_period_millis == 0 { + Vec::new() + } else { + vec![(if stream.grace_period_hold_stream { + AdmissionStrategy::GraceHoldStream + } else { + AdmissionStrategy::GraceInstantStream + }) + .to_string()] + } +} + +pub(crate) fn displayed_admission_strategy_tags( + state: &AdmissionStrategiesDto, + stream: &StreamConfigDto, +) -> Vec { + filter_disabled_grace_strategy_tags( + state.strategies.clone().unwrap_or_else(|| { + admission_strategy_tags(stream.admission_strategies.as_ref()) + .unwrap_or_else(|| legacy_admission_strategy_tags(stream)) + }), + stream.grace_period_millis, + ) +} + +pub(crate) fn available_admission_strategies( + selected_tags: &[String], + grace_period_millis: u64, +) -> Vec { + let has_grace = selected_tags.iter().filter_map(|tag| AdmissionStrategy::from_str(tag).ok()).any(is_grace_strategy); + + [ + AdmissionStrategy::EvictUserSameIpOldest, + AdmissionStrategy::EvictUserSameIpLatest, + AdmissionStrategy::EvictUserOldest, + AdmissionStrategy::EvictUserLatest, + AdmissionStrategy::GraceInstantStream, + AdmissionStrategy::GraceHoldStream, + ] + .into_iter() + .filter(|strategy| { + let tag = (*strategy).to_string(); + let grace_available = grace_period_millis > 0 || !is_grace_strategy(*strategy); + !selected_tags.iter().any(|selected| selected == &tag) + && grace_available + && (!has_grace || !is_grace_strategy(*strategy)) + }) + .collect() +} + +pub(crate) fn add_admission_strategy_tag(current: &[String], strategy: AdmissionStrategy) -> Vec { + let mut next = current.to_vec(); + let tag = strategy.to_string(); + if next.iter().any(|selected| selected == &tag) { + return next; + } + // When adding a same-IP eviction rule, insert it before any broader user-wide + // rule (if present) so the backend ordering validation is satisfied. + let is_narrower = + matches!(strategy, AdmissionStrategy::EvictUserSameIpOldest | AdmissionStrategy::EvictUserSameIpLatest); + if is_narrower { + let broader_oldest = AdmissionStrategy::EvictUserOldest.to_string(); + let broader_latest = AdmissionStrategy::EvictUserLatest.to_string(); + + let earliest_pos = next.iter().position(|t| t == &broader_oldest || t == &broader_latest); + if let Some(pos) = earliest_pos { + next.insert(pos, tag); + return next; + } + } + next.push(tag); + next +} + +pub(crate) fn remove_admission_strategy_tag(current: &[String], index: usize) -> Vec { + let mut next = current.to_vec(); + if index < next.len() { + next.remove(index); + } + next +} + +pub(crate) fn move_admission_strategy_tag(current: &[String], index: usize, delta: isize) -> Vec { + let mut next = current.to_vec(); + if let Some(target_index) = index.checked_add_signed(delta) { + if index < next.len() && target_index < next.len() { + next.swap(index, target_index); + // Reject the move if it would create an invalid ordering (broader before narrower). + let strategy_dtos: Vec = + next.iter().filter_map(|t| AdmissionStrategy::from_str(t).ok()).collect(); + if !shared::model::is_valid_admission_strategy_order(&strategy_dtos) { + next.swap(index, target_index); // revert + } + } + } + next +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn admission_strategy_tags_roundtrip() { + let tags = admission_strategy_tags(Some(&vec![ + AdmissionStrategy::EvictUserOldest, + AdmissionStrategy::GraceHoldStream, + ])) + .unwrap_or_default(); + assert_eq!( + parse_admission_strategy_tags(Some(&tags)), + Some(vec![AdmissionStrategy::EvictUserOldest, AdmissionStrategy::GraceHoldStream,]) + ); + } + + #[test] + fn admission_strategy_tags_roundtrip_evict_user_latest() { + let tags = admission_strategy_tags(Some(&vec![AdmissionStrategy::EvictUserLatest])).unwrap_or_default(); + assert_eq!(parse_admission_strategy_tags(Some(&tags)), Some(vec![AdmissionStrategy::EvictUserLatest])); + } + + #[test] + fn invalid_admission_strategy_tags_are_ignored() { + let tags = vec!["evict_user_latest".to_string(), "not-a-strategy".to_string(), "evict_user_latest".to_string()]; + assert_eq!(parse_admission_strategy_tags(Some(&tags)), Some(vec![AdmissionStrategy::EvictUserLatest])); + } + + #[test] + fn displayed_admission_strategies_fall_back_to_legacy_grace() { + let state = AdmissionStrategiesDto::default(); + let stream = StreamConfigDto { + grace_period_millis: 2_000, + grace_period_hold_stream: true, + ..StreamConfigDto::default() + }; + + assert_eq!(displayed_admission_strategy_tags(&state, &stream), vec!["grace_hold_stream".to_string()]); + } + + #[test] + fn available_admission_strategies_hide_second_grace_option() { + let available = available_admission_strategies(&["grace_hold_stream".to_string()], 2_000); + assert!(!available.contains(&AdmissionStrategy::GraceInstantStream)); + assert!(!available.contains(&AdmissionStrategy::GraceHoldStream)); + assert!(available.contains(&AdmissionStrategy::EvictUserSameIpOldest)); + assert!(available.contains(&AdmissionStrategy::EvictUserOldest)); + } + + #[test] + fn available_admission_strategies_hide_grace_when_disabled() { + let available = available_admission_strategies(&[], 0); + assert!(!available.contains(&AdmissionStrategy::GraceInstantStream)); + assert!(!available.contains(&AdmissionStrategy::GraceHoldStream)); + assert!(available.contains(&AdmissionStrategy::EvictUserSameIpOldest)); + assert!(available.contains(&AdmissionStrategy::EvictUserSameIpLatest)); + assert!(available.contains(&AdmissionStrategy::EvictUserOldest)); + assert!(available.contains(&AdmissionStrategy::EvictUserLatest)); + } + + #[test] + fn displayed_admission_strategies_hide_disabled_grace_tags() { + let state = AdmissionStrategiesDto { strategies: Some(vec!["grace_hold_stream".to_string()]) }; + let stream = StreamConfigDto { grace_period_millis: 0, ..StreamConfigDto::default() }; + + assert_eq!(displayed_admission_strategy_tags(&state, &stream), Vec::::new()); + } + + #[test] + fn filtered_admission_strategies_drop_grace_when_disabled() { + let parsed = parse_admission_strategy_tags(Some(&[ + "evict_user_same_ip_oldest".to_string(), + "grace_hold_stream".to_string(), + ])); + + assert_eq!(filter_disabled_grace_strategies(parsed, 0), Some(vec![AdmissionStrategy::EvictUserSameIpOldest])); + } + + #[test] + fn add_admission_strategy_enforces_narrower_before_broader() { + let current = vec!["evict_user_oldest".to_string()]; + let new_tags = add_admission_strategy_tag(¤t, AdmissionStrategy::EvictUserSameIpOldest); + assert_eq!(new_tags, vec!["evict_user_same_ip_oldest".to_string(), "evict_user_oldest".to_string()]); + } + + #[test] + fn add_admission_strategy_inserts_between_existing_narrower_and_broader_rules() { + let current = + vec![AdmissionStrategy::EvictUserSameIpOldest.to_string(), AdmissionStrategy::EvictUserOldest.to_string()]; + + let new_tags = add_admission_strategy_tag(¤t, AdmissionStrategy::EvictUserSameIpLatest); + + assert_eq!( + new_tags, + vec![ + AdmissionStrategy::EvictUserSameIpOldest.to_string(), + AdmissionStrategy::EvictUserSameIpLatest.to_string(), + AdmissionStrategy::EvictUserOldest.to_string(), + ] + ); + } + + #[test] + fn move_admission_strategy_reverts_invalid_order() { + let current = vec!["evict_user_same_ip_oldest".to_string(), "evict_user_oldest".to_string()]; + // Attempt to move broader EvictUserOldest up before narrower EvictUserSameIpOldest + let next = move_admission_strategy_tag(¤t, 1, -1); + assert_eq!(next, current); + } +} diff --git a/frontend/src/app/components/config/macros.rs b/frontend/src/app/components/config/macros.rs index d219688cc..e41b21b4c 100644 --- a/frontend/src/app/components/config/macros.rs +++ b/frontend/src/app/components/config/macros.rs @@ -693,6 +693,27 @@ macro_rules! edit_field_list_option { } }}; + ($state:expr, $label:expr, $field_id:expr, $placeholder:expr, $create_tag:expr) => {{ + let state = $state.clone(); + let create_tag = $create_tag.clone(); + html! { +
+ <$crate::app::components::FieldLabel + label={$label.to_string()} + field_id={$field_id.to_string()} + /> + <$crate::app::components::TagList + tags={(*state).clone()} + placeholder={$placeholder} + readonly={false} + create_tag={create_tag} + on_change={Callback::from(move |value: Vec>| { + state.set(value); + })} + /> +
+ } + }}; } #[macro_export] diff --git a/frontend/src/app/components/config/mod.rs b/frontend/src/app/components/config/mod.rs index 0850e2a80..79ae9de8f 100644 --- a/frontend/src/app/components/config/mod.rs +++ b/frontend/src/app/components/config/mod.rs @@ -1,5 +1,6 @@ mod macros; +mod admission_strategies; mod api_config_view; mod config_page; mod config_update; @@ -21,6 +22,7 @@ mod schedules_config_view; mod video_config_view; mod webui_config_view; +pub(crate) use admission_strategies::*; pub use api_config_view::*; pub use config_page::*; pub use config_view::*; diff --git a/frontend/src/app/components/config/reverse_proxy_config_view.rs b/frontend/src/app/components/config/reverse_proxy_config_view.rs index 1febb9f62..3011c2b00 100644 --- a/frontend/src/app/components/config/reverse_proxy_config_view.rs +++ b/frontend/src/app/components/config/reverse_proxy_config_view.rs @@ -4,11 +4,15 @@ use crate::{ app::{ components::{ config::{ + add_admission_strategy_tag, admission_strategy_label, admission_strategy_tag_label, + admission_strategy_tags, available_admission_strategies, config_page::{ConfigForm, LABEL_REVERSE_PROXY_CONFIG}, config_view_context::ConfigViewContext, - use_emit_mapped_option, + displayed_admission_strategy_tags, filter_disabled_grace_strategies, move_admission_strategy_tag, + parse_admission_strategy_tags, remove_admission_strategy_tag, use_emit_mapped_option, + AdmissionStrategiesDto, }, - Card, Chip, IconButton, TextButton, + Card, Chip, IconButton, RadioButtonGroup, TextButton, }, context::ConfigContext, }, @@ -16,18 +20,18 @@ use crate::{ edit_field_bool, edit_field_list, edit_field_number, edit_field_number_f64, edit_field_number_u16, edit_field_number_u64, edit_field_number_usize, edit_field_text, edit_field_text_option, generate_form_reducer, i18n::{use_translation, YewI18n}, - utils::t_safe, }; +use enum_iterator::all; use shared::{ model::{ - AdmissionStrategy, CacheConfigDto, GeoIpConfigDto, QosAggregationConfigDto, RateLimitConfigDto, + CacheConfigDto, GeoIpConfigDto, GeoIpUnavailablePolicy, QosAggregationConfigDto, RateLimitConfigDto, ResourceRetryConfigDto, ReverseProxyConfigDto, ReverseProxyDisabledHeaderConfigDto, StreamBufferConfigDto, StreamConfigDto, StreamHistoryConfigDto, }, utils::{default_secret, format_float_localized}, }; +use std::{rc::Rc, str::FromStr}; use yew::prelude::*; - const LABEL_CACHE: &str = "LABEL.CACHE"; const LABEL_ENABLED: &str = "LABEL.ENABLED"; const LABEL_SIZE: &str = "LABEL.SIZE"; @@ -71,6 +75,7 @@ const LABEL_CF_HEADER: &str = "LABEL.CF_HEADER"; const LABEL_CUSTOM_HEADERS: &str = "LABEL.CUSTOM_HEADERS"; const LABEL_ADD_HEADER: &str = "LABEL.ADD_HEADER"; const LABEL_GEOIP: &str = "LABEL.GEOIP"; +const LABEL_GEOIP_UNAVAILABLE_POLICY: &str = "LABEL.GEOIP_UNAVAILABLE_POLICY"; const LABEL_URL: &str = "LABEL.URL"; const LABEL_STREAM_HISTORY: &str = "LABEL.STREAM_HISTORY"; @@ -80,12 +85,6 @@ const LABEL_STREAM_HISTORY_RETENTION_DAYS: &str = "LABEL.STREAM_HISTORY_RETENTIO const LABEL_QOS_AGGREGATION: &str = "LABEL.QOS_AGGREGATION"; const LABEL_QOS_AGGREGATION_ENABLED: &str = "LABEL.QOS_AGGREGATION_ENABLED"; const LABEL_INTERVAL_SECS: &str = "LABEL.INTERVAL_SECS"; -const LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_OLDEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_SAME_IP_OLDEST"; -const LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_LATEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_SAME_IP_LATEST"; -const LABEL_ADMISSION_STRATEGY_EVICT_USER_OLDEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_OLDEST"; -const LABEL_ADMISSION_STRATEGY_EVICT_USER_LATEST: &str = "LABEL.ADMISSION_STRATEGY_EVICT_USER_LATEST"; -const LABEL_ADMISSION_STRATEGY_GRACE_INSTANT_STREAM: &str = "LABEL.ADMISSION_STRATEGY_GRACE_INSTANT_STREAM"; -const LABEL_ADMISSION_STRATEGY_GRACE_HOLD_STREAM: &str = "LABEL.ADMISSION_STRATEGY_GRACE_HOLD_STREAM"; generate_form_reducer!( state: CacheConfigFormState { form: CacheConfigDto }, @@ -127,202 +126,6 @@ impl FailoverPatternsDto { pub fn is_empty(&self) -> bool { self.patterns.is_empty() } } -#[derive(Debug, Clone, Default, PartialEq)] -pub struct AdmissionStrategiesDto { - pub strategies: Option>, -} - -fn admission_strategy_tag(strategy: AdmissionStrategy) -> &'static str { - match strategy { - AdmissionStrategy::EvictUserSameIpOldest => "evict_user_same_ip_oldest", - AdmissionStrategy::EvictUserSameIpLatest => "evict_user_same_ip_latest", - AdmissionStrategy::EvictUserOldest => "evict_user_oldest", - AdmissionStrategy::EvictUserLatest => "evict_user_latest", - AdmissionStrategy::GraceInstantStream => "grace_instant_stream", - AdmissionStrategy::GraceHoldStream => "grace_hold_stream", - } -} - -fn parse_admission_strategy_tag(tag: &str) -> Option { - match tag.trim() { - "evict_user_same_ip_oldest" => Some(AdmissionStrategy::EvictUserSameIpOldest), - "evict_user_same_ip_latest" => Some(AdmissionStrategy::EvictUserSameIpLatest), - "evict_user_oldest" => Some(AdmissionStrategy::EvictUserOldest), - "evict_user_latest" => Some(AdmissionStrategy::EvictUserLatest), - "grace_instant_stream" => Some(AdmissionStrategy::GraceInstantStream), - "grace_hold_stream" => Some(AdmissionStrategy::GraceHoldStream), - _ => None, - } -} - -fn admission_strategy_label_key(strategy: AdmissionStrategy) -> &'static str { - match strategy { - AdmissionStrategy::EvictUserSameIpOldest => LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_OLDEST, - AdmissionStrategy::EvictUserSameIpLatest => LABEL_ADMISSION_STRATEGY_EVICT_USER_SAME_IP_LATEST, - AdmissionStrategy::EvictUserOldest => LABEL_ADMISSION_STRATEGY_EVICT_USER_OLDEST, - AdmissionStrategy::EvictUserLatest => LABEL_ADMISSION_STRATEGY_EVICT_USER_LATEST, - AdmissionStrategy::GraceInstantStream => LABEL_ADMISSION_STRATEGY_GRACE_INSTANT_STREAM, - AdmissionStrategy::GraceHoldStream => LABEL_ADMISSION_STRATEGY_GRACE_HOLD_STREAM, - } -} - -fn admission_strategy_label(translate: &YewI18n, strategy: AdmissionStrategy) -> String { - t_safe(translate, admission_strategy_label_key(strategy)).unwrap_or_else(|| match strategy { - AdmissionStrategy::EvictUserSameIpOldest => "Evict same-IP oldest stream".to_string(), - AdmissionStrategy::EvictUserSameIpLatest => "Evict same-IP latest stream".to_string(), - AdmissionStrategy::EvictUserOldest => "Evict user oldest stream".to_string(), - AdmissionStrategy::EvictUserLatest => "Evict user latest stream".to_string(), - AdmissionStrategy::GraceInstantStream => "Grace instant stream".to_string(), - AdmissionStrategy::GraceHoldStream => "Grace hold stream".to_string(), - }) -} - -fn is_grace_strategy(strategy: AdmissionStrategy) -> bool { - matches!(strategy, AdmissionStrategy::GraceInstantStream | AdmissionStrategy::GraceHoldStream) -} - -fn is_grace_strategy_tag(tag: &str) -> bool { - let tag = tag.trim(); - tag.starts_with("grace_") || parse_admission_strategy_tag(tag).is_some_and(is_grace_strategy) -} - -fn admission_strategy_tags(strategies: Option<&Vec>) -> Option> { - strategies.map(|entries| entries.iter().map(|entry| admission_strategy_tag(*entry).to_string()).collect()) -} - -fn parse_admission_strategy_tags(tags: Option<&[String]>) -> Option> { - let tags = tags?; - let mut parsed = Vec::new(); - for tag in tags { - if let Some(strategy) = parse_admission_strategy_tag(tag) { - if !parsed.contains(&strategy) { - parsed.push(strategy); - } - } - } - Some(parsed) -} - -fn filter_disabled_grace_strategy_tags(tags: Vec, grace_period_millis: u64) -> Vec { - if grace_period_millis == 0 { - tags.into_iter().filter(|tag| !is_grace_strategy_tag(tag)).collect() - } else { - tags - } -} - -fn filter_disabled_grace_strategies( - strategies: Option>, - grace_period_millis: u64, -) -> Option> { - strategies.map(|entries| { - if grace_period_millis == 0 { - entries.into_iter().filter(|strategy| !is_grace_strategy(*strategy)).collect() - } else { - entries - } - }) -} - -fn admission_strategy_tag_label(translate: &YewI18n, tag: &str) -> String { - parse_admission_strategy_tag(tag) - .map(|strategy| admission_strategy_label(translate, strategy)) - .unwrap_or_else(|| tag.to_string()) -} - -fn legacy_admission_strategy_tags(stream: &StreamConfigDto) -> Vec { - if stream.grace_period_millis == 0 { - Vec::new() - } else { - vec![admission_strategy_tag(if stream.grace_period_hold_stream { - AdmissionStrategy::GraceHoldStream - } else { - AdmissionStrategy::GraceInstantStream - }) - .to_string()] - } -} - -fn displayed_admission_strategy_tags(state: &AdmissionStrategiesDto, stream: &StreamConfigDto) -> Vec { - filter_disabled_grace_strategy_tags( - state.strategies.clone().unwrap_or_else(|| { - admission_strategy_tags(stream.admission_strategies.as_ref()) - .unwrap_or_else(|| legacy_admission_strategy_tags(stream)) - }), - stream.grace_period_millis, - ) -} - -fn available_admission_strategies(selected_tags: &[String], grace_period_millis: u64) -> Vec { - let has_grace = selected_tags.iter().filter_map(|tag| parse_admission_strategy_tag(tag)).any(is_grace_strategy); - - [ - AdmissionStrategy::EvictUserSameIpOldest, - AdmissionStrategy::EvictUserSameIpLatest, - AdmissionStrategy::EvictUserOldest, - AdmissionStrategy::EvictUserLatest, - AdmissionStrategy::GraceInstantStream, - AdmissionStrategy::GraceHoldStream, - ] - .into_iter() - .filter(|strategy| { - let tag = admission_strategy_tag(*strategy); - let grace_available = grace_period_millis > 0 || !is_grace_strategy(*strategy); - !selected_tags.iter().any(|selected| selected == tag) - && grace_available - && (!has_grace || !is_grace_strategy(*strategy)) - }) - .collect() -} - -fn add_admission_strategy_tag(current: &[String], strategy: AdmissionStrategy) -> Vec { - let mut next = current.to_vec(); - let tag = admission_strategy_tag(strategy).to_string(); - if next.iter().any(|selected| selected == &tag) { - return next; - } - // When adding a same-IP eviction rule, insert it before any broader user-wide - // rule (if present) so the backend ordering validation is satisfied. - let is_narrower = - matches!(strategy, AdmissionStrategy::EvictUserSameIpOldest | AdmissionStrategy::EvictUserSameIpLatest); - if is_narrower { - let broader_oldest = admission_strategy_tag(AdmissionStrategy::EvictUserOldest); - let broader_latest = admission_strategy_tag(AdmissionStrategy::EvictUserLatest); - - let earliest_pos = next.iter().position(|t| t == broader_oldest || t == broader_latest); - if let Some(pos) = earliest_pos { - next.insert(pos, tag); - return next; - } - } - next.push(tag); - next -} - -fn remove_admission_strategy_tag(current: &[String], index: usize) -> Vec { - let mut next = current.to_vec(); - if index < next.len() { - next.remove(index); - } - next -} - -fn move_admission_strategy_tag(current: &[String], index: usize, delta: isize) -> Vec { - let mut next = current.to_vec(); - if let Some(target_index) = index.checked_add_signed(delta) { - if index < next.len() && target_index < next.len() { - next.swap(index, target_index); - // Reject the move if it would create an invalid ordering (broader before narrower). - let strategy_dtos: Vec = - next.iter().filter_map(|t| parse_admission_strategy_tag(t)).collect(); - if !shared::model::is_valid_admission_strategy_order(&strategy_dtos) { - next.swap(index, target_index); // revert - } - } - } - next -} - generate_form_reducer!( state: FailoverPatternsFormState { form: FailoverPatternsDto }, action_name: FailoverPatternsFormAction, @@ -371,6 +174,7 @@ generate_form_reducer!( fields { Enabled => enabled: bool, Url => url: String, + UnavailablePolicy => unavailable_policy: GeoIpUnavailablePolicy, } ); @@ -414,6 +218,24 @@ generate_form_reducer!( } ); +fn geoip_unavailable_policy_options() -> Rc> { + Rc::new(all::().map(|policy| policy.to_string()).collect()) +} + +pub(crate) fn geoip_unavailable_policy_label(translate: &YewI18n, policy: GeoIpUnavailablePolicy) -> String { + match policy { + GeoIpUnavailablePolicy::Deny => translate.t("LABEL.GEOIP_UNAVAILABLE_POLICY_DENY"), + GeoIpUnavailablePolicy::Allow => translate.t("LABEL.GEOIP_UNAVAILABLE_POLICY_ALLOW"), + } +} + +fn geoip_unavailable_policy_labels(translate: &YewI18n) -> Rc> { + Rc::new(vec![ + geoip_unavailable_policy_label(translate, GeoIpUnavailablePolicy::Deny), + geoip_unavailable_policy_label(translate, GeoIpUnavailablePolicy::Allow), + ]) +} + #[component] pub fn ReverseProxyConfigView() -> Html { let translate = use_translation(); @@ -816,6 +638,13 @@ pub fn ReverseProxyConfigView() -> Html {

{translate.t(LABEL_GEOIP)}

{ config_field_bool!(geoip_state.form, translate.t(LABEL_ENABLED), enabled) } { config_field!(geoip_state.form, translate.t(LABEL_URL), url) } + { config_field_child!(translate.t(LABEL_GEOIP_UNAVAILABLE_POLICY), "GEO_IP_CONFIG.UNAVAILABLE_POLICY", { + html! { + + {geoip_unavailable_policy_label(&translate, geoip_state.form.unavailable_policy)} + + } + }) } } }; @@ -910,11 +739,31 @@ pub fn ReverseProxyConfigView() -> Html { }; let render_geoip_edit = || { + let geoip_policy_state = geoip_state.clone(); + let selected_policy = Rc::new(vec![geoip_state.form.unavailable_policy.to_string()]); html! {

{translate.t(LABEL_GEOIP)}

{ edit_field_bool!(geoip_state, translate.t(LABEL_ENABLED), enabled, GeoIpConfigFormAction::Enabled) } { edit_field_text!(geoip_state, translate.t(LABEL_URL), url, GeoIpConfigFormAction::Url) } + { config_field_child!(translate.t(LABEL_GEOIP_UNAVAILABLE_POLICY), "GEO_IP_CONFIG.UNAVAILABLE_POLICY", { + html! { + >| { + if let Some(selection) = selections.first() { + geoip_policy_state.dispatch(GeoIpConfigFormAction::UnavailablePolicy( + GeoIpUnavailablePolicy::from_str(selection).unwrap_or(GeoIpUnavailablePolicy::Deny), + )); + } + })} + /> + } + }) }
} }; @@ -1156,92 +1005,10 @@ mod tests { use super::*; #[test] - fn admission_strategy_tags_roundtrip() { - let tags = admission_strategy_tags(Some(&vec![ - AdmissionStrategy::EvictUserOldest, - AdmissionStrategy::GraceHoldStream, - ])) - .unwrap_or_default(); - assert_eq!( - parse_admission_strategy_tags(Some(&tags)), - Some(vec![AdmissionStrategy::EvictUserOldest, AdmissionStrategy::GraceHoldStream,]) - ); - } - - #[test] - fn admission_strategy_tags_roundtrip_evict_user_latest() { - let tags = admission_strategy_tags(Some(&vec![AdmissionStrategy::EvictUserLatest])).unwrap_or_default(); - assert_eq!(parse_admission_strategy_tags(Some(&tags)), Some(vec![AdmissionStrategy::EvictUserLatest])); - } - - #[test] - fn invalid_admission_strategy_tags_are_ignored() { - let tags = vec!["evict_user_latest".to_string(), "not-a-strategy".to_string(), "evict_user_latest".to_string()]; - assert_eq!(parse_admission_strategy_tags(Some(&tags)), Some(vec![AdmissionStrategy::EvictUserLatest])); - } - - #[test] - fn displayed_admission_strategies_fall_back_to_legacy_grace() { - let state = AdmissionStrategiesDto::default(); - let stream = StreamConfigDto { - grace_period_millis: 2_000, - grace_period_hold_stream: true, - ..StreamConfigDto::default() - }; - - assert_eq!(displayed_admission_strategy_tags(&state, &stream), vec!["grace_hold_stream".to_string()]); - } - - #[test] - fn available_admission_strategies_hide_second_grace_option() { - let available = available_admission_strategies(&["grace_hold_stream".to_string()], 2_000); - assert!(!available.contains(&AdmissionStrategy::GraceInstantStream)); - assert!(!available.contains(&AdmissionStrategy::GraceHoldStream)); - assert!(available.contains(&AdmissionStrategy::EvictUserSameIpOldest)); - assert!(available.contains(&AdmissionStrategy::EvictUserOldest)); - } - - #[test] - fn available_admission_strategies_hide_grace_when_disabled() { - let available = available_admission_strategies(&[], 0); - assert!(!available.contains(&AdmissionStrategy::GraceInstantStream)); - assert!(!available.contains(&AdmissionStrategy::GraceHoldStream)); - assert!(available.contains(&AdmissionStrategy::EvictUserSameIpOldest)); - assert!(available.contains(&AdmissionStrategy::EvictUserSameIpLatest)); - assert!(available.contains(&AdmissionStrategy::EvictUserOldest)); - assert!(available.contains(&AdmissionStrategy::EvictUserLatest)); - } - - #[test] - fn displayed_admission_strategies_hide_disabled_grace_tags() { - let state = AdmissionStrategiesDto { strategies: Some(vec!["grace_hold_stream".to_string()]) }; - let stream = StreamConfigDto { grace_period_millis: 0, ..StreamConfigDto::default() }; - - assert_eq!(displayed_admission_strategy_tags(&state, &stream), Vec::::new()); - } - - #[test] - fn filtered_admission_strategies_drop_grace_when_disabled() { - let parsed = parse_admission_strategy_tags(Some(&[ - "evict_user_same_ip_oldest".to_string(), - "grace_hold_stream".to_string(), - ])); - - assert_eq!(filter_disabled_grace_strategies(parsed, 0), Some(vec![AdmissionStrategy::EvictUserSameIpOldest])); - } - - #[test] - fn add_admission_strategy_enforces_narrower_before_broader() { - let current = vec!["evict_user_oldest".to_string()]; - let new_tags = add_admission_strategy_tag(¤t, AdmissionStrategy::EvictUserSameIpOldest); - assert_eq!(new_tags, vec!["evict_user_same_ip_oldest".to_string(), "evict_user_oldest".to_string()]); - } - - #[test] - fn move_admission_strategy_reverts_invalid_order() { - let current = vec!["evict_user_same_ip_oldest".to_string(), "evict_user_oldest".to_string()]; - // Attempt to move broader EvictUserOldest up before narrower EvictUserSameIpOldest - let next = move_admission_strategy_tag(¤t, 1, -1); - assert_eq!(next, current); + fn geoip_unavailable_policy_roundtrips_through_string_representation() { + for policy in all::() { + let parsed = GeoIpUnavailablePolicy::from_str(&policy.to_string()).unwrap_or(GeoIpUnavailablePolicy::Deny); + assert_eq!(parsed, policy); + } } } diff --git a/frontend/src/app/components/radio_button_group.rs b/frontend/src/app/components/radio_button_group.rs index d9890c442..d647f1245 100644 --- a/frontend/src/app/components/radio_button_group.rs +++ b/frontend/src/app/components/radio_button_group.rs @@ -7,6 +7,8 @@ pub struct RadioButtonGroupProps { pub options: Rc>, pub selected: Rc>, pub on_select: Callback>>, + #[prop_or_default] + pub labels: Option>>, // optional localized labels, same length as options #[prop_or(false)] pub multi_select: bool, #[prop_or(false)] @@ -62,14 +64,26 @@ pub fn RadioButtonGroup(props: &RadioButtonGroupProps) -> Html { }) }; + let display_label = |option: &String| -> String { + if let Some(labels) = &props.labels { + if let Some(pos) = props.options.iter().position(|o| o == option) { + if let Some(label) = labels.get(pos) { + return label.clone(); + } + } + } + option.clone() + }; + html! {
{ for props.options.iter().map(|option| { let is_selected = (*selections).contains(option); let class = if is_selected { "primary" } else { "" }; let onclick = on_click.clone(); + let label = display_label(option); html! { - + } }) }
diff --git a/frontend/src/app/components/tag_list.rs b/frontend/src/app/components/tag_list.rs index 59f2d388b..d5955f195 100644 --- a/frontend/src/app/components/tag_list.rs +++ b/frontend/src/app/components/tag_list.rs @@ -9,11 +9,15 @@ pub struct Tag { pub class: Option, } +fn default_create_tag(value: String) -> Option { Some(Tag { label: value, class: None }) } + #[derive(Properties, Clone, PartialEq)] pub struct TagListProps { pub tags: Vec>, #[prop_or_else(Callback::noop)] pub on_change: Callback>>, + #[prop_or_else(|| Callback::from(default_create_tag))] + pub create_tag: Callback>, #[prop_or(true)] pub readonly: bool, #[prop_or_else(|| "Add tag...".to_string())] @@ -22,7 +26,7 @@ pub struct TagListProps { #[component] pub fn TagList(props: &TagListProps) -> Html { - let TagListProps { tags, on_change, readonly, placeholder } = props.clone(); + let TagListProps { tags, on_change, create_tag, readonly, placeholder } = props.clone(); let tag_state = use_state(|| tags.clone()); let new_tag = use_state(String::default); @@ -61,15 +65,22 @@ pub fn TagList(props: &TagListProps) -> Html { let new_tag = new_tag.clone(); let tag_state = tag_state.clone(); let on_change = on_change.clone(); + let create_tag = create_tag.clone(); Callback::from(move |()| { let val = (*new_tag).trim().to_string(); - if !val.is_empty() && !tag_state.iter().any(|t| t.label == val) { - let mut updated = (*tag_state).clone(); - updated.push(Rc::new(Tag { label: val.clone(), class: None })); - on_change.emit(updated.clone()); - tag_state.set(updated); + if val.is_empty() { + return; + } + + if let Some(next_tag) = create_tag.emit(val.clone()) { + if !tag_state.iter().any(|t| t.label == next_tag.label) { + let mut updated = (*tag_state).clone(); + updated.push(Rc::new(next_tag)); + on_change.emit(updated.clone()); + tag_state.set(updated); + new_tag.set(String::new()); + } } - new_tag.set(String::new()); }) }; diff --git a/frontend/src/app/components/userlist/proxy_user_credentials_form.rs b/frontend/src/app/components/userlist/proxy_user_credentials_form.rs index 132297fcb..a1f9337d3 100644 --- a/frontend/src/app/components/userlist/proxy_user_credentials_form.rs +++ b/frontend/src/app/components/userlist/proxy_user_credentials_form.rs @@ -2,12 +2,13 @@ use crate::{ app::{ components::{ config::HasFormData, select::Select, userlist::proxy_type_input::ProxyTypeInput, ClusterFlagsInput, - ClusterFlagsInputMode, DropDownOption, DropDownSelection, TextButton, UserStatus, + ClusterFlagsInputMode, DropDownOption, DropDownSelection, Tag, TextButton, UserStatus, }, TargetUser, }, - config_field_child, config_field_custom, edit_field_bool, edit_field_date, edit_field_number, edit_field_number_i8, - edit_field_number_u16, edit_field_text, edit_field_text_option, generate_form_reducer, + config_field_child, config_field_custom, edit_field_bool, edit_field_date, edit_field_list_option, + edit_field_number, edit_field_number_i8, edit_field_number_u16, edit_field_text, edit_field_text_option, + generate_form_reducer, hooks::use_service_context, html_if, i18n::use_translation, @@ -15,17 +16,83 @@ use crate::{ use chrono::{Duration, Utc}; use shared::{ model::{ - permission::Permission, ApiProxyServerInfoDto, ClusterFlags, ConfigTargetDto, ProxyType, + permission::Permission, ApiProxyServerInfoDto, ClusterFlags, ConfigTargetDto, NetworkAccessDto, ProxyType, ProxyUserCredentialsDto, ProxyUserStatus, }, utils::generate_random_string, }; -use std::rc::Rc; +use std::{net::IpAddr, rc::Rc}; use yew::prelude::*; const DEFAULT_MAX_CONNECTIONS: u32 = 1; const DEFAULT_EXPIRATION_DAYS: i64 = 365; +fn normalize_country_entry(input: &str) -> Result { + let normalized = input.trim().to_ascii_uppercase(); + if normalized.len() == 2 && normalized.chars().all(|ch| ch.is_ascii_alphabetic()) { + Ok(normalized) + } else { + Err("MESSAGES.VALIDATION.NETWORK_ACCESS_COUNTRIES") + } +} + +fn normalize_network_entry(input: &str) -> Result { + let normalized = input.trim().to_string(); + let Some((address, prefix)) = normalized.split_once('/') else { + return Err("MESSAGES.VALIDATION.NETWORK_ACCESS_NETWORKS"); + }; + let address = address.trim().to_ascii_lowercase(); + let prefix = prefix.trim(); + let Ok(ip) = address.parse::() else { + return Err("MESSAGES.VALIDATION.NETWORK_ACCESS_NETWORKS"); + }; + let Ok(prefix) = prefix.parse::() else { + return Err("MESSAGES.VALIDATION.NETWORK_ACCESS_NETWORKS"); + }; + let prefix_valid = match ip { + IpAddr::V4(_) => prefix <= 32, + IpAddr::V6(_) => prefix <= 128, + }; + if prefix_valid { + Ok(normalized) + } else { + Err("MESSAGES.VALIDATION.NETWORK_ACCESS_NETWORKS") + } +} + +fn build_network_access(countries: &[String], networks: &[String]) -> Result, String> { + let countries_list = countries.iter().filter(|s| !s.trim().is_empty()).cloned().collect::>(); + let networks_list = networks.iter().filter(|s| !s.trim().is_empty()).cloned().collect::>(); + if countries_list.is_empty() && networks_list.is_empty() { + return Ok(None); + } + let mut dto = NetworkAccessDto { + allowed_countries: if countries_list.is_empty() { None } else { Some(countries_list) }, + allowed_networks: if networks_list.is_empty() { None } else { Some(networks_list) }, + }; + dto.prepare().map_err(|e| e.to_string())?; + Ok(Some(dto)) +} + +fn network_access_changed( + original: Option<&NetworkAccessDto>, + countries: &[String], + networks: &[String], +) -> Result { + let built = build_network_access(countries, networks)?; + Ok(original != built.as_ref()) +} + +fn validate_network_access(countries: &[String], networks: &[String]) -> Result<(), &'static str> { + for country in countries { + normalize_country_entry(country)?; + } + for network in networks { + normalize_network_entry(network)?; + } + Ok(()) +} + generate_form_reducer!( state: UserFormState { form: ProxyUserCredentialsDto }, action_name: UserFormAction, @@ -64,6 +131,8 @@ pub fn ProxyUserCredentialsForm(props: &ProxyUserCredentialsFormProps) -> Html { let service_ctx = use_service_context(); let selected_target = use_state(|| None); let update = use_state(|| false); + let allowed_countries = use_state(Vec::>::new); + let allowed_networks = use_state(Vec::>::new); let form_state: UseReducerHandle = use_reducer(|| UserFormState { form: ProxyUserCredentialsDto::default(), modified: false }); @@ -104,14 +173,30 @@ pub fn ProxyUserCredentialsForm(props: &ProxyUserCredentialsFormProps) -> Html { let form_state = form_state.clone(); let set_selected_target = selected_target.clone(); let set_update = update.clone(); + let set_allowed_countries = allowed_countries.clone(); + let set_allowed_networks = allowed_networks.clone(); use_effect_with((props.user.clone(), props.server.clone()), move |(user, server)| { if let Some(u) = user.clone() { set_update.set(true); set_selected_target.set(Some(u.target.clone())); - form_state.dispatch(UserFormAction::SetAll((*u.credentials).clone())); + let creds = (*u.credentials).clone(); + if let Some(na) = &creds.network_access { + set_allowed_countries.set(na.allowed_countries.as_ref().map_or_else(Vec::new, |countries| { + countries.iter().map(|country| Rc::new(Tag { label: country.clone(), class: None })).collect() + })); + set_allowed_networks.set(na.allowed_networks.as_ref().map_or_else(Vec::new, |networks| { + networks.iter().map(|network| Rc::new(Tag { label: network.clone(), class: None })).collect() + })); + } else { + set_allowed_countries.set(Vec::new()); + set_allowed_networks.set(Vec::new()); + } + form_state.dispatch(UserFormAction::SetAll(creds)); } else { set_update.set(false); set_selected_target.set(None); + set_allowed_countries.set(Vec::new()); + set_allowed_networks.set(Vec::new()); let mut user = ProxyUserCredentialsDto::default(); if let Some(api_server) = (*server).first() { user.server = Some(api_server.name.clone()); @@ -154,21 +239,47 @@ pub fn ProxyUserCredentialsForm(props: &ProxyUserCredentialsFormProps) -> Html { let target = selected_target.clone(); let onsave = props.on_save.clone(); let is_update = update.clone(); + let countries = allowed_countries.clone(); + let networks = allowed_networks.clone(); Callback::from(move |_| { let nothing_to_save = || services.toastr.warning(translate_clone.t("MESSAGES.SAVE.USER.NOTHING_TO_SAVE")); if let Some(target_name) = (*target).as_ref().cloned() { let original_target = original.as_ref().map(|u| u.target.clone()).unwrap_or_default(); let target_changed = target_name != original_target; - if target_changed || user.modified() { - let user = user.data(); + let countries_value = countries.iter().map(|tag| tag.label.clone()).collect::>(); + let networks_value = networks.iter().map(|tag| tag.label.clone()).collect::>(); + let network_access_changed = match network_access_changed( + original.as_ref().and_then(|u| u.credentials.network_access.as_ref()), + &countries_value, + &networks_value, + ) { + Ok(changed) => changed, + Err(err) => { + services.toastr.error(err); + return; + } + }; + if target_changed || user.modified() || network_access_changed { + let mut user = user.data().clone(); + if let Err(message_key) = validate_network_access(&countries_value, &networks_value) { + services.toastr.error(translate_clone.t(message_key)); + return; + } + user.network_access = match build_network_access(&countries_value, &networks_value) { + Ok(na) => na, + Err(err) => { + services.toastr.error(err); + return; + } + }; if let Err(err) = user.validate() { services.toastr.error(err.to_string()); } else { match original.as_ref().map(|t| t.credentials.clone()) { - None => onsave.emit((*is_update, target_name, user.clone())), + None => onsave.emit((*is_update, target_name, user)), Some(original_user) => { - if target_changed || &(*original_user) != user { - onsave.emit((*is_update, target_name, user.clone())); + if target_changed || (*original_user) != user { + onsave.emit((*is_update, target_name, user)); } else { nothing_to_save(); } @@ -190,6 +301,24 @@ pub fn ProxyUserCredentialsForm(props: &ProxyUserCredentialsFormProps) -> Html { let instance_proxy = form_state.clone(); let instance_server = form_state.clone(); let instance_output_clusters = form_state.clone(); + let country_services = service_ctx.clone(); + let country_translate = translate.clone(); + let create_country_tag = Callback::from(move |value: String| match normalize_country_entry(&value) { + Ok(normalized) => Some(Tag { label: normalized, class: None }), + Err(message_key) => { + country_services.toastr.error(country_translate.t(message_key)); + None + } + }); + let network_services = service_ctx.clone(); + let network_translate = translate.clone(); + let create_network_tag = Callback::from(move |value: String| match normalize_network_entry(&value) { + Ok(normalized) => Some(Tag { label: normalized, class: None }), + Err(message_key) => { + network_services.toastr.error(network_translate.t(message_key)); + None + } + }); html! {
@@ -199,7 +328,7 @@ pub fn ProxyUserCredentialsForm(props: &ProxyUserCredentialsFormProps) -> Html {
None, DropDownSelection::Single(option) => option.parse::().ok(), @@ -256,7 +385,7 @@ pub fn ProxyUserCredentialsForm(props: &ProxyUserCredentialsFormProps) -> Html { html! {