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
This commit is contained in:
knylbyte
2026-03-02 12:24:05 +01:00
committed by GitHub
parent 68849cea70
commit faac89325b
19 changed files with 1749 additions and 169 deletions
+1
View File
@@ -21,6 +21,7 @@ docker/binaries
AGENTS.md
PLANS.md
STEPS.md
CLAUDE.md
*.m3u
Generated
+133 -1
View File
@@ -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"
+56 -7
View File
@@ -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://<provider_name>/...` 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://<provider_name>/...` 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
+3
View File
@@ -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 = []
+11 -10
View File
@@ -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<AppState>,
) -> Result<Option<Vec<PatternTemplate>>, 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<Option<Vec<PatternTemplate>>, 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<AppState>) -> Result<(), TuliproxError> {
let prepared_templates = Self::load_prepared_global_templates(app_state)?;
async fn load_mapping(app_state: &Arc<AppState>) -> 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<PreparedSourcesReload, TuliproxError> {
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 => {
+159 -3
View File
@@ -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<Arc<AppState>>,
axum::extract::Json(sources): axum::extract::Json<SourcesConfigDto>,
axum::extract::Json(mut sources): axum::extract::Json<SourcesConfigDto>,
) -> 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<Arc<AppState>>) -> 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<Arc<AppSta
error!("Failed to prepare users: {err}");
internal_server_error!()
} else {
let runtime_sources = app_state.app_config.sources.load();
inject_provider_dns_resolved(&mut app_config.sources, &runtime_sources);
axum::response::Json(app_config).into_response()
}
}
@@ -243,3 +279,123 @@ pub fn v1_api_config_register(router: Router<Arc<AppState>>) -> axum::Router<Arc
.route("/config/apiproxy", axum::routing::get(get_config_api_proxy_config))
.route("/config/apiproxy", axum::routing::put(save_config_api_proxy_config))
}
#[cfg(test)]
mod tests {
use super::{inject_provider_dns_resolved, strip_provider_dns_resolved};
use crate::model::{ConfigProvider, SourcesConfig};
use shared::model::{ConfigProviderDto, DnsScheme, ProviderDnsDto, SourcesConfigDto};
use std::{collections::HashMap, net::IpAddr, sync::Arc};
#[test]
fn inject_provider_dns_resolved_populates_runtime_snapshot() {
let mut dto = SourcesConfigDto {
provider: Some(vec![ConfigProviderDto {
name: "p1".into(),
urls: vec!["http://example.com".into()],
dns: Some(ProviderDnsDto {
enabled: true,
schemes: Some(vec![DnsScheme::Http]),
..ProviderDnsDto::default()
}),
}]),
..SourcesConfigDto::default()
};
let runtime_provider = Arc::new(ConfigProvider::from(&ConfigProviderDto {
name: "p1".into(),
urls: vec!["http://example.com".into()],
dns: Some(ProviderDnsDto {
enabled: true,
schemes: Some(vec![DnsScheme::Http]),
..ProviderDnsDto::default()
}),
}));
runtime_provider.store_resolved(
"example.com",
vec!["203.0.113.10".parse::<IpAddr>().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::<IpAddr>().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::<IpAddr>().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::<IpAddr>().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");
}
}
+11 -5
View File
@@ -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<AppConfig>, targets: Arc<ProcessTarget
// Keep using the original `app_state` below, which is valid because `Arc::clone` borrows.
let shared_data = Arc::clone(&app_state);
let (cancel_token_scheduler, cancel_token_hdhomerun, cancel_token_file_watch) = {
let (cancel_token_scheduler, cancel_token_hdhomerun, cancel_token_file_watch, cancel_token_provider_dns) = {
let cancel_tokens = app_state.cancel_tokens.load();
(cancel_tokens.scheduler.clone(), cancel_tokens.hdhomerun.clone(), cancel_tokens.file_watch.clone())
(
cancel_tokens.scheduler.clone(),
cancel_tokens.hdhomerun.clone(),
cancel_tokens.file_watch.clone(),
cancel_tokens.provider_dns.clone(),
)
};
if let Err(err) = load_playlists_into_memory_cache(&app_state).await {
@@ -311,6 +316,7 @@ pub async fn start_server(app_config: Arc<AppConfig>, targets: Arc<ProcessTarget
exec_interner_prune(&app_state);
exec_config_watch(&app_state, &cancel_token_file_watch);
exec_provider_dns(&app_state, &cancel_token_provider_dns);
let web_auth_enabled = is_web_auth_enabled(&cfg, web_ui_enabled);
+36 -6
View File
@@ -1,6 +1,7 @@
use crate::{
api::{
config_watch::exec_config_watch,
model::provider_dns_manager::exec_provider_dns,
model::{
metadata_update_manager::MetadataUpdateManager, ActiveProviderManager, ActiveUserManager,
ConnectionManager, DownloadQueue, EventManager, PlaylistStorage, PlaylistStorageState, SharedStreamManager,
@@ -9,8 +10,8 @@ use crate::{
scheduler::exec_scheduler,
},
model::{
AppConfig, Config, ConfigTarget, GracePeriodOptions, HdHomeRunConfig, HdHomeRunDeviceConfig, ProcessTargets,
ReverseProxyDisabledHeaderConfig, ScheduleConfig, SourcesConfig,
AppConfig, Config, ConfigProvider, ConfigTarget, GracePeriodOptions, HdHomeRunConfig, HdHomeRunDeviceConfig,
ProcessTargets, ReverseProxyDisabledHeaderConfig, ScheduleConfig, SourcesConfig,
},
repository::{get_geoip_path, load_target_into_memory_cache},
tools::lru_cache::LRUResourceCache,
@@ -69,7 +70,7 @@ struct TargetChanges {
target: Arc<ConfigTarget>,
}
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<AppState>, 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<AppState>, 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<ConfigProvider>], b: &[Arc<ConfigProvider>]) -> 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<AppState>,
+2 -1
View File
@@ -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::*,
@@ -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<IpAddr>, prefer: DnsPrefer) -> Vec<IpAddr> {
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<IpAddr>) -> Vec<IpAddr> {
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<usize>) -> std::io::Result<Vec<IpAddr>> {
let addrs = lookup_host((hostname, 0)).await?;
let mut ips: Vec<IpAddr> = 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<ConfigProvider>) -> 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<String> {
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<AppState>, provider: &Arc<ConfigProvider>) {
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<AppState>, provider_name: Arc<str>, 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::<u64>() % (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<AppState>, 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());
}
}
+3 -2
View File
@@ -1120,7 +1120,7 @@ async fn patch_source_yml_add_alias(
password: &str,
exp_date: Option<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}"),
};
@@ -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<str>,
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}"),
};
+6 -4
View File
@@ -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(),))?;
+4 -4
View File
@@ -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)]),
+225 -7
View File
@@ -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<usize>,
pub schemes: Vec<DnsScheme>,
pub keep_vhost: bool,
pub overrides: HashMap<String, Vec<IpAddr>>,
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<IpAddr>,
pub rr_index: AtomicUsize,
pub last_ok: Option<SystemTime>,
pub last_err: Option<String>,
}
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<HashMap<String, ProviderDnsCacheEntry>>,
}
impl ProviderDnsCache {
pub fn select_ip_from(&self, host: &str, ips: &[IpAddr]) -> Option<IpAddr> {
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<IpAddr> {
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<IpAddr>) {
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<String>) {
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<String, Vec<IpAddr>> {
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<str>,
pub urls: Vec<Arc<str>>,
pub current_url_index: AtomicUsize,
pub dns: Option<ProviderDnsConfig>,
pub dns_cache: Arc<ProviderDnsCache>,
}
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<IpAddr> {
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<String> {
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::<IpAddr>().is_ok() {
continue;
}
hostnames.insert(host.to_ascii_lowercase());
}
hostnames
}
pub fn store_resolved(&self, host: &str, ips: Vec<IpAddr>) { 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<String>) { self.dns_cache.mark_resolve_error(host, err); }
pub fn snapshot_resolved(&self) -> HashMap<String, Vec<IpAddr>> { 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<ConfigSource, TuliproxError> {
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 {
+22 -29
View File
@@ -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<SourcesConfigDto, TuliproxError> {
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<SourcesConfigDto, _> =
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<SourcesConfigDto, TuliproxError> {
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<SourcesConfigDto, TuliproxError> {
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<SourcesConfigDto, TuliproxError> {
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<AppState>,
dto: &SourcesConfigDto,
) -> Result<Option<TemplateDefinitionDto>, 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<AppState>,
dto: &SourcesConfigDto,
) -> Result<Option<TemplateDefinitionDto>, 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<AppState>,
dto: SourcesConfigDto,
) -> Result<SourcesConfigDto, TuliproxError> {
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?;
+465 -89
View File
@@ -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<ConfigProvi
}
}
#[derive(Debug, Clone)]
struct AttemptTarget {
request_url: Url,
effective_url: Url,
host_header: Option<String>,
sni_host: Option<String>,
connect_ip: Option<IpAddr>,
dns_host: Option<String>,
}
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::<IpAddr>().is_ok() }
fn format_host_header_with_port(host: &str, port: Option<u16>) -> String {
match port {
Some(port) => format!("{host}:{port}"),
None => host.to_string(),
}
}
fn format_ip_host_header_with_port(ip: IpAddr, port: Option<u16>) -> 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<ConfigProvider>>) -> 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<ConfigProvider>>,
target: &AttemptTarget,
attempted_ips: &mut HashSet<IpAddr>,
) -> 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<AppConfig>,
sni_host: &str,
connect_ip: IpAddr,
connect_port: u16,
) -> Result<reqwest::Client, reqwest::Error> {
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<AppConfig>,
base_client: reqwest::Client,
request: reqwest::Request,
target: &AttemptTarget,
) -> Result<reqwest::Response, reqwest::Error> {
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<AppConfig> {
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<ConfigProvider> {
let parsed_ips = ips
.into_iter()
.map(|raw| raw.parse().expect("ip must parse"))
.collect::<Vec<_>>();
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<AtomicUsize>, 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();
}
}
+111
View File
@@ -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<CertifiedKey>,
}
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<Arc<CertifiedKey>> {
(client_hello.server_name() == Some(self.expected_host.as_str())).then(|| Arc::clone(&self.cert))
}
}
fn create_tls_acceptor(expected_host: &str) -> io::Result<TlsAcceptor> {
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();
}
+194
View File
@@ -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<str>,
#[serde(with = "arc_str_vec_serde")]
pub urls: Vec<Arc<str>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dns: Option<ProviderDnsDto>,
}
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<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub schemes: Option<Vec<DnsScheme>>,
#[serde(default, skip_serializing_if = "is_false")]
pub keep_vhost: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub overrides: Option<HashMap<String, Vec<IpAddr>>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub resolved: Option<HashMap<String, Vec<IpAddr>>>,
#[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<String, Vec<IpAddr>> = 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::<IpAddr>().expect("valid ip"),
"203.0.113.10".parse::<IpAddr>().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::<IpAddr>().expect("valid ip")]));
}
}
+3 -1
View File
@@ -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();
}