Runtime Ownership hardened

Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
This commit is contained in:
Alexey
2026-08-30 08:38:03 +03:00
parent 1bb6b0bdda
commit 281f63f940
91 changed files with 3972 additions and 1239 deletions
+41 -34
View File
@@ -2,23 +2,26 @@
use std::collections::HashMap;
use std::net::{IpAddr, Ipv6Addr, SocketAddr};
use std::sync::{OnceLock, RwLock};
use std::sync::Arc;
use arc_swap::ArcSwap;
use crate::error::{ProxyError, Result};
type OverrideMap = HashMap<(String, u16), IpAddr>;
const DNS_OVERRIDE_MAX_ENTRIES: usize = 4096;
/// Immutable DNS override snapshot owned by one runtime generation.
#[derive(Debug, Clone, Default)]
pub struct DnsOverrides {
entries: std::sync::Arc<OverrideMap>,
entries: Arc<OverrideMap>,
}
impl DnsOverrides {
/// Parses a validated generation-local override snapshot.
pub fn from_entries(entries: &[String]) -> Result<Self> {
Ok(Self {
entries: std::sync::Arc::new(parse_entries(entries)?),
entries: Arc::new(parse_entries(entries)?),
})
}
@@ -35,10 +38,31 @@ impl DnsOverrides {
}
}
static DNS_OVERRIDES: OnceLock<RwLock<OverrideMap>> = OnceLock::new();
/// Atomically published DNS override snapshot owned by one runtime generation.
#[derive(Debug, Default)]
pub struct GenerationDnsResolver {
snapshot: ArcSwap<DnsOverrides>,
}
fn overrides_store() -> &'static RwLock<OverrideMap> {
DNS_OVERRIDES.get_or_init(|| RwLock::new(HashMap::new()))
impl GenerationDnsResolver {
/// Creates one resolver from a validated immutable entry set.
pub fn from_entries(entries: &[String]) -> Result<Self> {
Ok(Self {
snapshot: ArcSwap::from_pointee(DnsOverrides::from_entries(entries)?),
})
}
/// Validates and atomically publishes a new generation-local snapshot.
pub fn apply_entries(&self, entries: &[String]) -> Result<()> {
let snapshot = DnsOverrides::from_entries(entries)?;
self.snapshot.store(Arc::new(snapshot));
Ok(())
}
/// Resolves one configured override without consulting system DNS.
pub fn resolve_socket_addr(&self, host: &str, port: u16) -> Option<SocketAddr> {
self.snapshot.load().resolve_socket_addr(host, port)
}
}
fn parse_ip_spec(ip_spec: &str) -> Result<IpAddr> {
@@ -111,6 +135,11 @@ fn parse_entry(entry: &str) -> Result<((String, u16), IpAddr)> {
}
fn parse_entries(entries: &[String]) -> Result<OverrideMap> {
if entries.len() > DNS_OVERRIDE_MAX_ENTRIES {
return Err(ProxyError::Config(format!(
"network.dns_overrides exceeds maximum entry count {DNS_OVERRIDE_MAX_ENTRIES}"
)));
}
let mut parsed = HashMap::new();
for entry in entries {
let (key, ip) = parse_entry(entry)?;
@@ -125,30 +154,6 @@ pub fn validate_entries(entries: &[String]) -> Result<()> {
Ok(())
}
/// Replace runtime DNS overrides with a new validated snapshot.
pub fn install_entries(entries: &[String]) -> Result<()> {
let parsed = parse_entries(entries)?;
let mut guard = overrides_store().write().map_err(|_| {
ProxyError::Config("network.dns_overrides runtime lock is poisoned".to_string())
})?;
*guard = parsed;
Ok(())
}
/// Resolve a hostname override for `(host, port)` if present.
pub fn resolve(host: &str, port: u16) -> Option<IpAddr> {
let key = (host.to_ascii_lowercase(), port);
overrides_store()
.read()
.ok()
.and_then(|guard| guard.get(&key).copied())
}
/// Resolve a hostname override and construct a socket address when present.
pub fn resolve_socket_addr(host: &str, port: u16) -> Option<SocketAddr> {
resolve(host, port).map(|ip| SocketAddr::new(ip, port))
}
/// Parse a runtime endpoint in `host:port` format.
///
/// Supports:
@@ -199,12 +204,14 @@ mod tests {
}
#[test]
fn install_and_resolve_are_case_insensitive_for_host() {
fn generation_resolver_updates_are_case_insensitive_for_host() {
let entries = vec!["MyPetrovich.ru:8443:127.0.0.1".to_string()];
install_entries(&entries).unwrap();
let resolver = GenerationDnsResolver::from_entries(&entries).unwrap();
let resolved = resolve("mypetrovich.ru", 8443);
assert_eq!(resolved, Some("127.0.0.1".parse().unwrap()));
assert_eq!(
resolver.resolve_socket_addr("mypetrovich.ru", 8443),
Some("127.0.0.1:8443".parse().unwrap())
);
}
#[test]
+16 -3
View File
@@ -3,6 +3,7 @@
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket};
use std::sync::Arc;
use std::time::Duration;
use tokio::task::JoinSet;
@@ -12,7 +13,8 @@ use tracing::{debug, info, warn};
use crate::config::{NetworkConfig, UpstreamConfig, UpstreamType};
use crate::error::Result;
use crate::network::stun::{
DualStunResult, IpFamily, StunProbeResult, stun_probe_family_with_bind_and_tcp_fallback,
DualStunResult, IpFamily, StunProbeResult,
stun_probe_family_with_bind_tcp_fallback_and_resolver,
};
use crate::transport::UpstreamManager;
@@ -67,6 +69,11 @@ pub async fn run_probe(
stun_nat_probe_concurrency: usize,
) -> Result<NetworkProbe> {
let mut probe = NetworkProbe::default();
let dns_resolver = Arc::new(
crate::network::dns_overrides::GenerationDnsResolver::from_entries(
&config.dns_overrides,
)?,
);
let servers = collect_stun_servers(config);
let mut detected_ipv4 = detect_local_ip_v4();
let mut detected_ipv6 = detect_local_ip_v6();
@@ -88,6 +95,7 @@ pub async fn run_probe(
None,
None,
config.stun_tcp_fallback,
Arc::clone(&dns_resolver),
)
.await
}
@@ -171,6 +179,7 @@ pub async fn run_probe(
bind_v4,
bind_v6,
config.stun_tcp_fallback,
Arc::clone(&dns_resolver),
)
.await;
if let Some(reflected) = direct_stun_res.v4.map(|r| r.reflected_addr) {
@@ -286,6 +295,7 @@ async fn probe_stun_servers_parallel(
bind_v4: Option<IpAddr>,
bind_v6: Option<IpAddr>,
tcp_fallback: bool,
dns_resolver: Arc<crate::network::dns_overrides::GenerationDnsResolver>,
) -> DualStunResult {
let mut join_set = JoinSet::new();
let mut next_idx = 0usize;
@@ -295,6 +305,7 @@ async fn probe_stun_servers_parallel(
while next_idx < servers.len() || !join_set.is_empty() {
while next_idx < servers.len() && join_set.len() < concurrency {
let stun_addr = servers[next_idx].clone();
let dns_resolver = Arc::clone(&dns_resolver);
next_idx += 1;
join_set.spawn(async move {
let batch_timeout = if tcp_fallback {
@@ -303,18 +314,20 @@ async fn probe_stun_servers_parallel(
STUN_BATCH_TIMEOUT
};
let res = timeout(batch_timeout, async {
let v4 = stun_probe_family_with_bind_and_tcp_fallback(
let v4 = stun_probe_family_with_bind_tcp_fallback_and_resolver(
&stun_addr,
IpFamily::V4,
bind_v4,
tcp_fallback,
Some(dns_resolver.as_ref()),
)
.await?;
let v6 = stun_probe_family_with_bind_and_tcp_fallback(
let v6 = stun_probe_family_with_bind_tcp_fallback_and_resolver(
&stun_addr,
IpFamily::V6,
bind_v6,
tcp_fallback,
Some(dns_resolver.as_ref()),
)
.await?;
Ok::<DualStunResult, crate::error::ProxyError>(DualStunResult { v4, v6 })
+33 -8
View File
@@ -10,7 +10,7 @@ use tokio::time::{Duration, sleep, timeout};
use crate::crypto::SecureRandom;
use crate::error::{ProxyError, Result};
use crate::network::dns_overrides::{resolve, split_host_port};
use crate::network::dns_overrides::{GenerationDnsResolver, split_host_port};
fn stun_rng() -> &'static SecureRandom {
static STUN_RNG: OnceLock<SecureRandom> = OnceLock::new();
@@ -80,13 +80,32 @@ pub async fn stun_probe_family_with_bind_and_tcp_fallback(
family: IpFamily,
bind_ip: Option<IpAddr>,
tcp_fallback: bool,
) -> Result<Option<StunProbeResult>> {
stun_probe_family_with_bind_tcp_fallback_and_resolver(
stun_addr,
family,
bind_ip,
tcp_fallback,
None,
)
.await
}
/// Probes one STUN family with an optional generation-owned DNS resolver.
pub async fn stun_probe_family_with_bind_tcp_fallback_and_resolver(
stun_addr: &str,
family: IpFamily,
bind_ip: Option<IpAddr>,
tcp_fallback: bool,
dns_resolver: Option<&GenerationDnsResolver>,
) -> Result<Option<StunProbeResult>> {
let udp_attempts = if tcp_fallback { 1 } else { 3 };
let udp_result = stun_probe_family_udp(stun_addr, family, bind_ip, udp_attempts).await?;
let udp_result =
stun_probe_family_udp(stun_addr, family, bind_ip, udp_attempts, dns_resolver).await?;
if udp_result.is_some() || !tcp_fallback {
return Ok(udp_result);
}
stun_probe_family_tcp(stun_addr, family, bind_ip).await
stun_probe_family_tcp(stun_addr, family, bind_ip, dns_resolver).await
}
async fn stun_probe_family_udp(
@@ -94,6 +113,7 @@ async fn stun_probe_family_udp(
family: IpFamily,
bind_ip: Option<IpAddr>,
max_attempts: u8,
dns_resolver: Option<&GenerationDnsResolver>,
) -> Result<Option<StunProbeResult>> {
let bind_addr = match (family, bind_ip) {
(IpFamily::V4, Some(IpAddr::V4(ip))) => SocketAddr::new(IpAddr::V4(ip), 0),
@@ -111,7 +131,7 @@ async fn stun_probe_family_udp(
Err(e) => return Err(ProxyError::Proxy(format!("STUN bind failed: {e}"))),
};
let target_addr = resolve_stun_addr(stun_addr, family).await?;
let target_addr = resolve_stun_addr(stun_addr, family, dns_resolver).await?;
if let Some(addr) = target_addr {
match socket.connect(addr).await {
Ok(()) => {}
@@ -182,8 +202,9 @@ async fn stun_probe_family_tcp(
stun_addr: &str,
family: IpFamily,
bind_ip: Option<IpAddr>,
dns_resolver: Option<&GenerationDnsResolver>,
) -> Result<Option<StunProbeResult>> {
let target_addr = match resolve_stun_addr(stun_addr, family).await? {
let target_addr = match resolve_stun_addr(stun_addr, family, dns_resolver).await? {
Some(addr) => addr,
None => return Ok(None),
};
@@ -360,7 +381,11 @@ fn parse_reflected_addr(buf: &[u8], txid: &[u8]) -> Option<SocketAddr> {
None
}
async fn resolve_stun_addr(stun_addr: &str, family: IpFamily) -> Result<Option<SocketAddr>> {
async fn resolve_stun_addr(
stun_addr: &str,
family: IpFamily,
dns_resolver: Option<&GenerationDnsResolver>,
) -> Result<Option<SocketAddr>> {
if let Ok(addr) = stun_addr.parse::<SocketAddr>() {
return Ok(match (addr.is_ipv4(), family) {
(true, IpFamily::V4) | (false, IpFamily::V6) => Some(addr),
@@ -369,9 +394,9 @@ async fn resolve_stun_addr(stun_addr: &str, family: IpFamily) -> Result<Option<S
}
if let Some((host, port)) = split_host_port(stun_addr)
&& let Some(ip) = resolve(&host, port)
&& let Some(addr) = dns_resolver
.and_then(|resolver| resolver.resolve_socket_addr(&host, port))
{
let addr = SocketAddr::new(ip, port);
return Ok(match (addr.is_ipv4(), family) {
(true, IpFamily::V4) | (false, IpFamily::V6) => Some(addr),
_ => None,