mirror of
https://github.com/euzu/tuliprox.git
synced 2026-10-01 21:42:06 +02:00
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:
@@ -21,6 +21,7 @@ docker/binaries
|
||||
AGENTS.md
|
||||
PLANS.md
|
||||
STEPS.md
|
||||
CLAUDE.md
|
||||
|
||||
*.m3u
|
||||
|
||||
|
||||
Generated
+133
-1
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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 => {
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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}"),
|
||||
};
|
||||
|
||||
@@ -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(),))?;
|
||||
|
||||
@@ -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)]),
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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?;
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
@@ -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")]));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user