From faac89325b2d7803509db576bb53fdf55fdfea3e Mon Sep 17 00:00:00 2001 From: knylbyte <40831653+knylbyte@users.noreply.github.com> Date: Mon, 2 Mar 2026 12:24:05 +0100 Subject: [PATCH] feat: add provider DNS IP-connect, resolved persistence, and SNI-safe failover (#624) * feat: add provider DNS IP-connect, resolved persistence, and SNI-safe failover --- .gitignore | 1 + Cargo.lock | 134 ++++- README.md | 63 +- backend/Cargo.toml | 3 + backend/src/api/config_file.rs | 21 +- backend/src/api/endpoints/v1_api_config.rs | 162 ++++- backend/src/api/main_api.rs | 16 +- backend/src/api/model/app_state.rs | 42 +- backend/src/api/model/mod.rs | 3 +- backend/src/api/model/provider_dns_manager.rs | 304 ++++++++++ backend/src/api/panel_api.rs | 5 +- backend/src/api/setup_api.rs | 10 +- backend/src/model/config/input.rs | 8 +- backend/src/model/config/source.rs | 232 +++++++- backend/src/utils/file/config_reader.rs | 51 +- backend/src/utils/network/request.rs | 554 +++++++++++++++--- backend/tests/provider_dns_https_sni.rs | 111 ++++ shared/src/model/config/input.rs | 194 ++++++ shared/src/utils/serde_utils.rs | 4 +- 19 files changed, 1749 insertions(+), 169 deletions(-) create mode 100644 backend/src/api/model/provider_dns_manager.rs create mode 100644 backend/tests/provider_dns_https_sni.rs diff --git a/.gitignore b/.gitignore index 2c7537d02..318a0878a 100644 --- a/.gitignore +++ b/.gitignore @@ -21,6 +21,7 @@ docker/binaries AGENTS.md PLANS.md STEPS.md +CLAUDE.md *.m3u diff --git a/Cargo.lock b/Cargo.lock index 54d6c4707..17bc15e47 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -154,6 +154,45 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" +[[package]] +name = "asn1-rs" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "56624a96882bb8c26d61312ae18cb45868e5a9992ea73c58e45c3101e56a1e60" +dependencies = [ + "asn1-rs-derive", + "asn1-rs-impl", + "displaydoc", + "nom 7.1.3", + "num-traits", + "rusticata-macros", + "thiserror 2.0.18", + "time", +] + +[[package]] +name = "asn1-rs-derive" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3109e49b1e4909e9db6515a30c633684d68cdeaa252f215214cb4fa1a5bfee2c" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", + "synstructure", +] + +[[package]] +name = "asn1-rs-impl" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b18050c2cd6fe86c3a76584ef5e0baf286d038cda203eb6223df2cc413565f7" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "async-compression" version = "0.4.40" @@ -816,6 +855,20 @@ dependencies = [ "zeroize", ] +[[package]] +name = "der-parser" +version = "10.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07da5016415d5a3c4dd39b11ed26f915f52fc4e0dc197d87908bc916e51bc1a6" +dependencies = [ + "asn1-rs", + "displaydoc", + "nom 7.1.3", + "num-bigint", + "num-traits", + "rusticata-macros", +] + [[package]] name = "deranged" version = "0.5.8" @@ -2430,6 +2483,12 @@ dependencies = [ "walkdir", ] +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + [[package]] name = "miniz_oxide" version = "0.8.9" @@ -2464,6 +2523,16 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2bf50223579dc7cdcfb3bfcacf7069ff68243f8c363f62ffa99cf000a6b9c451" +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", +] + [[package]] name = "nom" version = "8.0.0" @@ -2645,6 +2714,15 @@ dependencies = [ "objc2-core-foundation", ] +[[package]] +name = "oid-registry" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "12f40cff3dde1b6087cc5d5f5d4d65712f34016a03ed60e9c08dcc392736b5b7" +dependencies = [ + "asn1-rs", +] + [[package]] name = "once_cell" version = "1.21.3" @@ -3300,6 +3378,20 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "rcgen" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10b99e0098aa4082912d4c649628623db6aba77335e4f4569ff5083a6448b32e" +dependencies = [ + "pem", + "ring", + "rustls-pki-types", + "time", + "x509-parser", + "yasna", +] + [[package]] name = "redox_syscall" version = "0.5.18" @@ -3471,7 +3563,7 @@ dependencies = [ "either", "enum-iterator", "lazy_static", - "nom", + "nom 8.0.0", "regex", "serde", ] @@ -3533,6 +3625,15 @@ dependencies = [ "semver", ] +[[package]] +name = "rusticata-macros" +version = "4.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "faf0c4a6ece9950b9abdb62b1cfcf2a68b3b67a10ba445b3bb85be2a293d0632" +dependencies = [ + "nom 7.1.3", +] + [[package]] name = "rustix" version = "1.1.4" @@ -3553,6 +3654,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c665f33d38cea657d9614f766881e4d510e0eda4239891eea56b4cadcf01801b" dependencies = [ "aws-lc-rs", + "log", "once_cell", "rustls-pki-types", "rustls-webpki", @@ -4517,12 +4619,14 @@ dependencies = [ "quick-xml", "rand 0.9.2", "rayon", + "rcgen", "regex", "reqwest", "rmp-serde", "rpassword", "rphonetic", "rust-argon2", + "rustls", "serde", "serde-saphyr", "serde_json", @@ -4532,6 +4636,7 @@ dependencies = [ "sysinfo", "tempfile", "tokio", + "tokio-rustls", "tokio-stream", "tokio-util", "tower", @@ -5418,6 +5523,33 @@ version = "0.6.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9edde0db4769d2dc68579893f2306b26c6ecfbe0ef499b013d731b7b9247e0b9" +[[package]] +name = "x509-parser" +version = "0.18.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202" +dependencies = [ + "asn1-rs", + "data-encoding", + "der-parser", + "lazy_static", + "nom 7.1.3", + "oid-registry", + "ring", + "rusticata-macros", + "thiserror 2.0.18", + "time", +] + +[[package]] +name = "yasna" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e17bb3549cc1321ae1296b9cdc2698e2b6cb1992adfa19a8c72e5b7a738f44cd" +dependencies = [ + "time", +] + [[package]] name = "yew" version = "0.22.0" diff --git a/README.md b/README.md index a3024369c..7ff12c6a2 100644 --- a/README.md +++ b/README.md @@ -170,16 +170,18 @@ CLI overrides: ## 1.4 Provider Failover & Rotation -Tuliprox supports robust failover mechanisms for streaming providers. If a provider has multiple URLs defined (or aliases), Tuliprox can automatically -rotate between them in case of failures. +Tuliprox supports robust provider failover and DNS-aware rotation. +If a provider has multiple URLs defined (or aliases), Tuliprox can automatically rotate between them on failures. +Additionally, Tuliprox can periodically resolve provider hostnames to IPs and use those IPs for connection attempts. ### 1.4.1 `provider://` Scheme -You can use the special `provider:///...` URL scheme in your configurations. Tuliprox will automatically resolve this to the current -active URL of the specified provider. +You can use the special `provider:///...` URL scheme in your configurations. Tuliprox will automatically +resolve this to the current active URL or IP address of the specified provider. -- If the current URL fails (e.g., 5xx error, timeout), Tuliprox automatically rotates to the next available URL for that provider. -- It tracks failures and prevents infinite loops by limiting attempts to the number of available URLs. +- If the current URL | IP Address fails (e.g., 5xx error, timeout/connect error), Tuliprox automatically rotates to the + next available URL | IP Address for that provider. +- It tracks failures and prevents infinite loops by limiting attempts to the number of available URLs|IP Addresses. ### 1.4.2 Automatic Failover triggers @@ -192,7 +194,38 @@ Failover is triggered automatically on: It does **not** trigger on Authentication errors (401, 403), as those usually indicate invalid credentials rather than a server issue. -### 1.4.3 Provider Failover Configuration Example +### 1.4.3 Provider DNS Resolution (IP-Connect) + +Each provider can optionally enable a `dns` block: + +- `enabled` (default: `false`) +- `refresh_secs` (default: `300`, minimum effective value is `10`) +- `prefer`: `system` | `ipv4` | `ipv6` (default: `system`) +- `max_addrs`: limit number of resolved IPs per host +- `schemes`: list of `http` / `https` to which DNS-IP connect applies (default: `["http"]`) +- `keep_vhost` (default: `false`) +- `overrides`: static host -> IP list (used before DNS lookup) +- `on_resolve_error`: `keep_last_good` | `fallback_to_hostname` (default: `keep_last_good`) +- `on_connect_error`: `try_next_ip` | `rotate_provider_url` (default: `try_next_ip`) +- `resolved`: runtime-managed resolved snapshot per host + +Behavior: + +- A background task resolves hostnames from `provider.urls` periodically (`refresh_secs`). +- For HTTP attempts with a resolved IP, Tuliprox connects via IP. +- For HTTPS attempts with a resolved IP, Tuliprox connects via IP while keeping TLS SNI on the original hostname. +- `keep_vhost=false`: `Host` header uses `IP[:port]`. +- `keep_vhost=true`: `Host` header keeps `hostname[:port]`. +- On connect/timeout errors and `on_connect_error=try_next_ip`, Tuliprox tries the next IP for the same host before rotating provider URL. + +### 1.4.4 `dns.resolved` persistence and visibility + +- `dns.resolved` is runtime-managed. +- It is exposed in `GET /api/v1/config`. +- It is also written back to `source.yml` and overwritten on each DNS refresh cycle. +- Save endpoints treat `dns.resolved` as managed data and do not accept it as user-controlled input. + +### 1.4.5 Provider Failover + DNS Configuration Example Define a provider with multiple URLs and reference it from your inputs/sources. Tuliprox will resolve the active URL and rotate to the next entry on failover conditions. @@ -207,6 +240,22 @@ provider: - http://hello.provider.me - http://stable.golden-bridge.con - http://sleep.time.now.net + dns: + enabled: true + refresh_secs: 300 + prefer: ipv4 + schemes: [http, https] + keep_vhost: true + max_addrs: 2 + on_resolve_error: keep_last_good + on_connect_error: try_next_ip + overrides: + stable.golden-bridge.con: + - 203.0.113.10 + # runtime-managed, written by tuliprox: + resolved: + hello.provider.me: + - 203.0.113.20 inputs: - name: my_input type: xtream_batch diff --git a/backend/Cargo.toml b/backend/Cargo.toml index bcc15bd09..e096ad99c 100644 --- a/backend/Cargo.toml +++ b/backend/Cargo.toml @@ -94,6 +94,9 @@ vergen = { version = "9.1.0", features = ["build"] } [dev-dependencies] http-body-util = "0.1.3" tokio = { version = "1.49.0", features = ["test-util"] } +rcgen = "0.14.5" +rustls = "0.23.35" +tokio-rustls = "0.26.4" [features] proxy-auth-regression = [] diff --git a/backend/src/api/config_file.rs b/backend/src/api/config_file.rs index d7c13c942..ff44ec1c4 100644 --- a/backend/src/api/config_file.rs +++ b/backend/src/api/config_file.rs @@ -55,22 +55,22 @@ impl ConfigFile { // ----------------------------------------------------------------- /// Load and merge global templates using the current app state. - fn load_prepared_global_templates( + async fn load_prepared_global_templates( app_state: &Arc, ) -> Result>, TuliproxError> { let paths = app_state.app_config.paths.load(); let config = app_state.app_config.config.load(); - Self::load_prepared_global_templates_with_config(&paths, &config) + Self::load_prepared_global_templates_with_config(&paths, &config).await } /// Load and merge global templates using explicitly-provided config/paths. /// Used in the prepare phase when the new config has not yet been applied to `app_state`. - fn load_prepared_global_templates_with_config( + async fn load_prepared_global_templates_with_config( paths: &ConfigPaths, config: &Config, ) -> Result>, TuliproxError> { let sources_inline_templates = - read_sources_file(paths.sources_file_path.as_str(), false, false, None, None)?.templates; + read_sources_file(paths.sources_file_path.as_str(), false, false, None, None).await?.templates; // Use robust fallbacks for mapping and template paths let (effective_template_path, effective_mapping_path) = resolve_template_and_mapping_paths(paths, config.template_path.as_deref(), config.mapping_path.as_deref()); @@ -136,8 +136,8 @@ impl ConfigFile { Ok(()) } - fn load_mapping(app_state: &Arc) -> Result<(), TuliproxError> { - let prepared_templates = Self::load_prepared_global_templates(app_state)?; + async fn load_mapping(app_state: &Arc) -> Result<(), TuliproxError> { + let prepared_templates = Self::load_prepared_global_templates(app_state).await?; Self::load_mapping_with_templates(app_state, prepared_templates.as_deref()) } @@ -151,14 +151,15 @@ impl ConfigFile { paths: &ConfigPaths, ) -> Result { let sources_file = paths.sources_file_path.clone(); - let prepared_templates = Self::load_prepared_global_templates_with_config(paths, config)?; + let prepared_templates = Self::load_prepared_global_templates_with_config(paths, config).await?; let mut sources_dto = read_sources_file_from_path_with_templates( &PathBuf::from(sources_file.as_str()), true, true, config.get_hdhr_device_overview().as_ref(), prepared_templates.as_deref(), - )?; + ) + .await?; prepare_sources_batch(&mut sources_dto, true).await?; let sources: SourcesConfig = SourcesConfig::try_from(sources_dto)?; let prepared_mapping = @@ -253,7 +254,7 @@ impl ConfigFile { PreparedFollowUp::Sources(prepared) } else if mapping_changed { // Only mapping path changed; templates are the same → load templates once. - let prepared_templates = Self::load_prepared_global_templates_with_config(&effective_paths, &config)?; + let prepared_templates = Self::load_prepared_global_templates_with_config(&effective_paths, &config).await?; let prepared = Self::prepare_mapping_reload( effective_paths.mapping_file_path.as_deref(), prepared_templates.as_deref(), @@ -331,7 +332,7 @@ impl ConfigFile { app_state.event_manager.send_event(EventMessage::ConfigChange(ConfigType::ApiProxy)); } ConfigFile::Mapping => { - ConfigFile::load_mapping(app_state)?; + ConfigFile::load_mapping(app_state).await?; app_state.event_manager.send_event(EventMessage::ConfigChange(ConfigType::Mapping)); } ConfigFile::Template | ConfigFile::Sources => { diff --git a/backend/src/api/endpoints/v1_api_config.rs b/backend/src/api/endpoints/v1_api_config.rs index 538201037..7f3889e8d 100644 --- a/backend/src/api/endpoints/v1_api_config.rs +++ b/backend/src/api/endpoints/v1_api_config.rs @@ -20,6 +20,37 @@ use shared::{ }; use std::sync::Arc; +fn inject_provider_dns_resolved(sources_dto: &mut SourcesConfigDto, runtime_sources: &crate::model::SourcesConfig) { + let Some(provider_dtos) = sources_dto.provider.as_mut() else { + return; + }; + for provider_dto in provider_dtos { + let Some(runtime_provider) = runtime_sources.get_provider_by_name(provider_dto.name.as_ref()) else { + continue; + }; + let Some(dns_dto) = provider_dto.dns.as_mut() else { + continue; + }; + if runtime_provider.get_dns_config().is_none() { + dns_dto.resolved = None; + continue; + } + let snapshot = runtime_provider.snapshot_resolved(); + dns_dto.resolved = (!snapshot.is_empty()).then_some(snapshot); + } +} + +fn strip_provider_dns_resolved(sources_dto: &mut SourcesConfigDto) { + let Some(provider_dtos) = sources_dto.provider.as_mut() else { + return; + }; + for provider_dto in provider_dtos { + if let Some(dns_dto) = provider_dto.dns.as_mut() { + dns_dto.resolved = None; + } + } +} + pub(in crate::api::endpoints) async fn intern_save_config_api_proxy( backup_dir: &str, api_proxy: &ApiProxyConfigDto, @@ -78,9 +109,12 @@ async fn save_config_main( async fn save_config_sources( axum::extract::State(app_state): axum::extract::State>, - axum::extract::Json(sources): axum::extract::Json, + axum::extract::Json(mut sources): axum::extract::Json, ) -> impl axum::response::IntoResponse + Send { - let templates_to_persist = match utils::validate_source_config_for_persist(&app_state, &sources) { + // `dns.resolved` is runtime-managed and must not be accepted from API input. + strip_provider_dns_resolved(&mut sources); + + let templates_to_persist = match utils::validate_source_config_for_persist(&app_state, &sources).await { Ok(value) => value, Err(err) => { error!("Failed to validate source.yml {err}"); @@ -179,7 +213,7 @@ async fn save_config_api_proxy_config( async fn config(axum::extract::State(app_state): axum::extract::State>) -> impl IntoResponse + Send { let paths = app_state.app_config.paths.load(); - match utils::read_app_config_dto(&paths, true, false) { + match utils::read_app_config_dto(&paths, true, false).await { Ok(mut app_config) => { if let Err(err) = prepare_sources_batch(&mut app_config.sources, false).await { error!("Failed to prepare sources batch: {err}"); @@ -188,6 +222,8 @@ async fn config(axum::extract::State(app_state): axum::extract::State>) -> axum::Router().expect("ip parse should work")], + ); + let runtime_sources = SourcesConfig { + provider: vec![runtime_provider], + ..SourcesConfig::default() + }; + + inject_provider_dns_resolved(&mut dto, &runtime_sources); + + let resolved = dto.provider.as_ref() + .and_then(|providers| providers.first()) + .and_then(|provider| provider.dns.as_ref()) + .and_then(|dns| dns.resolved.as_ref()) + .expect("resolved dns snapshot should be present"); + assert_eq!( + resolved.get("example.com"), + Some(&vec!["203.0.113.10".parse::().expect("ip parse should work")]) + ); + } + + #[test] + fn inject_provider_dns_resolved_clears_value_when_runtime_dns_disabled() { + let mut dto = SourcesConfigDto { + provider: Some(vec![ConfigProviderDto { + name: "p1".into(), + urls: vec!["http://example.com".into()], + dns: Some(ProviderDnsDto { + enabled: true, + resolved: Some(HashMap::from([( + "example.com".to_string(), + vec!["203.0.113.10".parse::().expect("ip parse should work")], + )])), + ..ProviderDnsDto::default() + }), + }]), + ..SourcesConfigDto::default() + }; + + let runtime_provider = Arc::new(ConfigProvider::from(&ConfigProviderDto { + name: "p1".into(), + urls: vec!["http://example.com".into()], + dns: None, + })); + let runtime_sources = SourcesConfig { + provider: vec![runtime_provider], + ..SourcesConfig::default() + }; + + inject_provider_dns_resolved(&mut dto, &runtime_sources); + + let resolved = dto.provider.as_ref() + .and_then(|providers| providers.first()) + .and_then(|provider| provider.dns.as_ref()) + .and_then(|dns| dns.resolved.as_ref()); + assert!(resolved.is_none(), "resolved output must be empty when runtime dns is disabled"); + } + + #[test] + fn strip_provider_dns_resolved_removes_payload_values() { + let mut dto = SourcesConfigDto { + provider: Some(vec![ConfigProviderDto { + name: "p1".into(), + urls: vec!["http://example.com".into()], + dns: Some(ProviderDnsDto { + enabled: true, + resolved: Some(HashMap::from([( + "example.com".to_string(), + vec!["203.0.113.10".parse::().expect("ip parse should work")], + )])), + ..ProviderDnsDto::default() + }), + }]), + ..SourcesConfigDto::default() + }; + + strip_provider_dns_resolved(&mut dto); + + let resolved = dto + .provider + .as_ref() + .and_then(|providers| providers.first()) + .and_then(|provider| provider.dns.as_ref()) + .and_then(|dns| dns.resolved.as_ref()); + assert!(resolved.is_none(), "resolved must be stripped from incoming payload"); + } +} diff --git a/backend/src/api/main_api.rs b/backend/src/api/main_api.rs index ce19f8967..5fe6d7a60 100644 --- a/backend/src/api/main_api.rs +++ b/backend/src/api/main_api.rs @@ -16,9 +16,9 @@ use crate::{ hdhomerun_proprietary::spawn_proprietary_tasks, hdhomerun_ssdp::spawn_ssdp_discover_task, model::{ - create_cache, create_http_client, create_http_client_no_redirect, ActiveProviderManager, ActiveUserManager, - AppState, CancelTokens, ConnectionManager, DownloadQueue, EventManager, EventMessage, HdHomerunAppState, - MetadataUpdateManager, PlaylistStorageState, SharedStreamManager, UpdateGuard, + create_cache, create_http_client, create_http_client_no_redirect, exec_provider_dns, ActiveProviderManager, + ActiveUserManager, AppState, CancelTokens, ConnectionManager, DownloadQueue, EventManager, EventMessage, + HdHomerunAppState, MetadataUpdateManager, PlaylistStorageState, SharedStreamManager, UpdateGuard, }, panel_api::sync_panel_api_exp_dates_on_boot, scheduler::{exec_interner_prune, exec_scheduler}, @@ -287,9 +287,14 @@ pub async fn start_server(app_config: Arc, targets: Arc, targets: Arc, } -create_bitset!(u8, UpdateChangesFlags, Scheduler, Hdhomerun, FileWatch, Geoip); +create_bitset!(u8, UpdateChangesFlags, Scheduler, Hdhomerun, FileWatch, Geoip, ProviderDns); pub(in crate::api) struct UpdateChanges { flags: UpdateChangesFlagsSet, @@ -149,8 +150,14 @@ fn cancel_services(app_state: &Arc, changes: &UpdateChanges) { let scheduler = cancel_service!(scheduler, UpdateChangesFlags::Scheduler, changes, cancel_tokens); let hdhomerun = cancel_service!(hdhomerun, UpdateChangesFlags::Hdhomerun, changes, cancel_tokens); let file_watch = cancel_service!(file_watch, UpdateChangesFlags::FileWatch, changes, cancel_tokens); + let provider_dns = cancel_service!(provider_dns, UpdateChangesFlags::ProviderDns, changes, cancel_tokens); - let tokens = CancelTokens { scheduler, hdhomerun, file_watch }; + let tokens = CancelTokens { + scheduler, + hdhomerun, + file_watch, + provider_dns, + }; app_state.cancel_tokens.store(Arc::new(tokens)); } @@ -181,6 +188,10 @@ fn start_services(app_state: &Arc, changes: &UpdateChanges) { if changes.flags.contains(UpdateChangesFlags::FileWatch) { exec_config_watch(app_state, &app_state.cancel_tokens.load().file_watch); } + + if changes.flags.contains(UpdateChangesFlags::ProviderDns) { + exec_provider_dns(app_state, &app_state.cancel_tokens.load().provider_dns); + } } /// Creates the default HTTP client. @@ -285,6 +296,7 @@ pub struct CancelTokens { pub(crate) scheduler: CancellationToken, pub(crate) hdhomerun: CancellationToken, pub(crate) file_watch: CancellationToken, + pub(crate) provider_dns: CancellationToken, } impl Default for CancelTokens { fn default() -> Self { @@ -292,6 +304,7 @@ impl Default for CancelTokens { scheduler: CancellationToken::new(), hdhomerun: CancellationToken::new(), file_watch: CancellationToken::new(), + provider_dns: CancellationToken::new(), } } } @@ -425,9 +438,10 @@ impl AppState { } fn detect_changes_for_sources(&self, sources: &SourcesConfig) -> UpdateChanges { - let (file_watch_changed, target_changes) = { + let (file_watch_changed, provider_dns_changed, target_changes) = { let old_sources = self.app_config.sources.load(); let file_watch_changed = old_sources.get_input_files() != sources.get_input_files(); + let provider_dns_changed = providers_changed(&old_sources.provider, &sources.provider); let mut target_changes = HashMap::new(); for source in &old_sources.sources { @@ -477,11 +491,12 @@ impl AppState { } } - (file_watch_changed, target_changes) + (file_watch_changed, provider_dns_changed, target_changes) }; let mut changes = UpdateChanges { flags: UpdateChangesFlagsSet::new(), targets: Some(target_changes) }; changes.set_flag_if(file_watch_changed, UpdateChangesFlags::FileWatch); + changes.set_flag_if(provider_dns_changed, UpdateChangesFlags::ProviderDns); changes } @@ -590,6 +605,21 @@ fn hdhomerun_changed(a: &HdHomeRunConfig, b: &HdHomeRunConfig) -> bool { fn string_changed(a: &str, b: &str) -> bool { a != b } +fn providers_changed(a: &[Arc], b: &[Arc]) -> bool { + if a.len() != b.len() { + return true; + } + for lhs in a { + let Some(rhs) = b.iter().find(|candidate| candidate.name == lhs.name) else { + return true; + }; + if lhs.urls != rhs.urls || lhs.dns != rhs.dns { + return true; + } + } + false +} + #[derive(Clone)] pub struct HdHomerunAppState { pub app_state: Arc, diff --git a/backend/src/api/model/mod.rs b/backend/src/api/model/mod.rs index 550cd2551..9ef3133d7 100644 --- a/backend/src/api/model/mod.rs +++ b/backend/src/api/model/mod.rs @@ -8,6 +8,7 @@ mod metadata_update_manager; mod model_utils; mod playlist_mem_cache; mod provider_config; +mod provider_dns_manager; mod provider_lineup_manager; mod request; mod stream; @@ -19,7 +20,7 @@ mod xtream; pub(crate) use self::streams::*; pub use self::{ active_provider_manager::*, app_state::*, connection_manager::*, event_manager::*, metadata_update_manager::*, - playlist_mem_cache::*, provider_lineup_manager::*, stream::*, update_guard::*, + playlist_mem_cache::*, provider_dns_manager::*, provider_lineup_manager::*, stream::*, update_guard::*, }; pub(in crate::api) use self::{ active_user_manager::*, download::*, model_utils::*, provider_config::*, request::*, stream_error::*, xtream::*, diff --git a/backend/src/api/model/provider_dns_manager.rs b/backend/src/api/model/provider_dns_manager.rs new file mode 100644 index 000000000..bcc08e341 --- /dev/null +++ b/backend/src/api/model/provider_dns_manager.rs @@ -0,0 +1,304 @@ +use crate::api::model::AppState; +use crate::model::ConfigProvider; +use crate::utils::read_sources_file_from_path; +use log::{debug, warn}; +use shared::model::{DnsPrefer, OnResolveErrorPolicy, SourcesConfigDto}; +use std::collections::HashSet; +use std::io; +use std::net::IpAddr; +use std::path::{Path, PathBuf}; +use std::sync::Arc; +use std::time::{Duration, Instant}; +use tokio::fs; +use tokio::net::lookup_host; +use tokio_util::sync::CancellationToken; + +fn filter_by_preference(ips: Vec, prefer: DnsPrefer) -> Vec { + match prefer { + DnsPrefer::System => ips, + DnsPrefer::Ipv4 => ips.into_iter().filter(IpAddr::is_ipv4).collect(), + DnsPrefer::Ipv6 => ips.into_iter().filter(IpAddr::is_ipv6).collect(), + } +} + +fn dedup_keep_order(ips: Vec) -> Vec { + let mut seen = HashSet::new(); + ips.into_iter().filter(|ip| seen.insert(*ip)).collect() +} + +async fn resolve_hostname(hostname: &str, prefer: DnsPrefer, max_addrs: Option) -> std::io::Result> { + let addrs = lookup_host((hostname, 0)).await?; + let mut ips: Vec = addrs.map(|addr| addr.ip()).collect(); + ips = dedup_keep_order(ips); + ips = filter_by_preference(ips, prefer); + if let Some(max) = max_addrs.filter(|max| *max > 0) { + ips.truncate(max); + } + Ok(ips) +} + +#[derive(Debug, Default)] +struct ProviderResolveStats { + total: usize, + overridden: usize, + resolved: usize, + empty: usize, + failed: usize, +} + +async fn resolve_provider(provider: &Arc) -> ProviderResolveStats { + let mut stats = ProviderResolveStats::default(); + let Some(dns_cfg) = provider.get_dns_config().cloned() else { + return stats; + }; + if !dns_cfg.enabled { + return stats; + } + + let hostnames = provider.hostnames_from_urls(); + stats.total = hostnames.len(); + if hostnames.is_empty() { + debug!( + "Provider dns task '{}' found no hostname URLs to resolve (urls={:?})", + provider.name, provider.urls + ); + return stats; + } + + for host in hostnames { + if let Some(overridden) = dns_cfg.overrides.get(&host) { + provider.store_resolved(&host, overridden.clone()); + stats.overridden += 1; + stats.resolved += 1; + debug!("Provider dns '{}' host '{}' resolved from override: {:?}", provider.name, host, overridden); + continue; + } + + match resolve_hostname(&host, dns_cfg.prefer, dns_cfg.max_addrs).await { + Ok(ips) if !ips.is_empty() => { + debug!("Provider dns '{}' host '{}' resolved: {:?}", provider.name, host, ips); + provider.store_resolved(&host, ips); + stats.resolved += 1; + } + Ok(_) => { + stats.empty += 1; + provider.mark_resolve_error(&host, "DNS resolution returned no addresses"); + if dns_cfg.on_resolve_error == OnResolveErrorPolicy::FallbackToHostname { + provider.clear_resolved(&host); + } + warn!( + "Provider dns '{}' host '{}' returned empty address set (policy={:?})", + provider.name, host, dns_cfg.on_resolve_error + ); + } + Err(err) => { + stats.failed += 1; + provider.mark_resolve_error(&host, err.to_string()); + if dns_cfg.on_resolve_error == OnResolveErrorPolicy::FallbackToHostname { + provider.clear_resolved(&host); + } + warn!("provider dns resolve failed for '{}' host '{}': {err}", provider.name, host); + } + } + } + + stats +} + +fn serialize_sources_for_persist(sources: &SourcesConfigDto) -> io::Result { + let mut serialized = String::new(); + let options = serde_saphyr::SerializerOptions { + prefer_block_scalars: false, + ..Default::default() + }; + serde_saphyr::to_fmt_writer_with_options(&mut serialized, sources, options) + .map_err(|err| io::Error::other(format!("Could not serialize source.yml: {err}")))?; + Ok(serialized) +} + +async fn write_sources_file_force(path: &Path, sources: &SourcesConfigDto) -> io::Result<()> { + let serialized = serialize_sources_for_persist(sources)?; + let parent_dir = path.parent().ok_or_else(|| { + io::Error::other(format!( + "Could not write source.yml '{}': missing parent directory", + path.display() + )) + })?; + let dest_file_name = path.file_name().and_then(|s| s.to_str()).unwrap_or("source.yml"); + let mut tmp_path = parent_dir.to_path_buf(); + tmp_path.push(format!( + ".{dest_file_name}.tmp-{}-{}", + std::process::id(), + chrono::Local::now().timestamp_nanos_opt().unwrap_or_default() + )); + + fs::write(&tmp_path, serialized).await?; + match fs::rename(&tmp_path, path).await { + Ok(()) => Ok(()), + Err(err) => { + #[cfg(windows)] + { + // Try to rename again after removing destination (Windows often needs this) + if let Ok(()) = fs::remove_file(path).await { + if fs::rename(&tmp_path, path).await.is_ok() { + return Ok(()); + } + // Rename still failed after removing dest - fall through to clean up + } + } + let _ = fs::remove_file(&tmp_path).await; + Err(io::Error::other(format!( + "Could not replace '{}' with '{}': {err}", + path.display(), + tmp_path.display() + ))) + } + } +} + +async fn persist_provider_resolved_to_source_file(app_state: &Arc, provider: &Arc) { + let source_file = { + let paths = app_state.app_config.paths.load(); + paths.sources_file_path.clone() + }; + let source_path = PathBuf::from(&source_file); + let _lock = app_state.app_config.file_locks.write_lock(&source_path).await; + + let mut sources_dto = match read_sources_file_from_path(&source_path, false, false, None).await { + Ok(dto) => dto, + Err(err) => { + warn!( + "Provider dns '{}' failed to read source.yml '{}': {err}", + provider.name, + source_path.display() + ); + return; + } + }; + + let Some(provider_dtos) = sources_dto.provider.as_mut() else { + debug!( + "Provider dns '{}' source.yml '{}' has no provider section to persist resolved values", + provider.name, + source_path.display() + ); + return; + }; + + let Some(provider_dto) = provider_dtos.iter_mut().find(|dto| dto.name.as_ref() == provider.name.as_ref()) else { + warn!( + "Provider dns '{}' not found in source.yml '{}', cannot persist resolved values", + provider.name, + source_path.display() + ); + return; + }; + + let Some(dns_dto) = provider_dto.dns.as_mut() else { + debug!( + "Provider dns '{}' has no dns section in source.yml '{}', skipping resolved persist", + provider.name, + source_path.display() + ); + return; + }; + + let resolved_hosts = { + let snapshot = provider.snapshot_resolved(); + dns_dto.resolved = (!snapshot.is_empty()).then_some(snapshot); + dns_dto.resolved.as_ref().map_or(0, std::collections::HashMap::len) + }; + + match write_sources_file_force(&source_path, &sources_dto).await { + Ok(()) => { + debug!( + "Provider dns '{}' persisted dns.resolved to '{}' (hosts={resolved_hosts})", + provider.name, + source_path.display() + ); + } + Err(err) => { + warn!( + "Provider dns '{}' failed to persist dns.resolved to '{}': {err}", + provider.name, + source_path.display() + ); + } + } +} + +fn spawn_provider_dns_task(app_state: Arc, provider_name: Arc, cancel: CancellationToken) { + tokio::spawn(async move { + let mut refresh_secs = 300_u64; + { + let sources = app_state.app_config.sources.load(); + if let Some(provider) = sources.get_provider_by_name(provider_name.as_ref()) { + refresh_secs = provider.get_dns_config().map_or(300, |dns| dns.refresh_secs.max(10)); + } + } + // Add initial jitter (0-10% of refresh interval) to prevent thundering herd + let jitter_ms = (rand::random::() % (refresh_secs * 100)).min(5000); + tokio::time::sleep(Duration::from_millis(jitter_ms)).await; + debug!("Starting provider dns task for '{provider_name}' (refresh={refresh_secs}s)"); + loop { + tokio::select! { + () = cancel.cancelled() => { + debug!("Stopping provider dns task for '{provider_name}'"); + break; + } + () = async { + let start = Instant::now(); + debug!("Provider dns tick '{provider_name}' started"); + let provider = { + let sources = app_state.app_config.sources.load(); + sources.get_provider_by_name(provider_name.as_ref()).cloned() + }; + let Some(provider) = provider else { + warn!("Provider dns '{provider_name}' not found in runtime sources, retrying"); + tokio::time::sleep(Duration::from_secs(30)).await; + return; + }; + refresh_secs = provider.get_dns_config().map_or(300, |dns| dns.refresh_secs.max(10)); + let stats = resolve_provider(&provider).await; + persist_provider_resolved_to_source_file(&app_state, &provider).await; + let cache_hosts = provider.snapshot_resolved().len(); + debug!( + "Provider dns tick '{}' finished: total_hosts={} resolved={} overridden={} empty={} failed={} cache_hosts={} elapsed_ms={}", + provider.name, + stats.total, + stats.resolved, + stats.overridden, + stats.empty, + stats.failed, + cache_hosts, + start.elapsed().as_millis(), + ); + debug!("Provider dns '{}' next tick in {}s", provider.name, refresh_secs); + tokio::time::sleep(Duration::from_secs(refresh_secs)).await; + } => {} + } + } + }); +} + +pub fn exec_provider_dns(app_state: &Arc, cancel: &CancellationToken) { + let sources = app_state.app_config.sources.load(); + let provider_names: Vec<_> = sources + .provider + .iter() + .filter(|provider| provider.get_dns_config().is_some_and(|dns| dns.enabled)) + .map(|provider| provider.name.clone()) + .collect(); + drop(sources); + + if provider_names.is_empty() { + debug!("Provider dns manager: no enabled providers found"); + return; + } + + debug!("Provider dns manager: starting {} provider task(s)", provider_names.len()); + + for provider_name in provider_names { + spawn_provider_dns_task(Arc::clone(app_state), provider_name, cancel.clone()); + } +} diff --git a/backend/src/api/panel_api.rs b/backend/src/api/panel_api.rs index f6fd1cd33..11f78ec00 100644 --- a/backend/src/api/panel_api.rs +++ b/backend/src/api/panel_api.rs @@ -1120,7 +1120,7 @@ async fn patch_source_yml_add_alias( password: &str, exp_date: Option, ) -> Result<(), TuliproxError> { - let mut sources = match read_sources_file_from_path(source_file_path, false, false, None) { + let mut sources = match read_sources_file_from_path(source_file_path, false, false, None).await { Ok(sources) => sources, Err(e) => return info_err_res!("panel_api: failed to read source file: {e}"), }; @@ -1416,6 +1416,7 @@ async fn persist_sources_yml_with_patches( return Ok(false); } let mut sources = read_sources_file_from_path(sources_path, false, false, None) + .await .map_err(|e| info_err!("panel_api: failed to read source file: {e}"))?; let changed = apply_sources_yml_patches(&mut sources, patches)?; @@ -1434,7 +1435,7 @@ async fn patch_source_yml_update_exp_date( account_name: &Arc, exp_date: i64, ) -> Result<(), TuliproxError> { - let mut sources = match read_sources_file_from_path(source_file_path, false, false, None) { + let mut sources = match read_sources_file_from_path(source_file_path, false, false, None).await { Ok(sources) => sources, Err(e) => return info_err_res!("panel_api: failed to read source file: {e}"), }; diff --git a/backend/src/api/setup_api.rs b/backend/src/api/setup_api.rs index 51cfd07b8..51015c945 100644 --- a/backend/src/api/setup_api.rs +++ b/backend/src/api/setup_api.rs @@ -203,7 +203,7 @@ fn create_default_draft() -> AppConfigDto { } } -fn build_initial_draft(paths: &ConfigPaths) -> AppConfigDto { +async fn build_initial_draft(paths: &ConfigPaths) -> AppConfigDto { let mut draft = create_default_draft(); if file_exists(&paths.config_file_path) { @@ -219,8 +219,10 @@ fn build_initial_draft(paths: &ConfigPaths) -> AppConfigDto { false, false, draft.config.get_hdhr_device_overview().as_ref(), - None, - ) { + None + ) + .await + { Ok(src) => draft.sources = src, Err(err) => warn!("Setup mode: failed to load existing source.yml: {err}"), } @@ -980,7 +982,7 @@ fn create_compression_layer() -> tower_http::compression::CompressionLayer { } pub async fn start_setup_server(paths: &ConfigPaths, missing_files: &[String]) -> Result<(), TuliproxError> { - let draft = build_initial_draft(paths); + let draft = build_initial_draft(paths).await; let (host, port, web_root) = setup_bind_values(&draft); let web_dir = resolve_setup_web_dir(&web_root) .ok_or_else(|| info_err!("Setup mode requires a web directory. Tried '{}'", web_root.display(),))?; diff --git a/backend/src/model/config/input.rs b/backend/src/model/config/input.rs index a987b7cd2..a6ed18f97 100644 --- a/backend/src/model/config/input.rs +++ b/backend/src/model/config/input.rs @@ -619,8 +619,8 @@ fn assemble_provider_url(provider: &ConfigProvider, path_and_query: &str) -> Res mod tests { use super::*; use crate::model::ConfigProvider; + use shared::model::ConfigProviderDto; use std::borrow::Cow; - use std::sync::atomic::AtomicUsize; use std::sync::Arc; #[test] @@ -636,11 +636,11 @@ mod tests { #[test] fn test_resolve_url_provider() { - let provider = ConfigProvider { + let provider = ConfigProvider::from(&ConfigProviderDto { name: "myprovider".into(), urls: vec!["http://provider.com".into()], - current_url_index: AtomicUsize::new(0), - }; + dns: None, + }); let input = ConfigInput { name: "test_input".into(), provider_configs: Some(vec![Arc::new(provider)]), diff --git a/backend/src/model/config/source.rs b/backend/src/model/config/source.rs index ba9ed45a2..cfc408d9c 100644 --- a/backend/src/model/config/source.rs +++ b/backend/src/model/config/source.rs @@ -1,17 +1,166 @@ use crate::model::{macros, ConfigInput, ConfigTarget, ProcessTargets}; +use parking_lot::RwLock; use shared::error::{info_err_res, TuliproxError}; -use shared::model::{ConfigProviderDto, ConfigSourceDto, PatternTemplate, SourcesConfigDto}; +use shared::model::{ + ConfigProviderDto, ConfigSourceDto, DnsPrefer, DnsScheme, OnConnectErrorPolicy, OnResolveErrorPolicy, PatternTemplate, + SourcesConfigDto, +}; use std::borrow::Cow; use std::collections::{HashMap, HashSet}; +use std::net::IpAddr; use std::path::PathBuf; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; +use std::time::SystemTime; +use url::Url; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ProviderDnsConfig { + pub enabled: bool, + pub refresh_secs: u64, + pub prefer: DnsPrefer, + pub max_addrs: Option, + pub schemes: Vec, + pub keep_vhost: bool, + pub overrides: HashMap>, + pub on_resolve_error: OnResolveErrorPolicy, + pub on_connect_error: OnConnectErrorPolicy, +} + +impl ProviderDnsConfig { + pub fn supports_scheme(&self, scheme: &str) -> bool { + if !self.enabled { + return false; + } + match scheme.to_ascii_lowercase().as_str() { + "http" => self.schemes.contains(&DnsScheme::Http), + "https" => self.schemes.contains(&DnsScheme::Https), + _ => false, + } + } +} + +impl From<&shared::model::ProviderDnsDto> for ProviderDnsConfig { + fn from(dto: &shared::model::ProviderDnsDto) -> Self { + Self { + enabled: dto.enabled, + refresh_secs: dto.refresh_secs, + prefer: dto.prefer, + max_addrs: dto.max_addrs, + schemes: dto + .schemes + .as_ref() + .filter(|list| !list.is_empty()) + .cloned() + .unwrap_or_else(|| vec![DnsScheme::Http, DnsScheme::Https]), + keep_vhost: dto.keep_vhost, + overrides: dto.overrides.clone().unwrap_or_default(), + on_resolve_error: dto.on_resolve_error, + on_connect_error: dto.on_connect_error, + } + } +} + +#[derive(Debug, Default)] +pub struct ProviderDnsCacheEntry { + pub ips: Vec, + pub rr_index: AtomicUsize, + pub last_ok: Option, + pub last_err: Option, +} + +impl Clone for ProviderDnsCacheEntry { + fn clone(&self) -> Self { + Self { + ips: self.ips.clone(), + rr_index: AtomicUsize::new(self.rr_index.load(Ordering::Relaxed)), + last_ok: self.last_ok, + last_err: self.last_err.clone(), + } + } +} + +#[derive(Debug, Default)] +pub struct ProviderDnsCache { + by_host: RwLock>, +} + +impl ProviderDnsCache { + pub fn select_ip_from(&self, host: &str, ips: &[IpAddr]) -> Option { + if ips.is_empty() { + return None; + } + let len = ips.len(); + let mut guard = self.by_host.write(); + let entry = guard.entry(host.to_ascii_lowercase()).or_default(); + // `% len` guards against a stale index when the override list changed length. + let idx = entry.rr_index.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |i| Some((i + 1) % len)).unwrap_or_else(|i| i) % len; + Some(ips[idx]) + } + + pub fn select_cached_ip(&self, host: &str) -> Option { + let guard = self.by_host.read(); + let entry = guard.get(&host.to_ascii_lowercase())?; + if entry.ips.is_empty() { + return None; + } + let len = entry.ips.len(); + let idx = entry.rr_index.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |i| Some((i + 1) % len)).unwrap_or_else(|i| i) % len; + Some(entry.ips[idx]) + } + + pub fn store_resolved(&self, host: &str, ips: Vec) { + let mut guard = self.by_host.write(); + let entry = guard.entry(host.to_ascii_lowercase()).or_default(); + let new_len = ips.len(); + entry.ips = ips; + if new_len == 0 { + entry.rr_index.store(0, Ordering::Relaxed); + } else { + entry.rr_index.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |i| Some(i % new_len)).ok(); + } + entry.last_ok = Some(SystemTime::now()); + entry.last_err = None; + } + + pub fn clear_resolved(&self, host: &str) { + let mut guard = self.by_host.write(); + if let Some(entry) = guard.get_mut(&host.to_ascii_lowercase()) { + entry.ips.clear(); + entry.rr_index.store(0, Ordering::Relaxed); + entry.last_ok = None; + } + } + + pub fn mark_resolve_error(&self, host: &str, err: impl Into) { + let mut guard = self.by_host.write(); + let entry = guard.entry(host.to_ascii_lowercase()).or_default(); + entry.last_err = Some(err.into()); + } + + pub fn snapshot_resolved(&self) -> HashMap> { + let guard = self.by_host.read(); + guard + .iter() + .filter_map(|(host, entry)| (!entry.ips.is_empty()).then_some((host.clone(), entry.ips.clone()))) + .collect() + } + + pub fn ip_count(&self, host: &str) -> usize { + let guard = self.by_host.read(); + guard + .get(&host.to_ascii_lowercase()) + .map_or(0, |entry| entry.ips.len()) + } +} #[derive(Debug)] pub struct ConfigProvider { pub name: Arc, pub urls: Vec>, pub current_url_index: AtomicUsize, + pub dns: Option, + pub dns_cache: Arc, } impl Clone for ConfigProvider { @@ -20,6 +169,8 @@ impl Clone for ConfigProvider { name: self.name.clone(), urls: self.urls.clone(), current_url_index: AtomicUsize::new(self.current_url_index.load(Ordering::Relaxed)), + dns: self.dns.clone(), + dns_cache: Arc::clone(&self.dns_cache), } } } @@ -28,10 +179,21 @@ impl Clone for ConfigProvider { macros::from_impl!(ConfigProvider); impl From<&ConfigProviderDto> for ConfigProvider { fn from(dto: &ConfigProviderDto) -> Self { + let dns_cfg = dto.dns.as_ref().map(ProviderDnsConfig::from); + let dns_cache = Arc::new(ProviderDnsCache::default()); + if let Some(dns_dto) = dto.dns.as_ref() { + if let Some(resolved) = dns_dto.resolved.as_ref() { + for (host, ips) in resolved { + dns_cache.store_resolved(host, ips.clone()); + } + } + } Self { name: dto.name.clone(), urls: dto.urls.clone(), current_url_index: AtomicUsize::new(0), + dns: dns_cfg, + dns_cache, } } } @@ -54,6 +216,63 @@ impl ConfigProvider { self.current_url_index.load(Ordering::Relaxed) } + pub fn get_dns_config(&self) -> Option<&ProviderDnsConfig> { self.dns.as_ref() } + + pub fn dns_enabled_for_scheme(&self, scheme: &str) -> bool { + self.dns.as_ref().is_some_and(|cfg| cfg.supports_scheme(scheme)) + } + + pub fn select_ip_for_host(&self, host: &str) -> Option { + let dns = self.dns.as_ref()?; + let normalized = host.trim().to_ascii_lowercase(); + if let Some(ips) = dns.overrides.get(&normalized) { + return self.dns_cache.select_ip_from(&normalized, ips); + } + self.dns_cache.select_cached_ip(&normalized) + } + + pub fn ip_count_for_host(&self, host: &str) -> usize { + let Some(dns) = self.dns.as_ref() else { + return 0; + }; + let normalized = host.trim().to_ascii_lowercase(); + if let Some(ips) = dns.overrides.get(&normalized) { + return ips.len(); + } + self.dns_cache.ip_count(&normalized) + } + + pub fn hostnames_from_urls(&self) -> HashSet { + let mut hostnames = HashSet::new(); + for raw in &self.urls { + let raw = raw.as_ref(); + let candidate = if raw.contains("://") { + raw.to_string() + } else { + format!("http://{raw}") + }; + let Ok(parsed) = Url::parse(&candidate) else { + continue; + }; + let Some(host) = parsed.host_str() else { + continue; + }; + if host.parse::().is_ok() { + continue; + } + hostnames.insert(host.to_ascii_lowercase()); + } + hostnames + } + + pub fn store_resolved(&self, host: &str, ips: Vec) { self.dns_cache.store_resolved(host, ips); } + + pub fn clear_resolved(&self, host: &str) { self.dns_cache.clear_resolved(host); } + + pub fn mark_resolve_error(&self, host: &str, err: impl Into) { self.dns_cache.mark_resolve_error(host, err); } + + pub fn snapshot_resolved(&self) -> HashMap> { self.dns_cache.snapshot_resolved() } + /// Rotates to next URL, checking if a full cycle has been completed. /// Returns None if we've cycled back to the `start_index`, indicating all URLs were tried. /// @@ -97,13 +316,12 @@ impl ConfigSource { } } -// macros::try_from_impl!(ConfigSource); -impl ConfigSource { - pub fn from_dto(dto: &ConfigSourceDto) -> Result { - Ok(Self { +impl From<&ConfigSourceDto> for ConfigSource { + fn from(dto: &ConfigSourceDto) -> Self { + Self { inputs: dto.inputs.clone(), targets: dto.targets.iter().map(|c| Arc::new(ConfigTarget::from(c))).collect(), - }) + } } } @@ -145,7 +363,7 @@ impl TryFrom<&SourcesConfigDto> for SourcesConfig { return info_err_res!("Source references unknown input: {input_name}"); } } - sources.push(ConfigSource::from_dto(source_dto)?); + sources.push(ConfigSource::from(source_dto)); } Ok(Self { diff --git a/backend/src/utils/file/config_reader.rs b/backend/src/utils/file/config_reader.rs index 325828fd0..2af823ca9 100644 --- a/backend/src/utils/file/config_reader.rs +++ b/backend/src/utils/file/config_reader.rs @@ -86,11 +86,12 @@ pub async fn read_api_proxy_config( } } -fn parse_sources_file_from_path( +async fn parse_sources_file_from_path( sources_file: &Path, resolve_env: bool, ) -> Result { - match open_file(sources_file) { + let sources_file = sources_file.to_path_buf(); + tokio::task::spawn_blocking(move || match open_file(&sources_file) { Ok(file) => { let maybe_sources: Result = serde_saphyr::from_reader(config_file_reader(file, resolve_env)); @@ -106,7 +107,9 @@ fn parse_sources_file_from_path( "Can't read the sources-config file: {}: {err}", sources_file.display() ), - } + }) + .await + .map_err(|join_err| info_err!("Failed to read sources-config file: {join_err}"))? } pub fn resolve_template_and_mapping_paths<'a>(paths: &'a ConfigPaths, template_path: Option<&'a str>, mapping_path: Option<&'a str>) -> (Cow<'a, str>, Cow<'a, str>) { @@ -120,14 +123,14 @@ pub fn resolve_template_and_mapping_paths<'a>(paths: &'a ConfigPaths, template_p (effective_template_path, effective_mapping_path) } -pub fn read_sources_file_from_path_with_templates( +pub async fn read_sources_file_from_path_with_templates( sources_file: &Path, resolve_env: bool, include_computed: bool, hdhr_config: Option<&HdHomeRunDeviceOverview>, prepared_templates: Option<&[shared::model::PatternTemplate]>, ) -> Result { - let mut sources = parse_sources_file_from_path(sources_file, resolve_env)?; + let mut sources = parse_sources_file_from_path(sources_file, resolve_env).await?; if resolve_env { if let Err(err) = sources.prepare(include_computed, hdhr_config, prepared_templates) { return info_err_res!( @@ -139,35 +142,23 @@ pub fn read_sources_file_from_path_with_templates( Ok(sources) } -pub fn read_sources_file_from_path( +pub async fn read_sources_file_from_path( sources_file: &Path, resolve_env: bool, include_computed: bool, hdhr_config: Option<&HdHomeRunDeviceOverview>, ) -> Result { - read_sources_file_from_path_with_templates( - sources_file, - resolve_env, - include_computed, - hdhr_config, - None, - ) + read_sources_file_from_path_with_templates(sources_file, resolve_env, include_computed, hdhr_config, None).await } -pub fn read_sources_file( +pub async fn read_sources_file( sources_file: &str, resolve_env: bool, include_computed: bool, hdhr_config: Option<&HdHomeRunDeviceOverview>, prepared_templates: Option<&[shared::model::PatternTemplate]>, ) -> Result { - read_sources_file_from_path_with_templates( - &PathBuf::from(sources_file), - resolve_env, - include_computed, - hdhr_config, - prepared_templates, - ) + read_sources_file_from_path_with_templates(&PathBuf::from(sources_file), resolve_env, include_computed, hdhr_config, prepared_templates).await } pub fn read_config_file( @@ -319,7 +310,7 @@ pub(crate) fn read_templates( }) } -pub fn read_app_config_dto( +pub async fn read_app_config_dto( paths: &ConfigPaths, resolve_env: bool, include_computed: bool, @@ -333,7 +324,7 @@ pub fn read_app_config_dto( // Resolve effective paths for templates and mappings (with robust fallbacks) let (effective_template_path,effective_mapping_path) = resolve_template_and_mapping_paths(paths, config.template_path.as_deref(), config.mapping_path.as_deref()); - let mut sources = parse_sources_file_from_path(&PathBuf::from(sources_file), resolve_env)?; + let mut sources = parse_sources_file_from_path(&PathBuf::from(sources_file), resolve_env).await?; let mut mappings = read_mappings_file_unprepared(effective_mapping_path.as_ref(), resolve_env)? .map(|(_, mapping)| mapping); @@ -527,7 +518,7 @@ pub async fn read_initial_app_config( paths.template_file_path.replace(path); } - let mut sources_dto = parse_sources_file_from_path(&PathBuf::from(sources_file), resolve_env)?; + let mut sources_dto = parse_sources_file_from_path(&PathBuf::from(sources_file), resolve_env).await?; let (mapping_paths, mut mappings_dto) = if let Some(mappings_file) = &paths.mapping_file_path { match read_mappings_file_unprepared(mappings_file.as_str(), resolve_env) { @@ -769,7 +760,7 @@ pub async fn save_templates_config( write_config_file(file_path, backup_dir, config, TEMPLATE_FILE).await } -fn build_templates_to_persist( +async fn build_templates_to_persist( app_state: &Arc, dto: &SourcesConfigDto, ) -> Result, TuliproxError> { @@ -780,7 +771,9 @@ fn build_templates_to_persist( let existing_source_inline_templates = match parse_sources_file_from_path( Path::new(paths.sources_file_path.as_str()), true, - ) { + ) + .await + { Ok(existing_sources) => existing_sources.templates, Err(err) => { warn!( @@ -833,11 +826,11 @@ fn build_templates_to_persist( } } -pub fn validate_source_config_for_persist( +pub async fn validate_source_config_for_persist( app_state: &Arc, dto: &SourcesConfigDto, ) -> Result, TuliproxError> { - build_templates_to_persist(app_state, dto) + build_templates_to_persist(app_state, dto).await } pub async fn persist_templates_config( @@ -925,7 +918,7 @@ pub async fn validate_and_persist_source_config( app_state: &Arc, dto: SourcesConfigDto, ) -> Result { - let templates_to_persist = validate_source_config_for_persist(app_state, &dto)?; + let templates_to_persist = validate_source_config_for_persist(app_state, &dto).await?; if let Some(template_definition) = templates_to_persist.as_ref() { persist_templates_config(app_state, template_definition).await?; diff --git a/backend/src/utils/network/request.rs b/backend/src/utils/network/request.rs index 7a685988d..ccc087bf6 100644 --- a/backend/src/utils/network/request.rs +++ b/backend/src/utils/network/request.rs @@ -14,21 +14,22 @@ use futures::{StreamExt, TryStreamExt}; use log::{debug, error, log_enabled, trace, warn, Level}; use regex::Regex; use reqwest::{ - header::{HeaderMap, HeaderName, HeaderValue, CONTENT_ENCODING}, + header::{HeaderMap, HeaderName, HeaderValue, CONTENT_ENCODING, HOST}, redirect::Policy, StatusCode, }; use shared::{ error::{notify_err_res, string_to_io_error, TuliproxError}, - model::{format_elapsed_time, InputFetchMethod, DEFAULT_USER_AGENT}, + model::{format_elapsed_time, InputFetchMethod, OnConnectErrorPolicy, DEFAULT_USER_AGENT}, utils::{ filter_request_header, human_readable_byte_size, sanitize_sensitive_info, CONTENT_TYPE_JSON, ENCODING_DEFLATE, ENCODING_GZIP, }, }; use std::{ - collections::HashMap, + collections::{HashMap, HashSet}, io::{Error, ErrorKind}, + net::{IpAddr, SocketAddr}, path::{Path, PathBuf}, pin::Pin, sync::{Arc, Once}, @@ -158,6 +159,172 @@ fn resolve_provider_url_for_attempt(url: &Url, provider: Option<&Arc, + sni_host: Option, + connect_ip: Option, + dns_host: Option, +} + +impl AttemptTarget { + fn new(url: Url) -> Self { + Self { + request_url: url.clone(), + effective_url: url, + host_header: None, + sni_host: None, + connect_ip: None, + dns_host: None, + } + } +} + +fn is_ip_literal(host: &str) -> bool { host.parse::().is_ok() } + +fn format_host_header_with_port(host: &str, port: Option) -> String { + match port { + Some(port) => format!("{host}:{port}"), + None => host.to_string(), + } +} + +fn format_ip_host_header_with_port(ip: IpAddr, port: Option) -> String { + match (ip, port) { + (IpAddr::V4(addr), Some(port)) => format!("{addr}:{port}"), + (IpAddr::V4(addr), None) => addr.to_string(), + (IpAddr::V6(addr), Some(port)) => format!("[{addr}]:{port}"), + (IpAddr::V6(addr), None) => format!("[{addr}]"), + } +} + +fn resolve_attempt_target(url: &Url, provider: Option<&Arc>) -> AttemptTarget { + let resolved_url = resolve_provider_url_for_attempt(url, provider); + let Some(provider) = provider else { + return AttemptTarget::new(resolved_url); + }; + + let mut target = AttemptTarget::new(resolved_url.clone()); + let scheme = resolved_url.scheme(); + if !provider.dns_enabled_for_scheme(scheme) { + return target; + } + + let Some(host) = resolved_url.host_str() else { + return target; + }; + if is_ip_literal(host) { + return target; + } + + let Some(connect_ip) = provider.select_ip_for_host(host) else { + return target; + }; + let keep_vhost = provider.get_dns_config().is_some_and(|dns| dns.keep_vhost); + let host_header = if keep_vhost { + format_host_header_with_port(host, resolved_url.port()) + } else { + format_ip_host_header_with_port(connect_ip, resolved_url.port()) + }; + + target.host_header = Some(host_header); + target.connect_ip = Some(connect_ip); + target.dns_host = Some(host.to_ascii_lowercase()); + + if scheme.eq_ignore_ascii_case("https") { + target.sni_host = Some(host.to_string()); + return target; + } + + if scheme.eq_ignore_ascii_case("http") { + let mut effective = resolved_url.clone(); + if effective.set_host(Some(connect_ip.to_string().as_str())).is_ok() { + target.effective_url = effective; + } + } + + target +} + +fn should_try_next_ip_on_connect_error( + provider: Option<&Arc>, + target: &AttemptTarget, + attempted_ips: &mut HashSet, +) -> bool { + let Some(provider) = provider else { + return false; + }; + let Some(connect_ip) = target.connect_ip else { + return false; + }; + let Some(dns_host) = target.dns_host.as_ref() else { + return false; + }; + let Some(dns_cfg) = provider.get_dns_config() else { + return false; + }; + if dns_cfg.on_connect_error != OnConnectErrorPolicy::TryNextIp { + return false; + } + + let inserted = attempted_ips.insert(connect_ip); + if !inserted { + return false; + } + + let total_ips = provider.ip_count_for_host(dns_host); + total_ips > attempted_ips.len() +} + +fn apply_attempt_to_request( + request: &mut reqwest::Request, + target: &AttemptTarget, +) -> Result<(), std::io::Error> { + if request.url().as_str() != target.effective_url.as_str() { + *request.url_mut() = target.effective_url.clone(); + } + if let Some(host_header) = target.host_header.as_ref() { + let host = HeaderValue::from_str(host_header) + .map_err(|err| string_to_io_error(format!("Invalid host header '{host_header}': {err}")))?; + request.headers_mut().insert(HOST, host); + } + Ok(()) +} + +fn build_https_attempt_client( + app_config: &Arc, + sni_host: &str, + connect_ip: IpAddr, + connect_port: u16, +) -> Result { + let config = app_config.config.load(); + let mut builder = create_client(app_config).http1_only(); + if config.connect_timeout_secs > 0 { + builder = builder.connect_timeout(Duration::from_secs(u64::from(config.connect_timeout_secs))); + } + drop(config); + builder = builder.resolve_to_addrs(sni_host, &[SocketAddr::new(connect_ip, connect_port)]); + builder.build() +} + +async fn execute_attempt_request( + app_config: &Arc, + base_client: reqwest::Client, + request: reqwest::Request, + target: &AttemptTarget, +) -> Result { + if target.effective_url.scheme().eq_ignore_ascii_case("https") { + if let (Some(sni_host), Some(connect_ip)) = (target.sni_host.as_ref(), target.connect_ip) { + let connect_port = target.effective_url.port_or_known_default().unwrap_or(443); + let https_client = build_https_attempt_client(app_config, sni_host.as_str(), connect_ip, connect_port)?; + return https_client.execute(request).await; + } + } + base_client.execute(request).await +} + #[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss, clippy::cast_precision_loss)] pub fn calculate_retry_backoff(base_delay_ms: u64, multiplier: f64, attempt: u32) -> u64 { let base = base_delay_ms.max(1); @@ -208,102 +375,120 @@ pub async fn send_with_retry_and_provider( 'provider_loop: loop { // 2. Retry loop for the current URL - for attempt in 0..max_attempts { - let resolved_url = resolve_provider_url_for_attempt(url, provider); - // Reset the idle timer for a new attempt - idle.as_mut().reset(tokio::time::Instant::now() + idle_timeout); + 'attempt_loop: for attempt in 0..max_attempts { + let mut attempted_dns_ips = HashSet::new(); - tokio::select! { - () = &mut idle => { - warn!("Request idle for too long: {}", sanitize_sensitive_info(url.as_str())); - // 1. Try Provider Failover first - if max_provider_attempts > 1 && provider_attempts < max_provider_attempts { - if let Some(p) = provider { - if p.rotate_to_next_url_with_cycle_check(start_index).is_some() { - provider_attempts += 1; - let current_index = p.get_current_index(); - warn!("Provider '{}' idle timeout -> switching to index {}", p.name, current_index); - continue 'provider_loop; - } - } - } + 'ip_loop: loop { + let attempt_target = resolve_attempt_target(url, provider); + // Reset the idle timer for a new attempt + idle.as_mut().reset(tokio::time::Instant::now() + idle_timeout); - // 2. If no provider or rotation failed, check if we can retry the same URL - if attempt < max_attempts - 1 { - let delay = calculate_retry_backoff(backoff_ms, backoff_multiplier, attempt); - warn!("Idle timeout, retrying same URL in {}ms (attempt {})", delay, attempt + 1); - tokio::time::sleep(Duration::from_millis(delay)).await; - continue; // This will restart the 'for attempt' loop - } + let request_builder = send(&attempt_target.request_url); + let (base_client, request_result) = request_builder.build_split(); + let mut request = request_result.map_err(|err| { + string_to_io_error(format!("Failed to build request: {}", sanitize_sensitive_info(err.to_string().as_str()))) + })?; + apply_attempt_to_request(&mut request, &attempt_target)?; - return Err(string_to_io_error(format!("Request timed out and no retries left: {}", sanitize_sensitive_info(url.as_str())))); - } - - result = send(&resolved_url).send() => { - match result { - Ok(response) => { - let status = response.status(); - if allow_redirects && status.is_redirection() { - return Ok(response); - } - let is_failover = is_failover_redirect(response.url(), &failover_patterns); - if !is_failover && status.is_success() { - return Ok(response); - } - - // Failover check: Should we switch to the next provider URL? - if (is_failover || should_trigger_failover(status)) - && max_provider_attempts > 1 - && provider_attempts < max_provider_attempts - { - if let Some(p) = provider { - if p.rotate_to_next_url_with_cycle_check(start_index).is_some() { - provider_attempts += 1; - let current_index = p.get_current_index(); - warn!("Provider '{}' failover: status {} -> switching to URL index {current_index}", - p.name, format_http_status(status)); - continue 'provider_loop; - } + tokio::select! { + () = &mut idle => { + warn!("Request idle for too long: {}", sanitize_sensitive_info(url.as_str())); + // 1. Try Provider Failover first + if max_provider_attempts > 1 && provider_attempts < max_provider_attempts { + if let Some(p) = provider { + if p.rotate_to_next_url_with_cycle_check(start_index).is_some() { + provider_attempts += 1; + let current_index = p.get_current_index(); + warn!("Provider '{}' idle timeout -> switching to index {}", p.name, current_index); + continue 'provider_loop; } } - - // Standard retry check for the same URL - let is_retryable = status.is_server_error() - || matches!(status, StatusCode::TOO_MANY_REQUESTS | StatusCode::REQUEST_TIMEOUT); - - if attempt < max_attempts - 1 && is_retryable { - perform_backoff(attempt, backoff_ms, backoff_multiplier, &response).await; - continue; - } - - return Err(string_to_io_error(format!("Request failed ({}): {}", - format_http_status(status), sanitize_sensitive_info(url.as_str())))); } - Err(err) => { - // Connection errors (Timeout/Connect) trigger failover if provider exists - if (err.is_timeout() || err.is_connect()) - && max_provider_attempts > 1 - && provider_attempts < max_provider_attempts - { - if let Some(p) = provider { - if p.rotate_to_next_url_with_cycle_check(start_index).is_some() { - provider_attempts += 1; - let current_index = p.get_current_index(); - warn!("Provider '{}' failover: connection error -> switching to index {}", p.name, current_index); - continue 'provider_loop; + // 2. If no provider or rotation failed, check if we can retry the same URL + if attempt < max_attempts - 1 { + let delay = calculate_retry_backoff(backoff_ms, backoff_multiplier, attempt); + warn!("Idle timeout, retrying same URL in {}ms (attempt {})", delay, attempt + 1); + tokio::time::sleep(Duration::from_millis(delay)).await; + continue 'attempt_loop; + } + + return Err(string_to_io_error(format!("Request timed out and no retries left: {}", sanitize_sensitive_info(url.as_str())))); + } + + result = execute_attempt_request(app_config, base_client, request, &attempt_target) => { + match result { + Ok(response) => { + let status = response.status(); + if allow_redirects && status.is_redirection() { + return Ok(response); + } + let is_failover = is_failover_redirect(response.url(), &failover_patterns); + if !is_failover && status.is_success() { + return Ok(response); + } + + // Failover check: Should we switch to the next provider URL? + if (is_failover || should_trigger_failover(status)) + && max_provider_attempts > 1 + && provider_attempts < max_provider_attempts + { + if let Some(p) = provider { + if p.rotate_to_next_url_with_cycle_check(start_index).is_some() { + provider_attempts += 1; + let current_index = p.get_current_index(); + warn!("Provider '{}' failover: status {} -> switching to URL index {current_index}", + p.name, format_http_status(status)); + continue 'provider_loop; + } } } + + // Standard retry check for the same URL + let is_retryable = status.is_server_error() + || matches!(status, StatusCode::TOO_MANY_REQUESTS | StatusCode::REQUEST_TIMEOUT); + + if attempt < max_attempts - 1 && is_retryable { + perform_backoff(attempt, backoff_ms, backoff_multiplier, &response).await; + continue 'attempt_loop; + } + + return Err(string_to_io_error(format!("Request failed ({}): {}", + format_http_status(status), sanitize_sensitive_info(url.as_str())))); } - // If not a provider or rotation failed, try standard retry - if (err.is_timeout() || err.is_connect()) && attempt < max_attempts - 1 { - let delay = calculate_retry_backoff(backoff_ms, backoff_multiplier, attempt); - tokio::time::sleep(Duration::from_millis(delay)).await; - continue; - } + Err(err) => { + // For DNS IP-connect policy, attempt next IP before provider URL rotation. + if (err.is_timeout() || err.is_connect()) + && should_try_next_ip_on_connect_error(provider, &attempt_target, &mut attempted_dns_ips) + { + continue 'ip_loop; + } - return Err(string_to_io_error(format!("Request error: {}", sanitize_sensitive_info(err.to_string().as_str())))); + // Connection errors (Timeout/Connect) trigger failover if provider exists + if (err.is_timeout() || err.is_connect()) + && max_provider_attempts > 1 + && provider_attempts < max_provider_attempts + { + if let Some(p) = provider { + if p.rotate_to_next_url_with_cycle_check(start_index).is_some() { + provider_attempts += 1; + let current_index = p.get_current_index(); + warn!("Provider '{}' failover: connection error -> switching to index {}", p.name, current_index); + continue 'provider_loop; + } + } + } + + // If not a provider or rotation failed, try standard retry + if (err.is_timeout() || err.is_connect()) && attempt < max_attempts - 1 { + let delay = calculate_retry_backoff(backoff_ms, backoff_multiplier, attempt); + tokio::time::sleep(Duration::from_millis(delay)).await; + continue 'attempt_loop; + } + + return Err(string_to_io_error(format!("Request error: {}", sanitize_sensitive_info(err.to_string().as_str())))); + } } } } @@ -1312,11 +1497,105 @@ pub fn should_trigger_failover(status: StatusCode) -> bool { #[cfg(test)] mod tests { - use super::{same_origin, strip_sensitive_headers_for_cross_origin_redirect}; + use super::{ + resolve_attempt_target, same_origin, send_with_retry_and_provider, should_try_next_ip_on_connect_error, + strip_sensitive_headers_for_cross_origin_redirect, + }; + use crate::{ + model::{AppConfig, Config, ConfigProvider, ResourceRetryConfig, ReverseProxyConfig, SourcesConfig}, + utils::FileLockManager, + }; + use arc_swap::{ArcSwap, ArcSwapOption}; + use shared::model::{ + ConfigPaths, ConfigProviderDto, DnsScheme, OnConnectErrorPolicy, ProviderDnsDto, + }; use shared::utils::{get_base_url_from_str, replace_url_extension, sanitize_sensitive_info}; - use std::collections::HashMap; + use std::{ + collections::{HashMap, HashSet}, + net::SocketAddr, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }, + time::Duration, + }; + use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpListener, + }; use url::Url; + fn make_test_app_config(config: Config) -> Arc { + Arc::new(AppConfig { + config: Arc::new(ArcSwap::from_pointee(config)), + sources: Arc::new(ArcSwap::from_pointee(SourcesConfig::default())), + hdhomerun: Arc::new(ArcSwapOption::default()), + api_proxy: Arc::new(ArcSwapOption::default()), + file_locks: Arc::new(FileLockManager::default()), + paths: Arc::new(ArcSwap::from_pointee(ConfigPaths { + config_path: String::new(), + config_file_path: String::new(), + sources_file_path: String::new(), + mapping_file_path: None, + mapping_files_used: None, + template_file_path: None, + template_files_used: None, + api_proxy_file_path: String::new(), + custom_stream_response_path: None, + })), + custom_stream_response: Arc::new(ArcSwapOption::default()), + access_token_secret: [0; 32], + encrypt_secret: [0; 16], + ffprobe_available: Arc::default(), + }) + } + + fn make_provider_with_dns(keep_vhost: bool, on_connect_error: OnConnectErrorPolicy, ips: Vec<&str>) -> Arc { + let parsed_ips = ips + .into_iter() + .map(|raw| raw.parse().expect("ip must parse")) + .collect::>(); + let dto = ConfigProviderDto { + name: "provider-a".into(), + urls: vec!["http://example.com".into()], + dns: Some(ProviderDnsDto { + enabled: true, + schemes: Some(vec![DnsScheme::Http, DnsScheme::Https]), + keep_vhost, + overrides: Some(HashMap::from([("example.com".to_string(), parsed_ips)])), + on_connect_error, + ..ProviderDnsDto::default() + }), + }; + Arc::new(ConfigProvider::from(&dto)) + } + + async fn start_plain_http_server() -> (SocketAddr, Arc, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("tcp bind should work"); + let addr = listener.local_addr().expect("local addr should exist"); + let accepted = Arc::new(AtomicUsize::new(0)); + let accepted_clone = Arc::clone(&accepted); + + let handle = tokio::spawn(async move { + loop { + let Ok((mut socket, _)) = listener.accept().await else { + continue; + }; + accepted_clone.fetch_add(1, Ordering::SeqCst); + tokio::spawn(async move { + let mut buf = vec![0_u8; 2048]; + let _ = socket.read(&mut buf).await; + let _ = socket + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok") + .await; + let _ = socket.shutdown().await; + }); + } + }); + + (addr, accepted, handle) + } + #[test] fn test_url_mask() { // Replace with "***" @@ -1412,4 +1691,101 @@ mod tests { assert!(!headers.contains_key("Host")); assert_eq!(headers.get("X-Test").map(String::as_str), Some("ok")); } + + #[test] + fn test_keep_vhost_false_uses_ip_host_header_for_http() { + let provider = make_provider_with_dns(false, OnConnectErrorPolicy::TryNextIp, vec!["203.0.113.10"]); + let url = Url::parse("http://example.com:8080/stream").expect("url parse should work"); + + let target = resolve_attempt_target(&url, Some(&provider)); + assert_eq!(target.effective_url.host_str(), Some("203.0.113.10")); + assert_eq!(target.host_header.as_deref(), Some("203.0.113.10:8080")); + } + + #[test] + fn test_keep_vhost_true_uses_hostname_host_header_for_http() { + let provider = make_provider_with_dns(true, OnConnectErrorPolicy::TryNextIp, vec!["203.0.113.10"]); + let url = Url::parse("http://example.com:8080/stream").expect("url parse should work"); + + let target = resolve_attempt_target(&url, Some(&provider)); + assert_eq!(target.effective_url.host_str(), Some("203.0.113.10")); + assert_eq!(target.host_header.as_deref(), Some("example.com:8080")); + } + + #[test] + fn test_https_attempt_keeps_hostname_and_sets_sni() { + let provider = make_provider_with_dns(false, OnConnectErrorPolicy::TryNextIp, vec!["203.0.113.10"]); + let url = Url::parse("https://example.com/live").expect("url parse should work"); + + let target = resolve_attempt_target(&url, Some(&provider)); + assert_eq!(target.effective_url.host_str(), Some("example.com")); + assert_eq!(target.sni_host.as_deref(), Some("example.com")); + assert_eq!(target.connect_ip.map(|ip| ip.to_string()), Some("203.0.113.10".to_string())); + assert_eq!(target.host_header.as_deref(), Some("203.0.113.10")); + } + + #[test] + fn test_try_next_ip_policy_uses_next_ip_until_exhausted() { + let provider = make_provider_with_dns(false, OnConnectErrorPolicy::TryNextIp, vec!["203.0.113.10", "203.0.113.11"]); + let url = Url::parse("http://example.com/live").expect("url parse should work"); + let mut tried = HashSet::new(); + + let first = resolve_attempt_target(&url, Some(&provider)); + let second = resolve_attempt_target(&url, Some(&provider)); + + assert!(should_try_next_ip_on_connect_error(Some(&provider), &first, &mut tried)); + assert!(!should_try_next_ip_on_connect_error(Some(&provider), &second, &mut tried)); + } + + #[tokio::test] + async fn test_on_connect_error_try_next_ip_before_provider_rotation() { + let (addr, accepted, server_handle) = start_plain_http_server().await; + + let mut cfg = Config { + connect_timeout_secs: 1, + ..Config::default() + }; + cfg.accept_insecure_ssl_certificates = true; + cfg.reverse_proxy = Some(ReverseProxyConfig { + resource_rewrite_disabled: false, + rewrite_secret: [0; 16], + resource_retry: ResourceRetryConfig { + max_attempts: 1, + ..ResourceRetryConfig::default() + }, + disabled_header: None, + stream: None, + cache: None, + rate_limit: None, + geoip: None, + }); + let app_config = make_test_app_config(cfg); + let client = reqwest::Client::builder() + .no_proxy() + .connect_timeout(Duration::from_millis(400)) + .timeout(Duration::from_secs(2)) + .build() + .expect("http client should build"); + let url = Url::parse(format!("http://example.com:{}/ok", addr.port()).as_str()).expect("url parse should work"); + + let provider_rotate = + make_provider_with_dns(false, OnConnectErrorPolicy::RotateProviderUrl, vec!["203.0.113.1", "127.0.0.1"]); + let result_rotate = send_with_retry_and_provider(&app_config, &url, Some(&provider_rotate), false, |resolved_url| { + client.get(resolved_url.clone()) + }) + .await; + assert!(result_rotate.is_err(), "without try_next_ip policy the request should fail"); + + let provider_try_next = + make_provider_with_dns(false, OnConnectErrorPolicy::TryNextIp, vec!["203.0.113.1", "127.0.0.1"]); + let result_try_next = + send_with_retry_and_provider(&app_config, &url, Some(&provider_try_next), false, |resolved_url| { + client.get(resolved_url.clone()) + }) + .await; + assert!(result_try_next.is_ok(), "try_next_ip should succeed by trying the second IP"); + assert_eq!(accepted.load(Ordering::SeqCst), 1, "server should be reached exactly once"); + + server_handle.abort(); + } } diff --git a/backend/tests/provider_dns_https_sni.rs b/backend/tests/provider_dns_https_sni.rs new file mode 100644 index 000000000..dd3cc384e --- /dev/null +++ b/backend/tests/provider_dns_https_sni.rs @@ -0,0 +1,111 @@ +use rcgen::generate_simple_self_signed; +use rustls::{ + crypto::aws_lc_rs::sign::any_supported_type, + pki_types::PrivateKeyDer, + server::{ClientHello, ResolvesServerCert, ServerConfig}, + sign::CertifiedKey, +}; +use std::{fmt, io, net::SocketAddr, sync::Arc, time::Duration}; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpListener, +}; +use tokio_rustls::TlsAcceptor; + +#[derive(Clone)] +struct StrictSniResolver { + expected_host: String, + cert: Arc, +} + +impl fmt::Debug for StrictSniResolver { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("StrictSniResolver").field("expected_host", &self.expected_host).finish() + } +} + +impl ResolvesServerCert for StrictSniResolver { + fn resolve(&self, client_hello: ClientHello<'_>) -> Option> { + (client_hello.server_name() == Some(self.expected_host.as_str())).then(|| Arc::clone(&self.cert)) + } +} + +fn create_tls_acceptor(expected_host: &str) -> io::Result { + let generated = generate_simple_self_signed(vec![expected_host.to_string()]) + .map_err(|err| io::Error::other(format!("failed to create self-signed cert: {err}")))?; + + let cert_der = generated.cert.der().clone(); + let key_der = PrivateKeyDer::Pkcs8(generated.signing_key.serialize_der().into()); + let signing_key = + any_supported_type(&key_der).map_err(|err| io::Error::other(format!("failed to create signing key: {err}")))?; + let certified_key = Arc::new(CertifiedKey::new(vec![cert_der], signing_key)); + + let resolver = Arc::new(StrictSniResolver { expected_host: expected_host.to_string(), cert: certified_key }); + + let mut config = ServerConfig::builder().with_no_client_auth().with_cert_resolver(resolver); + config.alpn_protocols.push(b"http/1.1".to_vec()); + + Ok(TlsAcceptor::from(Arc::new(config))) +} + +async fn start_tls_server(expected_host: &str) -> io::Result<(SocketAddr, tokio::task::JoinHandle<()>)> { + let acceptor = create_tls_acceptor(expected_host)?; + let listener = TcpListener::bind("127.0.0.1:0").await?; + let addr = listener.local_addr()?; + + let handle = tokio::spawn(async move { + loop { + let Ok((socket, _)) = listener.accept().await else { + continue; + }; + let acceptor = acceptor.clone(); + tokio::spawn(async move { + let Ok(mut tls_stream) = acceptor.accept(socket).await else { + return; + }; + + let mut request = vec![0_u8; 4096]; + let _ = tls_stream.read(&mut request).await; + let _ = + tls_stream.write_all(b"HTTP/1.1 200 OK\r\ncontent-length: 2\r\nconnection: close\r\n\r\nok").await; + let _ = tls_stream.shutdown().await; + }); + } + }); + + Ok((addr, handle)) +} + +#[tokio::test] +async fn https_ip_connect_uses_hostname_sni_with_resolve_to_addrs() { + let expected_host = "sni-test.local"; + let wrong_host = "wrong-sni.local"; + + let (server_addr, handle) = start_tls_server(expected_host).await.expect("tls server should start"); + + let client_ok = reqwest::Client::builder() + .no_proxy() + .danger_accept_invalid_certs(true) + .timeout(Duration::from_secs(5)) + .resolve_to_addrs(expected_host, &[server_addr]) + .build() + .expect("reqwest client should build"); + + let ok_url = format!("https://{expected_host}:{}/health", server_addr.port()); + let ok_response = client_ok.get(ok_url).send().await.expect("request with matching SNI should succeed"); + assert_eq!(ok_response.status(), reqwest::StatusCode::OK); + + let client_wrong = reqwest::Client::builder() + .no_proxy() + .danger_accept_invalid_certs(true) + .timeout(Duration::from_secs(5)) + .resolve_to_addrs(wrong_host, &[server_addr]) + .build() + .expect("reqwest client should build"); + + let wrong_url = format!("https://{wrong_host}:{}/health", server_addr.port()); + let wrong_result = client_wrong.get(wrong_url).send().await; + assert!(wrong_result.is_err(), "request with wrong SNI must fail TLS handshake"); + + handle.abort(); +} diff --git a/shared/src/model/config/input.rs b/shared/src/model/config/input.rs index ca67b0447..e5d220f94 100644 --- a/shared/src/model/config/input.rs +++ b/shared/src/model/config/input.rs @@ -18,6 +18,7 @@ use log::warn; use std::{ collections::{HashMap, HashSet}, fmt::Display, + net::IpAddr, str::FromStr, sync::Arc, }; @@ -636,6 +637,8 @@ pub struct ConfigProviderDto { pub name: Arc, #[serde(with = "arc_str_vec_serde")] pub urls: Vec>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub dns: Option, } impl ConfigProviderDto { @@ -648,6 +651,144 @@ impl ConfigProviderDto { if self.urls.is_empty() { return info_err_res!("Urls for provider is mandatory"); } + if let Some(dns) = self.dns.as_mut() { + dns.prepare()?; + } + Ok(()) + } +} + +pub const fn default_provider_dns_refresh_secs() -> u64 { 300 } +pub const fn is_default_provider_dns_refresh_secs(v: &u64) -> bool { *v == default_provider_dns_refresh_secs() } +pub fn is_default_dns_prefer(v: &DnsPrefer) -> bool { *v == DnsPrefer::default() } +pub fn is_default_on_resolve_error(v: &OnResolveErrorPolicy) -> bool { *v == OnResolveErrorPolicy::default() } +pub fn is_default_on_connect_error(v: &OnConnectErrorPolicy) -> bool { *v == OnConnectErrorPolicy::default() } + +#[derive(Debug, Copy, Clone, serde::Serialize, serde::Deserialize, Sequence, PartialEq, Eq, Default)] +pub enum DnsPrefer { + #[serde(rename = "ipv4")] + Ipv4, + #[serde(rename = "ipv6")] + Ipv6, + #[serde(rename = "system")] + #[default] + System, +} + +#[derive(Debug, Copy, Clone, serde::Serialize, serde::Deserialize, Sequence, PartialEq, Eq)] +pub enum DnsScheme { + #[serde(rename = "http")] + Http, + #[serde(rename = "https")] + Https, +} + +#[derive(Debug, Copy, Clone, serde::Serialize, serde::Deserialize, Sequence, PartialEq, Eq, Default)] +pub enum OnResolveErrorPolicy { + #[serde(rename = "keep_last_good")] + #[default] + KeepLastGood, + #[serde(rename = "fallback_to_hostname")] + FallbackToHostname, +} + +#[derive(Debug, Copy, Clone, serde::Serialize, serde::Deserialize, Sequence, PartialEq, Eq, Default)] +pub enum OnConnectErrorPolicy { + #[serde(rename = "try_next_ip")] + #[default] + TryNextIp, + #[serde(rename = "rotate_provider_url")] + RotateProviderUrl, +} + +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq)] +#[serde(deny_unknown_fields)] +pub struct ProviderDnsDto { + #[serde(default, skip_serializing_if = "is_false")] + pub enabled: bool, + #[serde( + default = "default_provider_dns_refresh_secs", + skip_serializing_if = "is_default_provider_dns_refresh_secs" + )] + pub refresh_secs: u64, + #[serde(default, skip_serializing_if = "is_default_dns_prefer")] + pub prefer: DnsPrefer, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_addrs: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub schemes: Option>, + #[serde(default, skip_serializing_if = "is_false")] + pub keep_vhost: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub overrides: Option>>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub resolved: Option>>, + #[serde(default, skip_serializing_if = "is_default_on_resolve_error")] + pub on_resolve_error: OnResolveErrorPolicy, + #[serde(default, skip_serializing_if = "is_default_on_connect_error")] + pub on_connect_error: OnConnectErrorPolicy, +} + +impl Default for ProviderDnsDto { + fn default() -> Self { + Self { + enabled: false, + refresh_secs: default_provider_dns_refresh_secs(), + prefer: DnsPrefer::default(), + max_addrs: None, + schemes: None, + keep_vhost: false, + overrides: None, + resolved: None, + on_resolve_error: OnResolveErrorPolicy::default(), + on_connect_error: OnConnectErrorPolicy::default(), + } + } +} + +impl ProviderDnsDto { + pub fn prepare(&mut self) -> Result<(), TuliproxError> { + self.refresh_secs = self.refresh_secs.max(10); + if self.max_addrs == Some(0) { + return info_err_res!("Provider dns max_addrs must be >= 1 when set"); + } + if let Some(schemes) = self.schemes.as_mut() { + let mut unique = Vec::with_capacity(schemes.len()); + for scheme in schemes.drain(..) { + if !unique.contains(&scheme) { + unique.push(scheme); + } + } + *schemes = unique; + if schemes.is_empty() { + self.schemes = None; + } + } + + if let Some(overrides) = self.overrides.as_mut() { + let mut normalized: HashMap> = HashMap::new(); + for (host, ips) in std::mem::take(overrides) { + let host = host.trim().to_ascii_lowercase(); + if host.is_empty() { + return info_err_res!("Provider dns overrides hostname must not be empty"); + } + if ips.is_empty() { + return info_err_res!("Provider dns overrides for host '{host}' must not be empty"); + } + let entry = normalized.entry(host.clone()).or_default(); + for ip in ips { + if !entry.contains(&ip) { + entry.push(ip); + } + } + } + if normalized.is_empty() { + self.overrides = None; + } else { + *overrides = normalized; + } + } + Ok(()) } } @@ -758,4 +899,57 @@ mod tests { let result = dto.generate_auto_epg_url().unwrap(); assert_eq!(result, "provider://myprovider/xmltv.php?username=test&password=secret"); } + + #[test] + fn test_provider_dns_defaults() { + let dns = ProviderDnsDto::default(); + assert!(!dns.enabled); + assert_eq!(dns.refresh_secs, 300); + assert_eq!(dns.prefer, DnsPrefer::System); + assert_eq!(dns.on_resolve_error, OnResolveErrorPolicy::KeepLastGood); + assert_eq!(dns.on_connect_error, OnConnectErrorPolicy::TryNextIp); + assert!(dns.schemes.is_none()); + } + + #[test] + fn test_provider_dns_prepare_normalizes_overrides_and_clamps_refresh() { + let mut dns = ProviderDnsDto { + refresh_secs: 1, + schemes: Some(vec![DnsScheme::Http, DnsScheme::Http, DnsScheme::Https]), + overrides: Some(HashMap::from([( + " EXAMPLE.COM ".to_string(), + vec![ + "203.0.113.10".parse::().expect("valid ip"), + "203.0.113.10".parse::().expect("valid ip"), + ], + )])), + ..ProviderDnsDto::default() + }; + + dns.prepare().expect("dns prepare should succeed"); + + assert_eq!(dns.refresh_secs, 10); + assert_eq!(dns.schemes, Some(vec![DnsScheme::Http, DnsScheme::Https])); + let overrides = dns.overrides.expect("overrides should exist"); + assert_eq!(overrides.len(), 1); + assert!(overrides.contains_key("example.com")); + assert_eq!(overrides["example.com"].len(), 1); + } + + #[test] + fn test_provider_dns_resolved_deserializes() { + let json = r#"{ + "name":"p1", + "urls":["http://example.com"], + "dns":{ + "enabled":true, + "resolved":{"example.com":["203.0.113.10"]} + } + }"#; + + let dto: ConfigProviderDto = serde_json::from_str(json).expect("provider json should parse"); + let dns = dto.dns.expect("dns should be present"); + let resolved = dns.resolved.expect("resolved must be deserialized"); + assert_eq!(resolved.get("example.com"), Some(&vec!["203.0.113.10".parse::().expect("valid ip")])); + } } diff --git a/shared/src/utils/serde_utils.rs b/shared/src/utils/serde_utils.rs index a369472c6..08022f008 100644 --- a/shared/src/utils/serde_utils.rs +++ b/shared/src/utils/serde_utils.rs @@ -132,7 +132,9 @@ where // optional sign let mut out = String::new(); if matches!(it.peek(), Some('-' | '+')) { - out.push(it.next().unwrap()); + if let Some(val) = it.next() { + out.push(val); + } while matches!(it.peek(), Some(c) if c.is_whitespace()) { it.next(); }