Compare commits

..
Author SHA1 Message Date
fanyang 82c66acc23 fix: bound peer rpc packet queues 2026-06-23 21:33:37 +08:00
fanyang 65f487ba26 fix: make stats counters thread safe 2026-06-23 21:18:07 +08:00
16 changed files with 2507 additions and 1894 deletions
-3
View File
@@ -1,3 +0,0 @@
[advisories]
# openidconnect 4.0.1 depends on rsa 0.9.10, and RUSTSEC-2023-0071 has no fixed upgrade.
ignore = ["RUSTSEC-2023-0071"]
Generated
+2011 -1529
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -33,7 +33,7 @@ sea-orm-migration = { version = "1.1" }
sqlx = { version = "0.8", features = ["sqlite", "runtime-tokio-rustls", "chrono", "uuid"] }
# Validation
validator = { version = "0.20", features = ["derive"] }
validator = { version = "0.18", features = ["derive"] }
thiserror = "1.0"
jsonwebtoken = "9.0"
+1 -1
View File
@@ -15,7 +15,7 @@ dashmap = "6.1"
url = "2.2"
async-trait = "0.1"
maxminddb = "0.27"
maxminddb = "0.24"
once_cell = "1.18"
axum = { version = "0.7", features = ["macros"] }
+19 -35
View File
@@ -245,40 +245,32 @@ impl ClientManager {
}
let location = if let Some(db) = &*geoip_db {
match db.lookup(ip).and_then(|result| result.decode::<geoip2::City>()) {
Ok(Some(city)) => {
match db.lookup::<geoip2::City>(ip) {
Ok(city) => {
let country = city
.country
.names
.simplified_chinese
.or(city.country.names.english)
.map(|s| s.to_string())
.and_then(|c| c.names)
.and_then(|n| {
n.get("zh-CN")
.or_else(|| n.get("en"))
.map(|s| s.to_string())
})
.unwrap_or_else(|| "海外".to_string());
let city_name = city
.city
.names
.simplified_chinese
.or(city.city.names.english)
.map(|s| s.to_string());
let city_name = city.city.and_then(|c| c.names).and_then(|n| {
n.get("zh-CN")
.or_else(|| n.get("en"))
.map(|s| s.to_string())
});
let region = if city.subdivisions.is_empty() {
None
} else {
let region = city
.subdivisions
.iter()
.filter_map(|x| x.names.simplified_chinese.or(x.names.english))
let region = city.subdivisions.map(|r| {
r.iter()
.filter_map(|x| x.names.as_ref())
.filter_map(|x| x.get("zh-CN").or_else(|| x.get("en")))
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(",");
if region.is_empty() {
None
} else {
Some(region)
}
};
.join(",")
});
Location {
country,
@@ -286,14 +278,6 @@ impl ClientManager {
region,
}
}
Ok(None) => {
tracing::debug!("GeoIP data not found for {}", ip);
Location {
country: "海外".to_string(),
city: None,
region: None,
}
}
Err(err) => {
tracing::debug!("GeoIP lookup failed for {}: {}", ip, err);
Location {
+5 -4
View File
@@ -236,11 +236,12 @@ http_req = { git = "https://github.com/EasyTier/http_req.git", default-features
] }
# for dns connector
hickory-resolver = "0.26.1"
hickory-proto = "0.26.1"
hickory-resolver = "0.25.2"
hickory-proto = "0.25.2"
# for magic dns
hickory-server = { version = "0.26.1", features = [
hickory-client = { version = "0.25.2", optional = true }
hickory-server = { version = "0.25.2", features = [
"resolver",
], optional = true }
@@ -400,7 +401,7 @@ jemalloc-prof = [
"jemalloc-sys/stats",
]
tracing = ["tokio/tracing", "dep:console-subscriber"]
magic-dns = ["dep:hickory-server"]
magic-dns = ["dep:hickory-client", "dep:hickory-server"]
faketcp = ["dep:flume"]
zstd = ["dep:zstd"]
# For Network Extension on macOS
+22 -35
View File
@@ -3,40 +3,33 @@ use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use anyhow::Context;
use hickory_proto::rr::RData;
use hickory_resolver::config::{
ConnectionConfig, LookupIpStrategy, NameServerConfig, ResolverConfig, ResolverOpts,
};
use hickory_resolver::net::runtime::TokioRuntimeProvider;
use hickory_proto::runtime::TokioRuntimeProvider;
use hickory_proto::xfer::Protocol;
use hickory_resolver::config::{LookupIpStrategy, NameServerConfig, ResolverConfig, ResolverOpts};
use hickory_resolver::name_server::{GenericConnector, TokioConnectionProvider};
use hickory_resolver::system_conf::read_system_conf;
use hickory_resolver::TokioResolver;
use hickory_resolver::{Resolver, TokioResolver};
use once_cell::sync::Lazy;
use tokio::net::lookup_host;
use super::error::Error;
pub fn get_default_resolver_config() -> ResolverConfig {
ResolverConfig::from_parts(
None,
vec![],
vec![
NameServerConfig::new(
"223.5.5.5".parse().unwrap(),
true,
vec![ConnectionConfig::udp()],
),
NameServerConfig::new(
"180.184.1.1".parse().unwrap(),
true,
vec![ConnectionConfig::udp()],
),
],
)
let mut default_resolve_config = ResolverConfig::new();
default_resolve_config.add_name_server(NameServerConfig::new(
"223.5.5.5:53".parse().unwrap(),
Protocol::Udp,
));
default_resolve_config.add_name_server(NameServerConfig::new(
"180.184.1.1:53".parse().unwrap(),
Protocol::Udp,
));
default_resolve_config
}
pub static ALLOW_USE_SYSTEM_DNS_RESOLVER: Lazy<AtomicBool> = Lazy::new(|| AtomicBool::new(true));
pub static RESOLVER: Lazy<Arc<TokioResolver>> =
pub static RESOLVER: Lazy<Arc<Resolver<GenericConnector<TokioRuntimeProvider>>>> =
Lazy::new(|| {
let system_cfg = read_system_conf();
let mut cfg = get_default_resolver_config();
@@ -48,11 +41,9 @@ pub static RESOLVER: Lazy<Arc<TokioResolver>> =
opt = s.1;
}
opt.ip_strategy = LookupIpStrategy::Ipv4AndIpv6;
let resolver = TokioResolver::builder_with_config(cfg, TokioRuntimeProvider::default())
.with_options(opt)
.build()
.expect("failed to build DNS resolver");
Arc::new(resolver)
let builder = TokioResolver::builder_with_config(cfg, TokioConnectionProvider::default())
.with_options(opt);
Arc::new(builder.build())
});
pub async fn resolve_txt_record(domain_name: &str) -> Result<String, Error> {
@@ -62,16 +53,12 @@ pub async fn resolve_txt_record(domain_name: &str) -> Result<String, Error> {
.await
.with_context(|| format!("txt_lookup failed, domain_name: {}", domain_name))?;
let Some(RData::TXT(txt_record)) = response
.answers()
let txt_record = response
.iter()
.next()
.map(|record| &record.data)
else {
return Err(anyhow::anyhow!("no txt record found, domain_name: {}", domain_name).into());
};
.with_context(|| format!("no txt record found, domain_name: {}", domain_name))?;
let txt_data = String::from_utf8_lossy(&txt_record.txt_data[0]);
let txt_data = String::from_utf8_lossy(&txt_record.txt_data()[0]);
tracing::info!(?txt_data, ?domain_name, "get txt record");
Ok(txt_data.to_string())
+75 -103
View File
@@ -1,8 +1,9 @@
use dashmap::DashMap;
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use std::cell::UnsafeCell;
use std::fmt;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
use tokio::time::interval;
use tokio_util::task::AbortOnDropHandle;
@@ -22,6 +23,8 @@ pub enum MetricName {
PeerRpcDuration,
/// RPC errors
PeerRpcErrors,
/// RPC/control packets dropped because the peer RPC queue is unavailable
PeerRpcPacketQueueDrops,
/// Data-plane traffic bytes sent
TrafficBytesTx,
@@ -115,6 +118,7 @@ impl fmt::Display for MetricName {
MetricName::PeerRpcServerRx => write!(f, "peer_rpc_server_rx"),
MetricName::PeerRpcDuration => write!(f, "peer_rpc_duration_ms"),
MetricName::PeerRpcErrors => write!(f, "peer_rpc_errors"),
MetricName::PeerRpcPacketQueueDrops => write!(f, "peer_rpc_packet_queue_drops"),
MetricName::TrafficBytesTx => write!(f, "traffic_bytes_tx"),
MetricName::TrafficBytesTxByInstance => write!(f, "traffic_bytes_tx_by_instance"),
@@ -374,10 +378,10 @@ impl Default for LabelSet {
}
}
/// UnsafeCounter provides a high-performance counter using UnsafeCell
/// UnsafeCounter provides a high-performance atomic counter
#[derive(Debug)]
pub struct UnsafeCounter {
value: UnsafeCell<u64>,
value: AtomicU64,
}
impl Default for UnsafeCounter {
@@ -389,121 +393,79 @@ impl Default for UnsafeCounter {
impl UnsafeCounter {
pub fn new() -> Self {
Self {
value: UnsafeCell::new(0),
value: AtomicU64::new(0),
}
}
pub fn new_with_value(initial: u64) -> Self {
Self {
value: UnsafeCell::new(initial),
value: AtomicU64::new(initial),
}
}
/// Increment the counter by the given amount
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn add(&self, delta: u64) {
let ptr = self.value.get();
unsafe {
*ptr = (*ptr).saturating_add(delta);
}
pub fn add(&self, delta: u64) {
let _ = self
.value
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
Some(current.saturating_add(delta))
});
}
/// Increment the counter by 1
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn inc(&self) {
unsafe {
self.add(1);
}
pub fn inc(&self) {
self.add(1);
}
/// Get the current value of the counter
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is modifying this counter simultaneously.
pub unsafe fn get(&self) -> u64 {
let ptr = self.value.get();
unsafe { *ptr }
pub fn get(&self) -> u64 {
self.value.load(Ordering::Relaxed)
}
/// Reset the counter to zero
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn reset(&self) {
let ptr = self.value.get();
unsafe {
*ptr = 0;
}
pub fn reset(&self) {
self.value.store(0, Ordering::Relaxed);
}
/// Set the counter to a specific value
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn set(&self, value: u64) {
let ptr = self.value.get();
unsafe {
*ptr = value;
}
pub fn set(&self, value: u64) {
self.value.store(value, Ordering::Relaxed);
}
}
// UnsafeCounter is Send + Sync because the safety is guaranteed by the caller
unsafe impl Send for UnsafeCounter {}
unsafe impl Sync for UnsafeCounter {}
/// MetricData contains both the counter and last update timestamp
/// Uses UnsafeCell for lock-free access
#[derive(Debug)]
struct MetricData {
counter: UnsafeCounter,
last_updated: UnsafeCell<Instant>,
last_updated: Mutex<Instant>,
}
impl MetricData {
fn new() -> Self {
Self {
counter: UnsafeCounter::new(),
last_updated: UnsafeCell::new(Instant::now()),
last_updated: Mutex::new(Instant::now()),
}
}
fn new_with_value(initial: u64) -> Self {
Self {
counter: UnsafeCounter::new_with_value(initial),
last_updated: UnsafeCell::new(Instant::now()),
last_updated: Mutex::new(Instant::now()),
}
}
/// Update the last_updated timestamp
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this timestamp simultaneously.
unsafe fn touch(&self) {
let ptr = self.last_updated.get();
unsafe {
*ptr = Instant::now();
}
fn touch(&self) {
*self.last_updated.lock() = Instant::now();
}
/// Get the last updated timestamp
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is modifying this timestamp simultaneously.
unsafe fn get_last_updated(&self) -> Instant {
let ptr = self.last_updated.get();
unsafe { *ptr }
fn get_last_updated(&self) -> Instant {
*self.last_updated.lock()
}
}
// MetricData is Send + Sync because the safety is guaranteed by the caller
unsafe impl Send for MetricData {}
unsafe impl Sync for MetricData {}
/// MetricKey uniquely identifies a metric with its name and labels
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct MetricKey {
@@ -546,39 +508,31 @@ impl CounterHandle {
/// Increment the counter by the given amount
pub fn add(&self, delta: u64) {
unsafe {
self.metric_data.counter.add(delta);
self.metric_data.touch();
}
self.metric_data.counter.add(delta);
self.metric_data.touch();
}
/// Increment the counter by 1
pub fn inc(&self) {
unsafe {
self.metric_data.counter.inc();
self.metric_data.touch();
}
self.metric_data.counter.inc();
self.metric_data.touch();
}
/// Get the current value of the counter
pub fn get(&self) -> u64 {
unsafe { self.metric_data.counter.get() }
self.metric_data.counter.get()
}
/// Reset the counter to zero
pub fn reset(&self) {
unsafe {
self.metric_data.counter.reset();
self.metric_data.touch();
}
self.metric_data.counter.reset();
self.metric_data.touch();
}
/// Set the counter to a specific value
pub fn set(&self, value: u64) {
unsafe {
self.metric_data.counter.set(value);
self.metric_data.touch();
}
self.metric_data.counter.set(value);
self.metric_data.touch();
}
}
@@ -624,7 +578,7 @@ impl StatsManager {
counters.retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
|| metric_data.get_last_updated() > cutoff_time
});
counters.shrink_to_fit();
}
@@ -662,7 +616,7 @@ impl StatsManager {
let key = entry.key();
let metric_data = entry.value();
let value = unsafe { metric_data.counter.get() };
let value = metric_data.counter.get();
metrics.push(MetricSnapshot {
name: key.name,
@@ -695,7 +649,7 @@ impl StatsManager {
let key = MetricKey::new(name, labels.clone());
if let Some(metric_data) = self.counters.get(&key) {
let value = unsafe { metric_data.counter.get() };
let value = metric_data.counter.get();
Some(MetricSnapshot {
name,
labels: labels.clone(),
@@ -796,17 +750,15 @@ mod tests {
async fn test_unsafe_counter() {
let counter = UnsafeCounter::new();
unsafe {
assert_eq!(counter.get(), 0);
counter.inc();
assert_eq!(counter.get(), 1);
counter.add(5);
assert_eq!(counter.get(), 6);
counter.set(10);
assert_eq!(counter.get(), 10);
counter.reset();
assert_eq!(counter.get(), 0);
}
assert_eq!(counter.get(), 0);
counter.inc();
assert_eq!(counter.get(), 1);
counter.add(5);
assert_eq!(counter.get(), 6);
counter.set(10);
assert_eq!(counter.get(), 10);
counter.reset();
assert_eq!(counter.get(), 0);
}
#[tokio::test]
@@ -951,8 +903,7 @@ mod tests {
stats
.counters
.retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > cutoff_time
});
assert_eq!(stats.metric_count(), 1);
@@ -962,12 +913,33 @@ mod tests {
stats
.counters
.retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > cutoff_time
});
assert_eq!(stats.metric_count(), 0);
}
#[tokio::test]
async fn test_counter_handle_concurrent_increment() {
const THREADS: usize = 8;
const INCREMENTS_PER_THREAD: usize = 10_000;
let stats = StatsManager::new();
let counter = stats.get_simple_counter(MetricName::TrafficPacketsForwarded);
std::thread::scope(|scope| {
for _ in 0..THREADS {
let counter = counter.clone();
scope.spawn(move || {
for _ in 0..INCREMENTS_PER_THREAD {
counter.inc();
}
});
}
});
assert_eq!(counter.get(), (THREADS * INCREMENTS_PER_THREAD) as u64);
}
#[tokio::test]
async fn test_stats_rpc_data_structures() {
// Test GetStatsRequest
+7 -10
View File
@@ -13,7 +13,7 @@ use crate::{
};
use anyhow::Context;
use dashmap::DashSet;
use hickory_resolver::proto::rr::{RData, rdata::SRV};
use hickory_resolver::proto::rr::rdata::SRV;
use rand::{Rng as _, seq::SliceRandom};
use strum::VariantArray;
@@ -85,12 +85,12 @@ impl DnsTunnelConnector {
fn handle_one_srv_record(record: &SRV, protocol: IpScheme) -> Result<(url::Url, u64), Error> {
// port must be non-zero
if record.port == 0 {
if record.port() == 0 {
return Err(anyhow::anyhow!("port must be non-zero").into());
}
let connector_dst = record.target.to_utf8();
let dst_url = format!("{}://{}:{}", protocol, connector_dst, record.port);
let connector_dst = record.target().to_utf8();
let dst_url = format!("{}://{}:{}", protocol, connector_dst, record.port());
Ok((
dst_url.parse().with_context(|| {
@@ -98,11 +98,11 @@ impl DnsTunnelConnector {
"parse dst_url failed, protocol: {}, connector_dst: {}, port: {}, dst_url: {}",
protocol,
connector_dst,
record.port,
record.port(),
dst_url
)
})?,
record.priority as _,
record.priority() as _,
))
}
@@ -129,10 +129,7 @@ impl DnsTunnelConnector {
format!("srv_lookup failed, srv_domain: {}", srv_domain)
})?;
tracing::info!(?response, ?srv_domain, "srv_lookup response");
for record in response.answers() {
let RData::SRV(record) = &record.data else {
continue;
};
for record in response.iter() {
let parsed_record = Self::handle_one_srv_record(record, **protocol);
tracing::info!(?parsed_record, ?srv_domain, "parsed_record");
if let Err(e) = &parsed_record {
+6 -4
View File
@@ -92,11 +92,12 @@ impl TryFrom<&Record> for rr::Record {
fn try_from(value: &Record) -> Result<Self, Self::Error> {
let name = value.name()?;
let ttl = value.ttl.as_secs() as u32;
let mut record = Self::update0(name, value.ttl.as_secs() as u32, value.rr_type());
record.set_dns_class(rr::DNSClass::IN);
match value.rr_type {
RecordType::A => {
let addr: Ipv4Addr = value.value.parse()?;
Ok(Self::from_rdata(name, ttl, RData::A(rr::rdata::a::A(addr))))
record.set_data(RData::A(rr::rdata::a::A(addr)));
}
RecordType::SOA => {
let soa = value.value.split_whitespace().collect::<Vec<_>>();
@@ -110,7 +111,7 @@ impl TryFrom<&Record> for rr::Record {
let retry: u32 = soa[4].parse()?;
let expire: u32 = soa[5].parse()?;
let minimum: u32 = soa[6].parse()?;
Ok(Self::from_rdata(name, ttl, RData::SOA(rr::rdata::soa::SOA::new(
record.set_data(RData::SOA(rr::rdata::soa::SOA::new(
mname,
rname,
serial,
@@ -118,10 +119,11 @@ impl TryFrom<&Record> for rr::Record {
retry.try_into().unwrap(),
expire.try_into().unwrap(),
minimum,
))))
)));
}
_ => todo!(),
}
Ok(record)
}
}
+37 -59
View File
@@ -3,14 +3,14 @@ use hickory_proto::op::Edns;
use hickory_proto::rr;
use hickory_proto::rr::LowerName;
use hickory_resolver::config::ResolverOpts;
use hickory_resolver::net::runtime::TokioRuntimeProvider;
use hickory_resolver::name_server::TokioConnectionProvider;
use hickory_resolver::system_conf::read_system_conf;
use hickory_server::net::runtime::TokioTime;
use hickory_server::server::Server as HickoryServer;
use hickory_server::ServerFuture;
use hickory_server::authority::{AuthorityObject, Catalog, ZoneType};
use hickory_server::server::{Request, RequestHandler, ResponseHandler, ResponseInfo};
use hickory_server::store::forwarder::ForwardConfig;
use hickory_server::store::{forwarder::ForwardZoneHandler, in_memory::InMemoryZoneHandler};
use hickory_server::zone_handler::{AxfrPolicy, Catalog, ZoneHandler, ZoneType};
use hickory_server::store::{forwarder::ForwardAuthority, in_memory::InMemoryAuthority};
use std::io;
use std::net::SocketAddr;
use std::str::FromStr;
use std::sync::Arc;
@@ -24,7 +24,7 @@ use crate::common::dns::get_default_resolver_config;
use super::config::{GeneralConfig, Record, RunConfig};
pub struct Server {
server: HickoryServer<CatalogRequestHandler>,
server: ServerFuture<CatalogRequestHandler>,
catalog: Arc<RwLock<Catalog>>,
general_config: GeneralConfig,
udp_local_addr: Option<SocketAddr>,
@@ -52,7 +52,7 @@ impl CatalogRequestHandler {
#[async_trait::async_trait]
impl RequestHandler for CatalogRequestHandler {
async fn handle_request<R: ResponseHandler, T: hickory_server::net::runtime::Time>(
async fn handle_request<R: ResponseHandler>(
&self,
request: &Request,
response_handle: R,
@@ -60,14 +60,14 @@ impl RequestHandler for CatalogRequestHandler {
self.catalog
.read()
.await
.handle_request::<R, T>(request, response_handle)
.handle_request(request, response_handle)
.await
}
}
pub fn build_authority(domain: &str, records: &[Record]) -> Result<InMemoryZoneHandler> {
pub fn build_authority(domain: &str, records: &[Record]) -> Result<InMemoryAuthority> {
let zone = rr::Name::from_str(domain)?;
let mut authority = InMemoryZoneHandler::empty(zone, ZoneType::Primary, AxfrPolicy::Deny);
let mut authority = InMemoryAuthority::empty(zone, ZoneType::Primary, false);
for record in records.iter() {
let r = record.try_into()?;
authority.upsert_mut(r, 0);
@@ -97,16 +97,18 @@ impl Server {
.name_servers()
.iter()
.filter(|&x| {
!config.excluded_forward_nameservers().contains(&x.ip)
!config
.excluded_forward_nameservers()
.contains(&x.socket_addr.ip())
})
.cloned()
.collect::<Vec<_>>()
.into(),
options: Some(system_conf.1),
};
let auth = ForwardZoneHandler::builder_with_config(
let auth = ForwardAuthority::builder_with_config(
forward_config,
TokioRuntimeProvider::default(),
TokioConnectionProvider::default(),
)
.build()
.unwrap();
@@ -115,7 +117,7 @@ impl Server {
let catalog = Arc::new(RwLock::new(catalog));
let handler = CatalogRequestHandler::new(catalog.clone());
let server = HickoryServer::new(handler);
let server = ServerFuture::new(handler);
Ok(Self {
server,
@@ -185,7 +187,7 @@ impl Server {
.with_context(|| format!("DNS Server failed to bind TCP address {}", address))?;
self.tcp_local_addr = Some(tcp_listener.local_addr()?);
self.server
.register_listener(tcp_listener, Duration::from_secs(5), 1024);
.register_listener(tcp_listener, Duration::from_secs(5));
}
if let Some(address) = self.general_config.listen_udp() {
@@ -201,11 +203,11 @@ impl Server {
Ok(())
}
pub async fn upsert(&self, name: LowerName, authority: Arc<dyn ZoneHandler>) {
pub async fn upsert(&self, name: LowerName, authority: Arc<dyn AuthorityObject>) {
self.catalog.write().await.upsert(name, vec![authority]);
}
pub async fn remove(&self, name: &LowerName) -> Option<Vec<Arc<dyn ZoneHandler>>> {
pub async fn remove(&self, name: &LowerName) -> Option<Vec<Arc<dyn AuthorityObject>>> {
self.catalog.write().await.remove(name)
}
@@ -214,16 +216,11 @@ impl Server {
update: &Request,
response_edns: Option<Edns>,
response_handle: R,
) -> ResponseInfo {
) -> io::Result<ResponseInfo> {
self.catalog
.write()
.await
.update(
update,
response_edns.as_ref(),
<TokioTime as hickory_server::net::runtime::Time>::current_time(),
response_handle,
)
.update(update, response_edns, response_handle)
.await
}
@@ -240,12 +237,7 @@ impl Server {
self.catalog
.read()
.await
.lookup(
request,
response_edns.as_ref(),
<TokioTime as hickory_server::net::runtime::Time>::current_time(),
response_handle,
)
.lookup(request, response_edns, response_handle)
.await
}
@@ -265,14 +257,11 @@ mod tests {
GeneralConfigBuilder, RecordBuilder, RecordType, RunConfigBuilder,
};
use anyhow::Result;
use hickory_client::client::{Client, ClientHandle};
use hickory_proto::rr;
use hickory_resolver::TokioResolver;
use hickory_resolver::config::{
ConnectionConfig, NameServerConfig, ResolverConfig, ResolverOpts,
};
use hickory_resolver::net::runtime::TokioRuntimeProvider;
use hickory_proto::runtime::TokioRuntimeProvider;
use hickory_proto::udp::UdpClientStream;
use maplit::hashmap;
use std::net::Ipv4Addr;
use std::time::Duration;
#[tokio::test]
@@ -325,34 +314,23 @@ mod tests {
server.run().await?;
let local_addr = server.udp_local_addr().unwrap();
let mut connection = ConnectionConfig::udp();
connection.port = local_addr.port();
let resolver_config = ResolverConfig::from_parts(
None,
vec![],
vec![NameServerConfig::new(
local_addr.ip(),
true,
vec![connection],
)],
);
let resolver = TokioResolver::builder_with_config(
resolver_config,
TokioRuntimeProvider::default(),
)
.with_options(ResolverOpts::default())
.build()?;
let response = resolver
.lookup(rr::Name::from_str("www.et.internal")?, rr::RecordType::A)
let stream = UdpClientStream::builder(local_addr, TokioRuntimeProvider::default()).build();
let (mut client, background) = Client::connect(stream).await?;
let background_task = tokio::spawn(background);
let response = client
.query(
rr::Name::from_str("www.et.internal")?,
rr::DNSClass::IN,
rr::RecordType::A,
)
.await?;
drop(background_task);
println!("Response: {:?}", response);
assert_eq!(response.answers().len(), 1);
let Some(rr::RData::A(ip)) = response.answers().first().map(|record| &record.data) else {
panic!("unexpected DNS response: {response:?}");
};
assert_eq!(ip.0, Ipv4Addr::new(123, 123, 123, 123));
let expected_record: rr::Record = configured_record.try_into()?;
assert_eq!(response.answers().first().unwrap(), &expected_record);
server.shutdown().await?;
Ok(())
@@ -39,10 +39,9 @@ use anyhow::Context;
use cidr::Ipv4Inet;
use dashmap::DashMap;
use hickory_proto::rr::LowerName;
use hickory_proto::serialize::binary::BinEncoder;
use hickory_server::net::{NetError, udp as dns_udp, xfer::Protocol};
use hickory_proto::serialize::binary::{BinDecodable, BinEncoder};
use hickory_server::authority::{MessageRequest, MessageResponse};
use hickory_server::server::{Request, RequestHandler, ResponseHandler, ResponseInfo};
use hickory_server::zone_handler::MessageResponse;
use multimap::MultiMap;
use pnet::packet::icmp::{IcmpTypes, MutableIcmpPacket};
use pnet::packet::ipv4::Ipv4Packet;
@@ -55,7 +54,7 @@ use pnet::packet::{
};
use std::net::{SocketAddr, SocketAddrV4};
use std::sync::Mutex;
use std::{collections::BTreeMap, net::Ipv4Addr, str::FromStr, sync::Arc, time::Duration};
use std::{collections::BTreeMap, io, net::Ipv4Addr, str::FromStr, sync::Arc, time::Duration};
static NIC_PIPELINE_NAME: &str = "magic_dns_server";
@@ -267,25 +266,25 @@ impl ResponseHandler for ResponseWrapper {
impl RecordIter<'a>,
impl RecordIter<'a>,
>,
) -> Result<ResponseInfo, NetError> {
) -> io::Result<ResponseInfo> {
let mut buffer = self
.response
.lock()
.map_err(|_| NetError::Msg("lock poisoned".to_string()))?;
buffer.clear();
.map_err(|_| io::Error::other("lock poisoned"))?;
let mut encoder = BinEncoder::new(&mut buffer);
// `max_size` should be u16::MAX for protocol other than UDP.
let max_size = response
.edns()
.map(|edns| edns.max_payload())
.unwrap_or(dns_udp::MAX_RECEIVE_BUFFER_SIZE as u16);
let max_size = if let Some(edns) = response.get_edns() {
edns.max_payload()
} else {
hickory_proto::udp::MAX_RECEIVE_BUFFER_SIZE as u16
};
encoder.set_max_size(max_size);
response
.destructive_emit(&mut encoder)
.map_err(NetError::from)
.map_err(io::Error::other)
}
}
@@ -361,12 +360,11 @@ impl MagicDnsServerInstanceData {
(
src_port,
dst_port,
Request::from_bytes(
request_payload.to_vec(),
Request::new(
MessageRequest::from_bytes(request_payload).ok()?,
SocketAddr::from(SocketAddrV4::new(src_ip, src_port)),
Protocol::Udp,
)
.ok()?,
hickory_proto::xfer::Protocol::Udp,
),
request_payload.len(),
)
};
@@ -377,7 +375,7 @@ impl MagicDnsServerInstanceData {
self.dns_server
.read_catalog()
.await
.handle_request::<ResponseWrapper, hickory_server::net::runtime::TokioTime>(
.handle_request(
&request,
ResponseWrapper {
response: response_payload_arc.clone(),
+40 -44
View File
@@ -1,16 +1,13 @@
use std::net::Ipv4Addr;
use std::net::{Ipv4Addr, SocketAddr};
use std::str::FromStr as _;
use std::sync::Arc;
use std::time::Duration;
use cidr::Ipv4Inet;
use hickory_client::client::{Client, ClientHandle as _};
use hickory_proto::rr;
use hickory_resolver::TokioResolver;
use hickory_resolver::config::{
ConnectionConfig, NameServerConfig, ResolverConfig, ResolverOpts,
};
use hickory_resolver::net::runtime::TokioRuntimeProvider;
use hickory_resolver::net::{DnsError, NetError};
use hickory_proto::runtime::TokioRuntimeProvider;
use hickory_proto::udp::UdpClientStream;
use tokio::sync::Notify;
use tokio_util::sync::CancellationToken;
@@ -69,56 +66,55 @@ pub async fn prepare_env_with_tld_dns_zone(
}
pub async fn check_dns_record(fake_ip: &Ipv4Addr, domain: &str, expected_ip: &str) {
let resolver = build_test_resolver(fake_ip);
let response = resolver
.lookup(rr::Name::from_str(domain).unwrap(), rr::RecordType::A)
let stream = UdpClientStream::builder(
SocketAddr::new((*fake_ip).into(), 53),
TokioRuntimeProvider::default(),
)
.build();
let (mut client, background) = Client::connect(stream).await.unwrap();
let background_task = tokio::spawn(background);
let response = client
.query(
rr::Name::from_str(domain).unwrap(),
rr::DNSClass::IN,
rr::RecordType::A,
)
.await
.unwrap_or_else(|e| panic!("DNS query failed unexpectedly for domain '{domain}': {e}"));
background_task.abort();
let _ = background_task.await;
println!("Response: {:?}", response);
assert_eq!(response.answers().len(), 1, "{:?}", response);
assert_eq!(response.answers().len(), 1, "{:?}", response.answers());
let resp = response.answers().first().unwrap();
let rr::RData::A(ip) = &resp.data else {
panic!("unexpected DNS response: {response:?}");
};
assert_eq!(
ip.0,
resp.clone().into_parts().rdata.into_a().unwrap().0,
expected_ip.parse::<Ipv4Addr>().unwrap()
);
}
pub async fn check_dns_record_missing(fake_ip: &Ipv4Addr, domain: &str) {
let resolver = build_test_resolver(fake_ip);
let response = resolver
.lookup(rr::Name::from_str(domain).unwrap(), rr::RecordType::A)
.await;
match response {
Ok(response) => assert!(response.answers().is_empty(), "{:?}", response),
Err(NetError::Dns(DnsError::NoRecordsFound(_))) => {}
Err(e) => {
let stream = UdpClientStream::builder(
SocketAddr::new((*fake_ip).into(), 53),
TokioRuntimeProvider::default(),
)
.build();
let (mut client, background) = Client::connect(stream).await.unwrap();
let background_task = tokio::spawn(background);
let response = client
.query(
rr::Name::from_str(domain).unwrap(),
rr::DNSClass::IN,
rr::RecordType::A,
)
.await
.unwrap_or_else(|e| {
panic!("DNS query for missing record failed unexpectedly for domain '{domain}': {e}")
}
}
}
fn build_test_resolver(fake_ip: &Ipv4Addr) -> TokioResolver {
let mut connection = ConnectionConfig::udp();
connection.port = 53;
let config = ResolverConfig::from_parts(
None,
vec![],
vec![NameServerConfig::new(
(*fake_ip).into(),
true,
vec![connection],
)],
);
TokioResolver::builder_with_config(config, TokioRuntimeProvider::default())
.with_options(ResolverOpts::default())
.build()
.unwrap()
});
background_task.abort();
let _ = background_task.await;
assert!(response.answers().is_empty(), "{:?}", response.answers());
}
#[tokio::test]
+100 -8
View File
@@ -18,7 +18,7 @@ use guarden::{Guard, defer};
use tokio::{
sync::{
Mutex,
mpsc::{self, UnboundedReceiver, UnboundedSender},
mpsc::{self, Receiver, Sender, error::TrySendError},
},
task::JoinSet,
};
@@ -30,7 +30,7 @@ use crate::{
error::Error,
global_ctx::{ArcGlobalCtx, GlobalCtx, GlobalCtxEvent, NetworkIdentity, TrustedKeySource},
join_joinset_background, shrink_dashmap,
stats_manager::{LabelSet, LabelType, MetricName, StatsManager},
stats_manager::{CounterHandle, LabelSet, LabelType, MetricName, StatsManager},
token_bucket::TokenBucket,
},
peer_center::instance::{PeerCenterInstance, PeerMapWithPeerRpcManager},
@@ -64,6 +64,35 @@ use super::{
},
};
const PEER_RPC_PACKET_QUEUE_CAPACITY: usize = 1024;
fn try_enqueue_peer_rpc_packet(
sender: &Sender<ZCPacket>,
packet: ZCPacket,
dropped_packets: &CounterHandle,
queue_name: &'static str,
) -> bool {
match sender.try_send(packet) {
Ok(()) => true,
Err(TrySendError::Full(_)) => {
dropped_packets.inc();
tracing::warn!(
queue = queue_name,
"drop peer rpc/control packet because queue is full"
);
false
}
Err(TrySendError::Closed(_)) => {
dropped_packets.inc();
tracing::warn!(
queue = queue_name,
"drop peer rpc/control packet because receiver is closed"
);
false
}
}
}
#[async_trait::async_trait]
#[auto_impl::auto_impl(&, Box, Arc)]
pub trait GlobalForeignNetworkAccessor: Send + Sync + 'static {
@@ -87,7 +116,7 @@ struct ForeignNetworkEntry {
pm_packet_sender: Mutex<Option<PacketRecvChan>>,
peer_rpc: Arc<PeerRpcManager>,
rpc_sender: UnboundedSender<ZCPacket>,
rpc_sender: Sender<ZCPacket>,
packet_recv: Mutex<Option<PacketRecvChanReceiver>>,
@@ -312,12 +341,12 @@ impl ForeignNetworkEntry {
fn build_rpc_tspt(
my_peer_id: PeerId,
peer_map: Arc<PeerMap>,
) -> (Arc<PeerRpcManager>, UnboundedSender<ZCPacket>) {
) -> (Arc<PeerRpcManager>, Sender<ZCPacket>) {
struct RpcTransport {
my_peer_id: PeerId,
peer_map: Weak<PeerMap>,
packet_recv: Mutex<UnboundedReceiver<ZCPacket>>,
packet_recv: Mutex<Receiver<ZCPacket>>,
}
#[async_trait::async_trait]
@@ -359,7 +388,8 @@ impl ForeignNetworkEntry {
}
}
let (rpc_transport_sender, peer_rpc_tspt_recv) = mpsc::unbounded_channel();
let (rpc_transport_sender, peer_rpc_tspt_recv) =
mpsc::channel(PEER_RPC_PACKET_QUEUE_CAPACITY);
let tspt = RpcTransport {
my_peer_id,
peer_map: Arc::downgrade(&peer_map),
@@ -478,6 +508,9 @@ impl ForeignNetworkEntry {
let rx_packets = self
.stats_mgr
.get_counter(MetricName::TrafficPacketsRx, label_set.clone());
let rpc_queue_drops = self
.stats_mgr
.get_counter(MetricName::PeerRpcPacketQueueDrops, label_set.clone());
self.tasks.lock().await.spawn(async move {
while let Ok(mut zc_packet) = recv_packet_from_chan(&mut recv).await {
@@ -526,7 +559,12 @@ impl ForeignNetworkEntry {
{
rx_bytes.add(buf_len as u64);
rx_packets.inc();
rpc_sender.send(zc_packet).unwrap();
try_enqueue_peer_rpc_packet(
&rpc_sender,
zc_packet,
&rpc_queue_drops,
"foreign_network_peer_rpc",
);
continue;
}
tracing::trace!(
@@ -1236,7 +1274,7 @@ impl Drop for ForeignNetworkManager {
pub mod tests {
use crate::{
common::global_ctx::tests::get_mock_global_ctx_with_network,
common::stats_manager::{LabelSet, LabelType, MetricName},
common::stats_manager::{LabelSet, LabelType, MetricName, StatsManager},
connector::udp_hole_punch::tests::{
create_mock_peer_manager_with_mock_stun, replace_stun_info_collector,
},
@@ -1253,6 +1291,7 @@ pub mod tests {
},
};
use std::{collections::HashMap, time::Duration};
use tokio::sync::mpsc;
use super::*;
@@ -1265,6 +1304,59 @@ pub mod tests {
.unwrap_or(0)
}
#[tokio::test]
async fn peer_rpc_queue_helper_enqueues_when_available() {
let (sender, mut receiver) = mpsc::channel(1);
let stats_manager = StatsManager::new();
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
assert!(try_enqueue_peer_rpc_packet(
&sender,
ZCPacket::new_with_payload(b"rpc"),
&dropped_packets,
"test_foreign_peer_rpc",
));
assert_eq!(dropped_packets.get(), 0);
assert!(receiver.try_recv().is_ok());
}
#[tokio::test]
async fn peer_rpc_queue_helper_drops_when_full() {
let (sender, _receiver) = mpsc::channel(1);
sender
.try_send(ZCPacket::new_with_payload(b"existing"))
.unwrap();
let stats_manager = StatsManager::new();
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
assert!(!try_enqueue_peer_rpc_packet(
&sender,
ZCPacket::new_with_payload(b"overflow"),
&dropped_packets,
"test_foreign_peer_rpc",
));
assert_eq!(dropped_packets.get(), 1);
}
#[tokio::test]
async fn peer_rpc_queue_helper_drops_when_closed() {
let (sender, receiver) = mpsc::channel(1);
drop(receiver);
let stats_manager = StatsManager::new();
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
assert!(!try_enqueue_peer_rpc_packet(
&sender,
ZCPacket::new_with_payload(b"closed"),
&dropped_packets,
"test_foreign_peer_rpc",
));
assert_eq!(dropped_packets.get(), 1);
}
async fn create_mock_peer_manager_for_foreign_network_ext(
network: &str,
secret: &str,
+106 -8
View File
@@ -13,7 +13,7 @@ use std::{
use tokio::{
sync::{
Mutex, RwLock,
mpsc::{self, UnboundedReceiver, UnboundedSender},
mpsc::{self, Receiver, Sender, error::TrySendError},
},
task::JoinSet,
};
@@ -72,14 +72,43 @@ use super::{
route_trait::{ArcRoute, Route},
};
const PEER_RPC_PACKET_QUEUE_CAPACITY: usize = 1024;
fn try_enqueue_peer_rpc_packet(
sender: &Sender<ZCPacket>,
packet: ZCPacket,
dropped_packets: &CounterHandle,
queue_name: &'static str,
) -> bool {
match sender.try_send(packet) {
Ok(()) => true,
Err(TrySendError::Full(_)) => {
dropped_packets.inc();
tracing::warn!(
queue = queue_name,
"drop peer rpc/control packet because queue is full"
);
false
}
Err(TrySendError::Closed(_)) => {
dropped_packets.inc();
tracing::warn!(
queue = queue_name,
"drop peer rpc/control packet because receiver is closed"
);
false
}
}
}
struct RpcTransport {
my_peer_id: PeerId,
peers: Weak<PeerMap>,
// TODO: this seems can be removed
foreign_peers: Mutex<Option<Weak<ForeignNetworkClient>>>,
packet_recv: Mutex<UnboundedReceiver<ZCPacket>>,
peer_rpc_tspt_sender: UnboundedSender<ZCPacket>,
packet_recv: Mutex<Receiver<ZCPacket>>,
peer_rpc_tspt_sender: Sender<ZCPacket>,
encryptor: Arc<dyn Encryptor>,
is_secure_mode_enabled: bool,
@@ -273,7 +302,8 @@ impl PeerManager {
.unwrap_or(false);
// TODO: remove these because we have impl pipeline processor.
let (peer_rpc_tspt_sender, peer_rpc_tspt_recv) = mpsc::unbounded_channel();
let (peer_rpc_tspt_sender, peer_rpc_tspt_recv) =
mpsc::channel(PEER_RPC_PACKET_QUEUE_CAPACITY);
let rpc_tspt = Arc::new(RpcTransport {
my_peer_id,
peers: Arc::downgrade(&peers),
@@ -1243,7 +1273,8 @@ impl PeerManager {
// for peer rpc packet
struct PeerRpcPacketProcessor {
peer_rpc_tspt_sender: UnboundedSender<ZCPacket>,
peer_rpc_tspt_sender: Sender<ZCPacket>,
dropped_packets: CounterHandle,
}
#[async_trait::async_trait]
@@ -1254,15 +1285,27 @@ impl PeerManager {
|| hdr.packet_type == PacketType::RpcReq as u8
|| hdr.packet_type == PacketType::RpcResp as u8
{
self.peer_rpc_tspt_sender.send(packet).unwrap();
try_enqueue_peer_rpc_packet(
&self.peer_rpc_tspt_sender,
packet,
&self.dropped_packets,
"local_peer_rpc",
);
None
} else {
Some(packet)
}
}
}
let peer_rpc_queue_drops = self.global_ctx.stats_manager().get_counter(
MetricName::PeerRpcPacketQueueDrops,
LabelSet::new().with_label_type(LabelType::NetworkName(
self.global_ctx.get_network_name().to_string(),
)),
);
self.add_packet_process_pipeline(Box::new(PeerRpcPacketProcessor {
peer_rpc_tspt_sender: self.peer_rpc_tspt.peer_rpc_tspt_sender.clone(),
dropped_packets: peer_rpc_queue_drops,
}))
.await;
}
@@ -2209,7 +2252,7 @@ mod tests {
PeerId,
config::Flags,
global_ctx::{NetworkIdentity, tests::get_mock_global_ctx},
stats_manager::{LabelSet, LabelType, MetricName},
stats_manager::{LabelSet, LabelType, MetricName, StatsManager},
},
connector::{
create_connector_by_url, direct::PeerManagerForDirectConnector,
@@ -2240,7 +2283,9 @@ mod tests {
},
};
use super::PeerManager;
use tokio::sync::mpsc;
use super::{PeerManager, try_enqueue_peer_rpc_packet};
async fn create_lazy_peer_manager() -> Arc<PeerManager> {
let peer_mgr = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await;
@@ -2265,6 +2310,59 @@ mod tests {
))
}
#[tokio::test]
async fn peer_rpc_queue_helper_enqueues_when_available() {
let (sender, mut receiver) = mpsc::channel(1);
let stats_manager = StatsManager::new();
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
assert!(try_enqueue_peer_rpc_packet(
&sender,
ZCPacket::new_with_payload(b"rpc"),
&dropped_packets,
"test_peer_rpc",
));
assert_eq!(dropped_packets.get(), 0);
assert!(receiver.try_recv().is_ok());
}
#[tokio::test]
async fn peer_rpc_queue_helper_drops_when_full() {
let (sender, _receiver) = mpsc::channel(1);
sender
.try_send(ZCPacket::new_with_payload(b"existing"))
.unwrap();
let stats_manager = StatsManager::new();
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
assert!(!try_enqueue_peer_rpc_packet(
&sender,
ZCPacket::new_with_payload(b"overflow"),
&dropped_packets,
"test_peer_rpc",
));
assert_eq!(dropped_packets.get(), 1);
}
#[tokio::test]
async fn peer_rpc_queue_helper_drops_when_closed() {
let (sender, receiver) = mpsc::channel(1);
drop(receiver);
let stats_manager = StatsManager::new();
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
assert!(!try_enqueue_peer_rpc_packet(
&sender,
ZCPacket::new_with_payload(b"closed"),
&dropped_packets,
"test_peer_rpc",
));
assert_eq!(dropped_packets.get(), 1);
}
struct TestCostCalculator {
costs: HashMap<(PeerId, PeerId), i32>,
}
+61 -32
View File
@@ -1,7 +1,4 @@
use std::{
cell::UnsafeCell,
sync::atomic::{AtomicU32, Ordering::Relaxed},
};
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering::Relaxed};
pub struct WindowLatency {
latency_us_window: Vec<AtomicU32>,
@@ -63,34 +60,30 @@ impl WindowLatency {
#[derive(Debug)]
pub struct Throughput {
tx_bytes: UnsafeCell<u64>,
rx_bytes: UnsafeCell<u64>,
tx_packets: UnsafeCell<u64>,
rx_packets: UnsafeCell<u64>,
tx_bytes: AtomicU64,
rx_bytes: AtomicU64,
tx_packets: AtomicU64,
rx_packets: AtomicU64,
}
impl Clone for Throughput {
fn clone(&self) -> Self {
Self {
tx_bytes: UnsafeCell::new(unsafe { *self.tx_bytes.get() }),
rx_bytes: UnsafeCell::new(unsafe { *self.rx_bytes.get() }),
tx_packets: UnsafeCell::new(unsafe { *self.tx_packets.get() }),
rx_packets: UnsafeCell::new(unsafe { *self.rx_packets.get() }),
tx_bytes: AtomicU64::new(self.tx_bytes()),
rx_bytes: AtomicU64::new(self.rx_bytes()),
tx_packets: AtomicU64::new(self.tx_packets()),
rx_packets: AtomicU64::new(self.rx_packets()),
}
}
}
// add sync::Send and sync::Sync traits to Throughput
unsafe impl Send for Throughput {}
unsafe impl Sync for Throughput {}
impl Default for Throughput {
fn default() -> Self {
Self {
tx_bytes: UnsafeCell::new(0),
rx_bytes: UnsafeCell::new(0),
tx_packets: UnsafeCell::new(0),
rx_packets: UnsafeCell::new(0),
tx_bytes: AtomicU64::new(0),
rx_bytes: AtomicU64::new(0),
tx_packets: AtomicU64::new(0),
rx_packets: AtomicU64::new(0),
}
}
}
@@ -101,32 +94,68 @@ impl Throughput {
}
pub fn tx_bytes(&self) -> u64 {
unsafe { *self.tx_bytes.get() }
self.tx_bytes.load(Relaxed)
}
pub fn rx_bytes(&self) -> u64 {
unsafe { *self.rx_bytes.get() }
self.rx_bytes.load(Relaxed)
}
pub fn tx_packets(&self) -> u64 {
unsafe { *self.tx_packets.get() }
self.tx_packets.load(Relaxed)
}
pub fn rx_packets(&self) -> u64 {
unsafe { *self.rx_packets.get() }
self.rx_packets.load(Relaxed)
}
pub fn record_tx_bytes(&self, bytes: u64) {
unsafe {
*self.tx_bytes.get() += bytes;
*self.tx_packets.get() += 1;
}
self.tx_bytes.fetch_add(bytes, Relaxed);
self.tx_packets.fetch_add(1, Relaxed);
}
pub fn record_rx_bytes(&self, bytes: u64) {
unsafe {
*self.rx_bytes.get() += bytes;
*self.rx_packets.get() += 1;
}
self.rx_bytes.fetch_add(bytes, Relaxed);
self.rx_packets.fetch_add(1, Relaxed);
}
}
#[cfg(test)]
mod tests {
use super::Throughput;
use std::sync::Arc;
#[test]
fn throughput_records_concurrent_tx_and_rx() {
const THREADS: usize = 8;
const RECORDS_PER_THREAD: usize = 10_000;
const TX_BYTES_PER_RECORD: u64 = 3;
const RX_BYTES_PER_RECORD: u64 = 7;
let throughput = Arc::new(Throughput::new());
std::thread::scope(|scope| {
for _ in 0..THREADS {
let throughput = Arc::clone(&throughput);
scope.spawn(move || {
for _ in 0..RECORDS_PER_THREAD {
throughput.record_tx_bytes(TX_BYTES_PER_RECORD);
throughput.record_rx_bytes(RX_BYTES_PER_RECORD);
}
});
}
});
let expected_packets = (THREADS * RECORDS_PER_THREAD) as u64;
assert_eq!(throughput.tx_packets(), expected_packets);
assert_eq!(throughput.rx_packets(), expected_packets);
assert_eq!(
throughput.tx_bytes(),
expected_packets * TX_BYTES_PER_RECORD
);
assert_eq!(
throughput.rx_bytes(),
expected_packets * RX_BYTES_PER_RECORD
);
}
}