use std::net::{IpAddr, SocketAddr};
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use dashmap::DashMap;
#[repr(C, align(128))]
pub struct TunnelClient {
last_receive_tick: AtomicU64,
packet_count: AtomicU64,
bytes_received: AtomicU64,
bytes_sent: AtomicU64,
bandwidth_estimate: AtomicU64,
last_bandwidth_calc: AtomicU64,
last_bandwidth_bytes: AtomicU64,
timeout_seconds: u64,
priority_score: AtomicU64,
connection_start: AtomicU64,
priority_dirty: AtomicU32,
latency_ms: AtomicU32,
packet_loss_rate: AtomicU32,
is_slow_connection: AtomicU32,
pub remote_ep: Option<SocketAddr>,
}
impl TunnelClient {
#[must_use]
pub fn new(timeout_seconds: u64) -> Self {
let now = Self::current_timestamp();
Self {
last_receive_tick: AtomicU64::new(now),
timeout_seconds,
packet_count: AtomicU64::new(0),
bytes_received: AtomicU64::new(0),
bytes_sent: AtomicU64::new(0),
bandwidth_estimate: AtomicU64::new(0),
last_bandwidth_calc: AtomicU64::new(now),
last_bandwidth_bytes: AtomicU64::new(0),
priority_score: AtomicU64::new(1000),
priority_dirty: AtomicU32::new(0),
latency_ms: AtomicU32::new(0),
packet_loss_rate: AtomicU32::new(0),
connection_start: AtomicU64::new(now),
is_slow_connection: AtomicU32::new(0),
remote_ep: None,
}
}
#[must_use]
pub fn new_with_endpoint(addr: SocketAddr, timeout_seconds: u64) -> Self {
let now = Self::current_timestamp();
Self {
last_receive_tick: AtomicU64::new(now),
timeout_seconds,
packet_count: AtomicU64::new(0),
bytes_received: AtomicU64::new(0),
bytes_sent: AtomicU64::new(0),
bandwidth_estimate: AtomicU64::new(0),
last_bandwidth_calc: AtomicU64::new(now),
last_bandwidth_bytes: AtomicU64::new(0),
priority_score: AtomicU64::new(1000),
priority_dirty: AtomicU32::new(0),
latency_ms: AtomicU32::new(0),
packet_loss_rate: AtomicU32::new(0),
connection_start: AtomicU64::new(now),
is_slow_connection: AtomicU32::new(0),
remote_ep: Some(addr),
}
}
#[inline]
pub fn set_last_receive_tick_at(&self, now: u64) {
self.last_receive_tick.store(now, Ordering::Release);
}
#[inline]
pub fn set_last_receive_tick(&self) {
self.set_last_receive_tick_at(Self::current_timestamp());
}
#[inline]
#[must_use]
pub fn is_timed_out(&self) -> bool {
let current = Self::current_timestamp();
let last = self.last_receive_tick.load(Ordering::Acquire);
current.saturating_sub(last) >= self.timeout_seconds
}
#[inline]
pub fn update_stats(&self, bytes_in: usize, bytes_out: usize, now: u64) {
self.packet_count.fetch_add(1, Ordering::Relaxed);
self.bytes_received
.fetch_add(bytes_in as u64, Ordering::Relaxed);
self.bytes_sent
.fetch_add(bytes_out as u64, Ordering::Relaxed);
let last_calc = self.last_bandwidth_calc.load(Ordering::Relaxed);
if now > last_calc {
let time_delta = now - last_calc;
if time_delta >= 1 {
if self
.last_bandwidth_calc
.compare_exchange(last_calc, now, Ordering::AcqRel, Ordering::Relaxed)
.is_err()
{
return;
}
let current_total = self.bytes_received.load(Ordering::Relaxed)
+ self.bytes_sent.load(Ordering::Relaxed);
let previous_total =
self.last_bandwidth_bytes
.swap(current_total, Ordering::Relaxed);
let delta_bytes = current_total.saturating_sub(previous_total);
let bandwidth = delta_bytes / time_delta;
let _ = self.bandwidth_estimate.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|old_estimate| Some((old_estimate * 7 + bandwidth * 3) / 10),
);
self.mark_priority_dirty();
}
}
}
#[inline]
fn mark_priority_dirty(&self) {
self.priority_dirty.store(1, Ordering::Relaxed);
}
fn recompute_priority(&self) {
let now = Self::current_timestamp();
let bandwidth = self.bandwidth_estimate.load(Ordering::Relaxed);
let latency = u64::from(self.latency_ms.load(Ordering::Relaxed));
let loss_rate = u64::from(self.packet_loss_rate.load(Ordering::Relaxed));
let connection_age = now.saturating_sub(self.connection_start.load(Ordering::Relaxed));
let bw_denom = bandwidth.max(100) + 1024;
let bandwidth_score = (1 << 20) / bw_denom;
let latency_score = latency.min(500);
let loss_score = loss_rate;
let age_score = (connection_age / 10).min(100);
let priority = bandwidth_score
.saturating_add(latency_score)
.saturating_add(loss_score)
.saturating_add(age_score);
self.priority_score.store(priority, Ordering::Relaxed);
let is_slow = u32::from(bandwidth < 100_000);
self.is_slow_connection.store(is_slow, Ordering::Relaxed);
}
#[inline]
#[must_use]
pub fn get_priority(&self) -> u64 {
if self.priority_dirty.swap(0, Ordering::Relaxed) != 0 {
self.recompute_priority();
}
self.priority_score.load(Ordering::Relaxed)
}
#[inline]
pub fn set_latency(&self, latency_ms: u32) {
self.latency_ms.store(latency_ms, Ordering::Relaxed);
self.mark_priority_dirty();
}
#[inline]
pub fn set_packet_loss_rate(&self, rate: u32) {
self.packet_loss_rate
.store(rate.min(1000), Ordering::Relaxed);
self.mark_priority_dirty();
}
#[inline]
#[must_use]
pub fn is_slow_connection(&self) -> bool {
self.is_slow_connection.load(Ordering::Relaxed) == 1
}
#[inline]
#[must_use]
pub fn current_timestamp() -> u64 {
coarsetime::Clock::now_since_epoch().as_secs()
}
#[inline]
#[must_use]
pub fn recent_timestamp() -> u64 {
coarsetime::Clock::recent_since_epoch().as_secs()
}
#[inline]
pub fn update_clock() {
coarsetime::Clock::update();
}
pub fn reset_stats(&self) {
let now = Self::current_timestamp();
self.packet_count.store(0, Ordering::Relaxed);
self.bytes_received.store(0, Ordering::Relaxed);
self.bytes_sent.store(0, Ordering::Relaxed);
self.bandwidth_estimate.store(0, Ordering::Relaxed);
self.last_bandwidth_calc.store(now, Ordering::Relaxed);
self.last_bandwidth_bytes.store(0, Ordering::Relaxed);
self.connection_start.store(now, Ordering::Relaxed);
self.mark_priority_dirty();
}
}
pub struct QualityAnalyzer {
loss_count: AtomicU32,
total_count: AtomicU32,
}
impl QualityAnalyzer {
#[must_use]
pub fn new() -> Self {
Self {
loss_count: AtomicU32::new(0),
total_count: AtomicU32::new(0),
}
}
#[inline]
pub fn record_packet(&self, lost: bool) {
if lost {
self.loss_count.fetch_add(1, Ordering::Relaxed);
}
let prev = self.total_count.fetch_add(1, Ordering::Relaxed);
if prev + 1 > 10_000 {
let did_halve = self.total_count.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|t| if t > 10_000 { Some(t / 2) } else { None },
).is_ok();
if did_halve {
let _ = self.loss_count.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|l| Some(l / 2),
);
}
}
}
#[inline]
#[must_use]
pub fn get_packet_loss_rate(&self) -> u32 {
let loss = self.loss_count.load(Ordering::Relaxed);
let total = self.total_count.load(Ordering::Relaxed);
if total == 0 {
return 0;
}
((u64::from(loss) * 1000) / u64::from(total)).min(1000) as u32
}
}
impl Default for QualityAnalyzer {
fn default() -> Self {
Self::new()
}
}
#[inline]
#[must_use]
pub fn validate_address(addr: &SocketAddr) -> bool {
let ip_valid = match addr.ip() {
IpAddr::V4(v4) => {
!v4.is_loopback() && !v4.is_unspecified() && !v4.is_broadcast() && !v4.is_multicast()
}
IpAddr::V6(v6) => !v6.is_loopback() && !v6.is_unspecified() && !v6.is_multicast(),
};
ip_valid && addr.port() != 0
}
pub fn create_dashmap_with_capacity<K, V>(capacity: usize) -> DashMap<K, V>
where
K: Eq + std::hash::Hash,
{
#[allow(clippy::redundant_closure_for_method_calls)]
let detected_cpus = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
let shard_amount = (detected_cpus * 4).next_power_of_two().clamp(4, 256);
tracing::debug!(
"Creating DashMap with {} shards for {} CPUs (capacity: {})",
shard_amount,
detected_cpus,
capacity
);
DashMap::with_capacity_and_shard_amount(capacity, shard_amount)
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, Ipv6Addr};
use std::sync::Arc;
#[test]
fn new_with_endpoint_stores_address() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)), 9000);
let client = TunnelClient::new_with_endpoint(addr, 30);
assert_eq!(client.remote_ep, Some(addr));
}
#[test]
fn new_without_endpoint_is_none() {
let client = TunnelClient::new(30);
assert_eq!(client.remote_ep, None);
}
#[test]
fn tunnel_client_timeout() {
let client = TunnelClient::new(0);
assert!(client.is_timed_out());
}
#[test]
fn tunnel_client_not_timed_out() {
let client = TunnelClient::new(60);
assert!(!client.is_timed_out());
}
#[test]
fn set_last_receive_tick_resets_timeout() {
let client = TunnelClient::new(0);
assert!(client.is_timed_out());
client.set_last_receive_tick();
let client2 = TunnelClient::new(60);
client2.set_last_receive_tick_at(TunnelClient::current_timestamp());
assert!(!client2.is_timed_out());
}
#[test]
fn set_last_receive_tick_at_with_old_timestamp_times_out() {
let client = TunnelClient::new(5);
client.set_last_receive_tick_at(0);
assert!(client.is_timed_out());
}
#[test]
fn update_stats_increments_counters() {
let client = TunnelClient::new(60);
let now = TunnelClient::current_timestamp();
client.update_stats(100, 50, now);
assert_eq!(client.packet_count.load(Ordering::Relaxed), 1);
assert_eq!(client.bytes_received.load(Ordering::Relaxed), 100);
assert_eq!(client.bytes_sent.load(Ordering::Relaxed), 50);
}
#[test]
fn update_stats_accumulates_multiple_packets() {
let client = TunnelClient::new(60);
let now = TunnelClient::current_timestamp();
client.update_stats(100, 50, now);
client.update_stats(200, 75, now);
client.update_stats(300, 25, now);
assert_eq!(client.packet_count.load(Ordering::Relaxed), 3);
assert_eq!(client.bytes_received.load(Ordering::Relaxed), 600);
assert_eq!(client.bytes_sent.load(Ordering::Relaxed), 150);
}
#[test]
fn update_stats_computes_bandwidth_after_time_delta() {
let client = TunnelClient::new(60);
let t0 = TunnelClient::current_timestamp();
client.update_stats(5000, 5000, t0);
let t1 = t0 + 2;
client.update_stats(5000, 5000, t1);
let bw = client.bandwidth_estimate.load(Ordering::Relaxed);
assert!(bw > 0, "bandwidth should be non-zero after time delta, got {bw}");
}
#[test]
fn update_stats_skips_bandwidth_when_now_equals_last_calc() {
let client = TunnelClient::new(60);
let now = TunnelClient::current_timestamp();
client.update_stats(1000, 1000, now);
client.update_stats(1000, 1000, now);
let bw = client.bandwidth_estimate.load(Ordering::Relaxed);
assert_eq!(bw, 0, "bandwidth should not update when time_delta == 0");
}
#[test]
fn update_stats_skips_bandwidth_when_now_is_in_the_past() {
let client = TunnelClient::new(60);
let now = TunnelClient::current_timestamp();
client.update_stats(1000, 1000, now);
client.update_stats(1000, 1000, now.saturating_sub(10));
let bw = client.bandwidth_estimate.load(Ordering::Relaxed);
assert_eq!(bw, 0, "bandwidth should not update when now < last_calc");
}
#[test]
fn priority_is_lazy() {
let client = TunnelClient::new(60);
let initial = client.get_priority();
assert!(initial > 0);
client.set_latency(100);
let score_before_read = client.priority_score.load(Ordering::Relaxed);
let score_after_read = client.get_priority();
assert_eq!(score_before_read, initial);
assert_ne!(score_after_read, initial);
}
#[test]
fn set_latency_stores_value() {
let client = TunnelClient::new(60);
client.set_latency(42);
assert_eq!(client.latency_ms.load(Ordering::Relaxed), 42);
}
#[test]
fn set_packet_loss_rate_clamps_to_1000() {
let client = TunnelClient::new(60);
client.set_packet_loss_rate(9999);
assert_eq!(client.packet_loss_rate.load(Ordering::Relaxed), 1000);
}
#[test]
fn set_packet_loss_rate_stores_valid_value() {
let client = TunnelClient::new(60);
client.set_packet_loss_rate(250);
assert_eq!(client.packet_loss_rate.load(Ordering::Relaxed), 250);
}
#[test]
fn high_latency_increases_priority_score() {
let client = TunnelClient::new(60);
let low = {
client.set_latency(0);
client.get_priority()
};
let high = {
client.set_latency(400);
client.get_priority()
};
assert!(high > low, "higher latency should increase priority score (lower is better)");
}
#[test]
fn is_slow_connection_reflects_bandwidth() {
let client = TunnelClient::new(60);
client.set_latency(0);
let _ = client.get_priority();
assert!(client.is_slow_connection());
}
#[test]
fn reset_stats_zeroes_counters() {
let client = TunnelClient::new(60);
let now = TunnelClient::current_timestamp();
client.update_stats(1000, 500, now);
client.update_stats(1000, 500, now + 2);
client.reset_stats();
assert_eq!(client.packet_count.load(Ordering::Relaxed), 0);
assert_eq!(client.bytes_received.load(Ordering::Relaxed), 0);
assert_eq!(client.bytes_sent.load(Ordering::Relaxed), 0);
assert_eq!(client.bandwidth_estimate.load(Ordering::Relaxed), 0);
assert_eq!(client.last_bandwidth_bytes.load(Ordering::Relaxed), 0);
}
#[test]
fn reset_stats_marks_priority_dirty() {
let client = TunnelClient::new(60);
let _ = client.get_priority(); client.reset_stats();
assert_eq!(
client.priority_dirty.load(Ordering::Relaxed),
1,
"reset_stats should mark priority as dirty"
);
}
#[test]
fn recent_timestamp_works() {
TunnelClient::update_clock();
let recent = TunnelClient::recent_timestamp();
let now = TunnelClient::current_timestamp();
assert!(now.abs_diff(recent) <= 1);
}
#[test]
fn current_timestamp_is_plausible() {
let ts = TunnelClient::current_timestamp();
assert!(ts > 1_704_067_200, "timestamp too small: {ts}");
assert!(ts < 4_102_444_800, "timestamp too large: {ts}");
}
#[test]
fn tunnel_client_is_send_and_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<TunnelClient>();
}
#[test]
fn tunnel_client_concurrent_update_stats() {
let client = Arc::new(TunnelClient::new(60));
let threads: Vec<_> = (0..4)
.map(|i| {
let c = Arc::clone(&client);
std::thread::spawn(move || {
let base = TunnelClient::current_timestamp() + i;
for j in 0..1000 {
c.update_stats(100, 50, base + j);
}
})
})
.collect();
for t in threads {
t.join().unwrap();
}
let total_packets = client.packet_count.load(Ordering::Relaxed);
assert_eq!(total_packets, 4000, "all packets should be counted");
let total_in = client.bytes_received.load(Ordering::Relaxed);
assert_eq!(total_in, 400_000, "all bytes_in should be counted");
}
#[test]
fn quality_analyzer_no_loss() {
let qa = QualityAnalyzer::new();
for _ in 0..100 {
qa.record_packet(false);
}
assert_eq!(qa.get_packet_loss_rate(), 0);
}
#[test]
fn quality_analyzer_50_percent_loss() {
let qa = QualityAnalyzer::new();
for i in 0..100 {
qa.record_packet(i % 2 == 0);
}
assert_eq!(qa.get_packet_loss_rate(), 500);
}
#[test]
fn quality_analyzer_100_percent_loss() {
let qa = QualityAnalyzer::new();
for _ in 0..100 {
qa.record_packet(true);
}
assert_eq!(qa.get_packet_loss_rate(), 1000);
}
#[test]
fn quality_analyzer_empty_returns_zero() {
let qa = QualityAnalyzer::new();
assert_eq!(qa.get_packet_loss_rate(), 0);
}
#[test]
fn quality_analyzer_halving() {
let qa = QualityAnalyzer::new();
for _ in 0..10_001 {
qa.record_packet(false);
}
let total = qa.total_count.load(Ordering::Relaxed);
assert!(total <= 5_001, "expected halving, got total={total}");
}
#[test]
fn quality_analyzer_halving_preserves_rate() {
let qa = QualityAnalyzer::new();
for i in 0..10_001 {
qa.record_packet(i % 4 == 0);
}
let rate = qa.get_packet_loss_rate();
assert!(
(200..=300).contains(&rate),
"loss rate should be ~250 after halving, got {rate}"
);
}
#[test]
fn quality_analyzer_default_trait() {
let qa = QualityAnalyzer::default();
assert_eq!(qa.get_packet_loss_rate(), 0);
assert_eq!(qa.total_count.load(Ordering::Relaxed), 0);
}
#[test]
fn quality_analyzer_is_send_and_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<QualityAnalyzer>();
}
#[test]
fn quality_analyzer_concurrent_recording() {
let qa = Arc::new(QualityAnalyzer::new());
let threads: Vec<_> = (0..4)
.map(|_| {
let q = Arc::clone(&qa);
std::thread::spawn(move || {
for i in 0..2500 {
q.record_packet(i % 2 == 0);
}
})
})
.collect();
for t in threads {
t.join().unwrap();
}
let rate = qa.get_packet_loss_rate();
assert!(
(400..=600).contains(&rate),
"concurrent 50% loss rate should be ~500, got {rate}"
);
}
#[test]
fn validate_address_rejects_loopback() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 1234);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_rejects_zero_port() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4)), 0);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_rejects_unspecified() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 1234);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_rejects_broadcast() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::BROADCAST), 1234);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_rejects_multicast_v4() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(224, 0, 0, 1)), 1234);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_accepts_valid() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4)), 1234);
assert!(validate_address(&addr));
}
#[test]
fn validate_address_accepts_private_v4() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1)), 8080);
assert!(validate_address(&addr));
}
#[test]
fn validate_address_rejects_ipv6_loopback() {
let addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 1234);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_rejects_ipv6_unspecified() {
let addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), 1234);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_rejects_ipv6_multicast() {
let addr = SocketAddr::new(
IpAddr::V6(Ipv6Addr::new(0xff02, 0, 0, 0, 0, 0, 0, 1)),
1234,
);
assert!(!validate_address(&addr));
}
#[test]
fn validate_address_accepts_valid_ipv6() {
let addr = SocketAddr::new(
IpAddr::V6(Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 1)),
443,
);
assert!(validate_address(&addr));
}
#[test]
fn validate_address_rejects_ipv6_zero_port() {
let addr = SocketAddr::new(
IpAddr::V6(Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 1)),
0,
);
assert!(!validate_address(&addr));
}
#[test]
fn dashmap_creation() {
let map: DashMap<u32, String> = create_dashmap_with_capacity(100);
map.insert(1, "test".to_string());
assert_eq!(map.len(), 1);
}
#[test]
fn dashmap_creation_zero_capacity() {
let map: DashMap<u32, u32> = create_dashmap_with_capacity(0);
map.insert(1, 42);
assert_eq!(*map.get(&1).unwrap(), 42);
}
}