Revert "fix: eliminate unsafe code in packet construction and stats counters"

This reverts commit 3464cb801a.
This commit is contained in:
fanyang
2026-06-28 14:16:47 +08:00
parent 3464cb801a
commit d4ef9decd8
12 changed files with 275 additions and 712 deletions
-32
View File
@@ -1,32 +0,0 @@
# rust-analyzer config
# Skip the `easytier-gui` (Tauri) crate because its build scripts require
# system libraries webkit2gtk-4.1 / javascriptcoregtk-4.1 that are not
# installed on this host. Excluding it keeps rust-analyzer healthy for the
# rest of the workspace.
# Override the command used to run build scripts / collect build data.
[cargo.buildScripts]
overrideCommand = [
"cargo",
"check",
"--quiet",
"--workspace",
"--exclude",
"easytier-gui",
"--message-format=json",
"--all-targets",
"--keep-going",
]
# Override the command used for diagnostics-on-save.
[check]
overrideCommand = [
"cargo",
"check",
"--workspace",
"--exclude",
"easytier-gui",
"--message-format=json",
"--all-targets",
"--keep-going",
]
Generated
-121
View File
@@ -129,12 +129,6 @@ dependencies = [
"libc", "libc",
] ]
[[package]]
name = "anes"
version = "0.1.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299"
[[package]] [[package]]
name = "anstream" name = "anstream"
version = "0.6.15" version = "0.6.15"
@@ -1151,12 +1145,6 @@ dependencies = [
"toml 0.9.12+spec-1.1.0", "toml 0.9.12+spec-1.1.0",
] ]
[[package]]
name = "cast"
version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5"
[[package]] [[package]]
name = "cc" name = "cc"
version = "1.2.10" version = "1.2.10"
@@ -1272,33 +1260,6 @@ version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6e4de3bc4ea267985becf712dc6d9eed8b04c953b3fcfb339ebc87acd9804901" checksum = "6e4de3bc4ea267985becf712dc6d9eed8b04c953b3fcfb339ebc87acd9804901"
[[package]]
name = "ciborium"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "42e69ffd6f0917f5c029256a24d0161db17cea3997d185db0d35926308770f0e"
dependencies = [
"ciborium-io",
"ciborium-ll",
"serde",
]
[[package]]
name = "ciborium-io"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "05afea1e0a06c9be33d539b876f1ce3692f4afea2cb41f740e7743225ed1c757"
[[package]]
name = "ciborium-ll"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "57663b653d948a338bfb3eeba9bb2fd5fcfaecb9e199e87e1eda4d9e8b240fd9"
dependencies = [
"ciborium-io",
"half",
]
[[package]] [[package]]
name = "cidr" name = "cidr"
version = "0.3.1" version = "0.3.1"
@@ -1644,42 +1605,6 @@ dependencies = [
"cfg-if", "cfg-if",
] ]
[[package]]
name = "criterion"
version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f2b12d017a929603d80db1831cd3a24082f8137ce19c69e6447f54f5fc8d692f"
dependencies = [
"anes",
"cast",
"ciborium",
"clap",
"criterion-plot",
"is-terminal",
"itertools 0.10.5",
"num-traits",
"once_cell",
"oorandom",
"plotters",
"rayon",
"regex",
"serde",
"serde_derive",
"serde_json",
"tinytemplate",
"walkdir",
]
[[package]]
name = "criterion-plot"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6b50826342786a51a89e2da3a28f1c32b06e387201bc2d19791f622c673706b1"
dependencies = [
"cast",
"itertools 0.10.5",
]
[[package]] [[package]]
name = "critical-section" name = "critical-section"
version = "1.2.0" version = "1.2.0"
@@ -2374,7 +2299,6 @@ dependencies = [
"clap_complete", "clap_complete",
"clap_complete_nushell", "clap_complete_nushell",
"console-subscriber", "console-subscriber",
"criterion",
"crossbeam", "crossbeam",
"ctor 0.8.0", "ctor 0.8.0",
"dashmap", "dashmap",
@@ -2469,7 +2393,6 @@ dependencies = [
"tempfile", "tempfile",
"terminal_size", "terminal_size",
"thiserror 1.0.63", "thiserror 1.0.63",
"thread_local",
"thunk-rs", "thunk-rs",
"tikv-jemalloc-ctl", "tikv-jemalloc-ctl",
"tikv-jemalloc-sys", "tikv-jemalloc-sys",
@@ -6048,12 +5971,6 @@ dependencies = [
"portable-atomic", "portable-atomic",
] ]
[[package]]
name = "oorandom"
version = "11.1.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e"
[[package]] [[package]]
name = "opaque-debug" name = "opaque-debug"
version = "0.3.1" version = "0.3.1"
@@ -6684,34 +6601,6 @@ dependencies = [
"time", "time",
] ]
[[package]]
name = "plotters"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5aeb6f403d7a4911efb1e33402027fc44f29b5bf6def3effcc22d7bb75f2b747"
dependencies = [
"num-traits",
"plotters-backend",
"plotters-svg",
"wasm-bindgen",
"web-sys",
]
[[package]]
name = "plotters-backend"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df42e13c12958a16b3f7f4386b9ab1f3e7933914ecea48da7139435263a4172a"
[[package]]
name = "plotters-svg"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "51bae2ac328883f7acdfea3d66a7c35751187f870bc81f94563733a154d7a670"
dependencies = [
"plotters-backend",
]
[[package]] [[package]]
name = "pnet" name = "pnet"
version = "0.35.0" version = "0.35.0"
@@ -10132,16 +10021,6 @@ dependencies = [
"zerovec", "zerovec",
] ]
[[package]]
name = "tinytemplate"
version = "1.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "be4d6b5f19ff7664e8c98d03e2139cb510db9b0a60b55f8e8709b689d939b6bc"
dependencies = [
"serde",
"serde_json",
]
[[package]] [[package]]
name = "tinyvec" name = "tinyvec"
version = "1.8.0" version = "1.8.0"
-10
View File
@@ -217,7 +217,6 @@ smoltcp = { git = "https://github.com/smoltcp-rs/smoltcp.git", rev = "0a926767a6
"async", "async",
] } ] }
parking_lot = { version = "0.12.0" } parking_lot = { version = "0.12.0" }
thread_local = "1.1"
wildmatch = "2.3.4" wildmatch = "2.3.4"
@@ -344,15 +343,6 @@ futures-util = "0.3.31"
maplit = "1.0.2" maplit = "1.0.2"
tempfile = "3.22.0" tempfile = "3.22.0"
ctor = "0.8.0" ctor = "0.8.0"
criterion = { version = "0.5", features = ["html_reports"] }
[[bench]]
name = "acl_hotpath"
harness = false
[[bench]]
name = "zc_packet"
harness = false
[target.'cfg(target_os = "linux")'.dev-dependencies] [target.'cfg(target_os = "linux")'.dev-dependencies]
defguard_wireguard_rs = "0.4.2" defguard_wireguard_rs = "0.4.2"
-137
View File
@@ -1,137 +0,0 @@
use std::sync::Arc;
use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main};
use easytier::common::acl_processor::{AclProcessor, PacketInfo};
use easytier::proto::acl::*;
use std::net::{IpAddr, Ipv4Addr};
fn make_acl_config() -> Acl {
let mut acl_config = Acl::default();
let mut acl_v1 = AclV1::default();
let mut chain = Chain {
name: "bench_inbound".to_string(),
chain_type: ChainType::Inbound as i32,
enabled: true,
..Default::default()
};
chain.rules.push(Rule {
name: "allow_all".to_string(),
priority: 100,
enabled: true,
action: Action::Allow as i32,
protocol: Protocol::Any as i32,
..Default::default()
});
acl_v1.chains.push(chain);
acl_config.acl_v1 = Some(acl_v1);
acl_config
}
fn make_packet_info() -> PacketInfo {
PacketInfo {
src_ip: IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100)),
dst_ip: IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)),
src_port: Some(12345),
dst_port: Some(80),
protocol: Protocol::Tcp,
packet_size: 1024,
src_groups: Arc::new(vec![]),
dst_groups: Arc::new(vec![]),
}
}
fn bench_cache_hit_single(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let processor = rt.block_on(async { AclProcessor::new(make_acl_config()) });
let packet_info = make_packet_info();
// Prime the cache
let _ = processor.process_packet(&packet_info, ChainType::Inbound);
c.bench_function("acl_cache_hit_1t", |b| {
b.iter(|| {
std::hint::black_box(processor.process_packet(&packet_info, ChainType::Inbound));
});
});
}
fn bench_cache_hit_multi(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let mut group = c.benchmark_group("acl_cache_hit_multi");
for threads in [2, 4, 8] {
let processor = Arc::new(rt.block_on(async { AclProcessor::new(make_acl_config()) }));
let packet_info = Arc::new(make_packet_info());
// Prime the cache
let _ = processor.process_packet(&packet_info, ChainType::Inbound);
group.bench_with_input(
BenchmarkId::from_parameter(threads),
&threads,
|b, &threads| {
b.iter_custom(|iters| {
use std::sync::Barrier;
use std::thread;
let barrier = Arc::new(Barrier::new(threads + 1));
let per_thread = (iters / threads as u64) as usize;
let mut handles = Vec::with_capacity(threads);
for _ in 0..threads {
let processor = Arc::clone(&processor);
let packet_info = Arc::clone(&packet_info);
let barrier = Arc::clone(&barrier);
handles.push(thread::spawn(move || {
barrier.wait();
for _ in 0..per_thread {
std::hint::black_box(
processor.process_packet(&packet_info, ChainType::Inbound),
);
}
}));
}
let start = std::time::Instant::now();
barrier.wait();
for handle in handles {
handle.join().unwrap();
}
start.elapsed()
});
},
);
}
group.finish();
}
fn bench_unique_rule_match(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let processor = rt.block_on(async { AclProcessor::new(make_acl_config()) });
c.bench_function("acl_unique_rule_match_1t", |b| {
let mut i = 0usize;
b.iter(|| {
let mut packet_info = make_packet_info();
packet_info.src_port = Some((1024 + (i % 60_000)) as u16);
packet_info.src_ip = IpAddr::V4(Ipv4Addr::new(
10,
((i >> 16) & 0xff) as u8,
((i >> 8) & 0xff) as u8,
(i & 0xff) as u8,
));
std::hint::black_box(processor.process_packet(&packet_info, ChainType::Inbound));
i = i.wrapping_add(1);
});
});
}
criterion_group!(
benches,
bench_cache_hit_single,
bench_cache_hit_multi,
bench_unique_rule_match
);
criterion_main!(benches);
-51
View File
@@ -1,51 +0,0 @@
use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main};
use easytier::tunnel::packet_def::{ZCPacket, ZCPacketType};
fn bench_new_with_payload(c: &mut Criterion) {
let mut group = c.benchmark_group("zc_new_with_payload");
for size in [64usize, 1500] {
let payload = vec![0xabu8; size];
group.bench_with_input(BenchmarkId::from_parameter(size), &payload, |b, payload| {
b.iter(|| {
std::hint::black_box(ZCPacket::new_with_payload(std::hint::black_box(payload)));
});
});
}
group.finish();
}
fn bench_new_for_foreign_network(c: &mut Criterion) {
let payload = vec![0xabu8; 64];
let foreign_packet = ZCPacket::new_with_payload(&payload);
let network_name = "bench-network".to_string();
c.bench_function("zc_new_for_foreign_network_64b", |b| {
b.iter(|| {
std::hint::black_box(ZCPacket::new_for_foreign_network(
std::hint::black_box(&network_name),
42,
std::hint::black_box(&foreign_packet),
));
});
});
}
fn bench_convert_type(c: &mut Criterion) {
let payload = vec![0xabu8; 64];
let packet = ZCPacket::new_with_payload(&payload);
c.bench_function("zc_convert_type_tcp_64b", |b| {
b.iter(|| {
let p = std::hint::black_box(packet.clone());
std::hint::black_box(p.convert_type(ZCPacketType::TCP));
});
});
}
criterion_group!(
benches,
bench_new_with_payload,
bench_new_for_foreign_network,
bench_convert_type
);
criterion_main!(benches);
+68 -100
View File
@@ -3,8 +3,7 @@ use std::{
net::{IpAddr, SocketAddr}, net::{IpAddr, SocketAddr},
str::FromStr as _, str::FromStr as _,
sync::Arc, sync::Arc,
sync::atomic::{AtomicU64, Ordering::Relaxed}, time::{Duration, SystemTime, UNIX_EPOCH},
time::Duration,
}; };
use hotpath::instant::Instant; use hotpath::instant::Instant;
@@ -77,41 +76,7 @@ pub struct FastLookupRule {
pub stateful: bool, pub stateful: bool,
pub rate_limit: u32, pub rate_limit: u32,
pub burst_limit: u32, pub burst_limit: u32,
pub rule_stats: Arc<RuleStatsTracker>, pub rule_stats: Arc<RuleStats>,
}
#[derive(Debug)]
pub struct RuleStatsTracker {
rule: Option<Rule>,
packets: AtomicU64,
bytes: AtomicU64,
}
impl RuleStatsTracker {
fn new(rule: Option<Rule>) -> Self {
Self {
rule,
packets: AtomicU64::new(0),
bytes: AtomicU64::new(0),
}
}
#[inline]
#[cfg_attr(feature = "hotpath", hotpath::measure(impl_type = "RuleStatsTracker"))]
fn increment(&self, packet_size: usize) {
self.packets.fetch_add(1, Relaxed);
self.bytes.fetch_add(packet_size as u64, Relaxed);
}
fn snapshot(&self) -> RuleStats {
RuleStats {
rule: self.rule.clone(),
stat: Some(StatItem {
packet_count: self.packets.load(Relaxed),
byte_count: self.bytes.load(Relaxed),
}),
}
}
} }
// Cache key combining packet info and chain type // Cache key combining packet info and chain type
@@ -143,17 +108,17 @@ impl AclCacheKey {
} }
// Cache entry with timestamp for LRU cleanup // Cache entry with timestamp for LRU cleanup
#[derive(Debug)] #[derive(Debug, Clone)]
pub(crate) struct AclCacheEntry { pub(crate) struct AclCacheEntry {
pub action: Action, pub action: Action,
pub matched_rule: RuleId, pub matched_rule: RuleId,
pub last_access: AtomicU64, pub last_access: Instant,
// New fields to track rule characteristics for proper cache behavior // New fields to track rule characteristics for proper cache behavior
pub conn_track_key: Option<String>, pub conn_track_key: Option<String>,
pub rate_limit_keys: Vec<RateLimitKey>, pub rate_limit_keys: Vec<RateLimitKey>,
pub chain_type: ChainType, pub chain_type: ChainType,
pub acl_result: Option<AclResult>, pub acl_result: Option<AclResult>,
pub rule_stats_vec: Vec<Arc<RuleStatsTracker>>, pub rule_stats_vec: Vec<Arc<RuleStats>>,
} }
// Packet info extracted for ACL processing // Packet info extracted for ACL processing
@@ -246,7 +211,7 @@ pub struct AclProcessor {
default_outbound_action: Action, default_outbound_action: Action,
default_forward_action: Action, default_forward_action: Action,
default_rule_stats: Arc<RuleStatsTracker>, default_rule_stats: Arc<RuleStats>,
// Connection tracking table - shared across different processor instances if needed // Connection tracking table - shared across different processor instances if needed
conn_track: Arc<DashMap<String, ConnTrackEntry>>, conn_track: Arc<DashMap<String, ConnTrackEntry>>,
@@ -259,13 +224,6 @@ pub struct AclProcessor {
cache_max_size: usize, cache_max_size: usize,
cache_cleanup_interval: Duration, cache_cleanup_interval: Duration,
// Coarse monotonic timestamp updated by the cleanup task, used to avoid
// calling Instant::now() on every cache hit.
coarse_millis: Arc<AtomicU64>,
// Hot-path counters that bypass the DashMap stats table
cache_hits: AtomicU64,
// Statistics // Statistics
stats: Arc<DashMap<AclStatKey, u64>>, stats: Arc<DashMap<AclStatKey, u64>>,
@@ -301,14 +259,18 @@ impl AclProcessor {
default_outbound_action, default_outbound_action,
default_forward_action, default_forward_action,
default_rule_stats: Arc::new(RuleStatsTracker::new(None)), default_rule_stats: Arc::new(RuleStats {
rule: None,
stat: Some(StatItem {
packet_count: 0,
byte_count: 0,
}),
}),
conn_track: conn_track.unwrap_or_else(|| Arc::new(DashMap::new())), conn_track: conn_track.unwrap_or_else(|| Arc::new(DashMap::new())),
rate_limiters: rate_limiters.unwrap_or_else(|| Arc::new(DashMap::new())), rate_limiters: rate_limiters.unwrap_or_else(|| Arc::new(DashMap::new())),
rule_cache: Arc::new(DashMap::new()), // Always start with fresh cache rule_cache: Arc::new(DashMap::new()), // Always start with fresh cache
cache_max_size: 1024, // Limit cache to 1k entries cache_max_size: 1024, // Limit cache to 1k entries
cache_cleanup_interval: Duration::from_secs(20), // Cleanup every 5 minutes cache_cleanup_interval: Duration::from_secs(20), // Cleanup every 5 minutes
coarse_millis: Arc::new(AtomicU64::new(0)),
cache_hits: AtomicU64::new(0),
stats: stats.unwrap_or_else(|| Arc::new(DashMap::new())), stats: stats.unwrap_or_else(|| Arc::new(DashMap::new())),
tasks, tasks,
}; };
@@ -412,14 +374,11 @@ impl AclProcessor {
let rule_cache = self.rule_cache.clone(); let rule_cache = self.rule_cache.clone();
let cache_max_size = self.cache_max_size; let cache_max_size = self.cache_max_size;
let cleanup_interval = self.cache_cleanup_interval; let cleanup_interval = self.cache_cleanup_interval;
let coarse_millis = self.coarse_millis.clone();
self.tasks.spawn(async move { self.tasks.spawn(async move {
let mut interval = tokio::time::interval(cleanup_interval); let mut interval = tokio::time::interval(cleanup_interval);
loop { loop {
interval.tick().await; interval.tick().await;
let now = crate::common::stats_manager::now_monotonic_millis();
coarse_millis.store(now, Relaxed);
Self::cleanup_cache(&rule_cache, cache_max_size); Self::cleanup_cache(&rule_cache, cache_max_size);
rule_cache.shrink_to_fit(); rule_cache.shrink_to_fit();
@@ -442,9 +401,10 @@ impl AclProcessor {
/// Clean up cache using LRU strategy /// Clean up cache using LRU strategy
fn cleanup_cache(cache: &DashMap<AclCacheKey, AclCacheEntry>, max_size: usize) { fn cleanup_cache(cache: &DashMap<AclCacheKey, AclCacheEntry>, max_size: usize) {
// remove cache not be used in last 15 second // remove cache not be used in last 15 second
let now = crate::common::stats_manager::now_monotonic_millis(); let expired_timepoint = Instant::now()
let cutoff = now.saturating_sub(15_000); .checked_sub(Duration::from_secs(15))
cache.retain(|_, entry| entry.last_access.load(Relaxed) > cutoff); .unwrap_or(Instant::now());
cache.retain(|_, entry| entry.last_access > expired_timepoint);
let current_size = cache.len(); let current_size = cache.len();
if current_size <= max_size { if current_size <= max_size {
@@ -452,9 +412,9 @@ impl AclProcessor {
} }
// Remove oldest entries (LRU cleanup) // Remove oldest entries (LRU cleanup)
let mut entries: Vec<(AclCacheKey, u64)> = cache let mut entries: Vec<(AclCacheKey, Instant)> = cache
.iter() .iter()
.map(|entry| (entry.key().clone(), entry.value().last_access.load(Relaxed))) .map(|entry| (entry.key().clone(), entry.value().last_access))
.collect(); .collect();
// Sort by last_access (oldest first) // Sort by last_access (oldest first)
@@ -473,7 +433,6 @@ impl AclProcessor {
); );
} }
#[cfg_attr(feature = "hotpath", hotpath::measure(impl_type = "AclProcessor"))]
pub(crate) fn process_packet_with_cache_entry( pub(crate) fn process_packet_with_cache_entry(
&self, &self,
packet_info: &PacketInfo, packet_info: &PacketInfo,
@@ -500,41 +459,42 @@ impl AclProcessor {
cache_entry.acl_result.clone().unwrap() cache_entry.acl_result.clone().unwrap()
} }
#[cfg_attr(feature = "hotpath", hotpath::measure(impl_type = "AclProcessor"))]
fn inc_cache_entry_stats(&self, cache_entry: &AclCacheEntry, packet_info: &PacketInfo) { fn inc_cache_entry_stats(&self, cache_entry: &AclCacheEntry, packet_info: &PacketInfo) {
for rule_stats in cache_entry.rule_stats_vec.iter() { for rule_stats in cache_entry.rule_stats_vec.iter() {
rule_stats.increment(packet_info.packet_size); // Use unsafe code to mutate the contents behind the Arc
let stat_ptr = rule_stats.stat.as_ref().unwrap() as *const StatItem as *mut StatItem;
unsafe {
(*stat_ptr).packet_count += 1;
(*stat_ptr).byte_count += packet_info.packet_size as u64;
}
} }
} }
pub fn get_rules_stats(&self) -> Vec<RuleStats> { pub fn get_rules_stats(&self) -> Vec<RuleStats> {
let mut stats: Vec<RuleStats> = Vec::new(); let mut stats: Vec<RuleStats> = Vec::new();
for rule in self.inbound_rules.iter() { for rule in self.inbound_rules.iter() {
stats.push(rule.rule_stats.snapshot()); stats.push((*rule.rule_stats).clone());
} }
for rule in self.outbound_rules.iter() { for rule in self.outbound_rules.iter() {
stats.push(rule.rule_stats.snapshot()); stats.push((*rule.rule_stats).clone());
} }
for rule in self.forward_rules.iter() { for rule in self.forward_rules.iter() {
stats.push(rule.rule_stats.snapshot()); stats.push((*rule.rule_stats).clone());
} }
stats stats
} }
/// Process a packet through ACL rules - Now lock-free! /// Process a packet through ACL rules - Now lock-free!
#[cfg_attr(feature = "hotpath", hotpath::measure(impl_type = "AclProcessor"))]
pub fn process_packet(&self, packet_info: &PacketInfo, chain_type: ChainType) -> AclResult { pub fn process_packet(&self, packet_info: &PacketInfo, chain_type: ChainType) -> AclResult {
// Check cache first for performance // Check cache first for performance
let cache_key = AclCacheKey::from_packet_info(packet_info, chain_type); let cache_key = AclCacheKey::from_packet_info(packet_info, chain_type);
// If cache hit and can skip checks, return cached result. // If cache hit and can skip checks, return cached result
// Use get() (read lock) instead of get_mut() (write lock) and update if let Some(mut cached) = self.rule_cache.get_mut(&cache_key) {
// last_access via AtomicU64, avoiding expensive Instant::now(). // Update last access time for LRU
if let Some(cached) = self.rule_cache.get(&cache_key) { cached.last_access = Instant::now();
cached
.last_access self.increment_stat(AclStatKey::CacheHits);
.store(self.coarse_millis.load(Relaxed), Relaxed);
self.cache_hits.fetch_add(1, Relaxed);
return self.process_packet_with_cache_entry(packet_info, &cached); return self.process_packet_with_cache_entry(packet_info, &cached);
} }
@@ -556,7 +516,7 @@ impl AclProcessor {
let mut cache_entry = AclCacheEntry { let mut cache_entry = AclCacheEntry {
action: Action::Allow, action: Action::Allow,
matched_rule: RuleId::Default, matched_rule: RuleId::Default,
last_access: AtomicU64::new(self.coarse_millis.load(Relaxed)), last_access: Instant::now(),
conn_track_key: None, conn_track_key: None,
rate_limit_keys: vec![], rate_limit_keys: vec![],
chain_type, chain_type,
@@ -621,9 +581,8 @@ impl AclProcessor {
// Cache the result with rule info // Cache the result with rule info
self.increment_stat(AclStatKey::RuleMatches); self.increment_stat(AclStatKey::RuleMatches);
self.inc_cache_entry_stats(&cache_entry, packet_info); self.inc_cache_entry_stats(&cache_entry, packet_info);
let result = cache_entry.acl_result.clone().unwrap(); self.cache_result(&cache_key, cache_entry.clone());
self.cache_result(&cache_key, cache_entry); return cache_entry.acl_result.clone().unwrap();
return result;
} }
let default_action = match chain_type { let default_action = match chain_type {
@@ -659,9 +618,8 @@ impl AclProcessor {
// Cache the default result (no rule info) // Cache the default result (no rule info)
self.inc_cache_entry_stats(&cache_entry, packet_info); self.inc_cache_entry_stats(&cache_entry, packet_info);
let result = cache_entry.acl_result.clone().unwrap(); self.cache_result(&cache_key, cache_entry.clone());
self.cache_result(&cache_key, cache_entry); cache_entry.acl_result.clone().unwrap()
result
} }
/// Get shared state for preserving across hot reloads /// Get shared state for preserving across hot reloads
@@ -785,11 +743,13 @@ impl AclProcessor {
/// Check connection state for stateful rules /// Check connection state for stateful rules
fn check_connection_state(&self, conn_track_key: &str, packet_info: &PacketInfo) { fn check_connection_state(&self, conn_track_key: &str, packet_info: &PacketInfo) {
let now = current_unix_secs();
self.conn_track self.conn_track
.entry(conn_track_key.to_string()) .entry(conn_track_key.to_string())
.and_modify(|x| { .and_modify(|x| {
x.last_seen = now; x.last_seen = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
x.packet_count += 1; x.packet_count += 1;
x.byte_count += packet_info.packet_size as u64; x.byte_count += packet_info.packet_size as u64;
x.state = ConnState::Established as i32; x.state = ConnState::Established as i32;
@@ -803,8 +763,14 @@ impl AclProcessor {
), ),
protocol: packet_info.protocol as i32, protocol: packet_info.protocol as i32,
state: ConnState::New as i32, state: ConnState::New as i32,
created_at: now, created_at: SystemTime::now()
last_seen: now, .duration_since(UNIX_EPOCH)
.unwrap()
.as_secs(),
last_seen: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs(),
packet_count: 1, packet_count: 1,
byte_count: packet_info.packet_size as u64, byte_count: packet_info.packet_size as u64,
}); });
@@ -896,7 +862,13 @@ impl AclProcessor {
stateful: rule.stateful, stateful: rule.stateful,
rate_limit: rule.rate_limit, rate_limit: rule.rate_limit,
burst_limit: rule.burst_limit, burst_limit: rule.burst_limit,
rule_stats: Arc::new(RuleStatsTracker::new(Some(rule.clone()))), rule_stats: Arc::new(RuleStats {
rule: Some(rule.clone()),
stat: Some(StatItem {
packet_count: 0,
byte_count: 0,
}),
}),
} }
} }
@@ -922,10 +894,6 @@ impl AclProcessor {
.collect::<HashMap<_, _>>(); .collect::<HashMap<_, _>>();
// Add cache statistics using enum keys // Add cache statistics using enum keys
stats.insert(
AclStatKey::CacheHits.as_str(),
self.cache_hits.load(Relaxed),
);
stats.insert(AclStatKey::CacheSize.as_str(), self.rule_cache.len() as u64); stats.insert(AclStatKey::CacheSize.as_str(), self.rule_cache.len() as u64);
stats.insert( stats.insert(
AclStatKey::CacheMaxSize.as_str(), AclStatKey::CacheMaxSize.as_str(),
@@ -940,11 +908,14 @@ impl AclProcessor {
conn_track: Arc<DashMap<String, ConnTrackEntry>>, conn_track: Arc<DashMap<String, ConnTrackEntry>>,
timeout_secs: u64, timeout_secs: u64,
) { ) {
let current_time = current_unix_secs(); let current_time = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let keys_to_remove: Vec<String> = conn_track let keys_to_remove: Vec<String> = conn_track
.iter() .iter()
.filter_map(|entry| { .filter_map(|entry| {
if current_time.saturating_sub(entry.last_seen) > timeout_secs { if current_time - entry.last_seen > timeout_secs {
Some(entry.key().clone()) Some(entry.key().clone())
} else { } else {
None None
@@ -959,7 +930,11 @@ impl AclProcessor {
/// Get cache hit rate /// Get cache hit rate
pub fn get_cache_hit_rate(&self) -> f64 { pub fn get_cache_hit_rate(&self) -> f64 {
let cache_hits = self.cache_hits.load(Relaxed); let cache_hits = self
.stats
.get(&AclStatKey::CacheHits)
.map(|v| *v.value())
.unwrap_or(0);
let total_requests = cache_hits let total_requests = cache_hits
+ self + self
.stats .stats
@@ -975,13 +950,6 @@ impl AclProcessor {
} }
} }
fn current_unix_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
}
// 新增辅助函数 // 新增辅助函数
fn parse_port_start(port_strs: &[String]) -> Option<u16> { fn parse_port_start(port_strs: &[String]) -> Option<u16> {
port_strs port_strs
-1
View File
@@ -22,7 +22,6 @@ pub mod machine_id;
pub mod netns; pub mod netns;
pub mod network; pub mod network;
pub mod os_info; pub mod os_info;
pub mod sharded_counter;
pub mod stats_manager; pub mod stats_manager;
pub mod stun; pub mod stun;
pub mod stun_codec_ext; pub mod stun_codec_ext;
-110
View File
@@ -1,110 +0,0 @@
use std::sync::atomic::{AtomicU64, Ordering::Relaxed};
const PUBLISH_INTERVAL: u64 = 256;
/// Counter optimized for long-lived worker threads.
///
/// Pending values below `PUBLISH_INTERVAL` live in `ThreadLocal` shards. The
/// shards are retained until this counter is dropped, so this is intended for
/// tokio workers or similarly long-lived threads rather than high-churn threads.
pub struct ShardedCounter {
published: AtomicU64,
locals: thread_local::ThreadLocal<AtomicU64>,
}
impl ShardedCounter {
pub fn new() -> Self {
Self {
published: AtomicU64::new(0),
locals: thread_local::ThreadLocal::new(),
}
}
#[inline]
#[cfg_attr(feature = "hotpath", hotpath::measure(impl_type = "ShardedCounter"))]
pub fn add(&self, delta: u64) {
let local = self.locals.get_or(|| AtomicU64::new(0));
let v = local.load(Relaxed).saturating_add(delta);
local.store(v, Relaxed);
if v >= PUBLISH_INTERVAL {
let pending = local.swap(0, Relaxed);
if pending > 0 {
self.published.fetch_add(pending, Relaxed);
}
}
}
#[inline]
pub fn inc(&self) {
self.add(1);
}
pub fn get(&self) -> u64 {
self.locals
.iter()
.fold(self.published.load(Relaxed), |total, local| {
total.saturating_add(local.load(Relaxed))
})
}
pub fn reset(&self) {
self.published.store(0, Relaxed);
for local in self.locals.iter() {
local.store(0, Relaxed);
}
}
}
impl Default for ShardedCounter {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for ShardedCounter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ShardedCounter")
.field("value", &self.get())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{sync::Arc, thread};
#[test]
fn sharded_counter_get_includes_other_thread_locals() {
let counter = Arc::new(ShardedCounter::new());
let thread_counter = Arc::clone(&counter);
thread::spawn(move || {
for _ in 0..10 {
thread_counter.inc();
}
})
.join()
.unwrap();
assert_eq!(counter.get(), 10);
}
#[test]
fn sharded_counter_reset_clears_other_thread_locals() {
let counter = Arc::new(ShardedCounter::new());
let thread_counter = Arc::clone(&counter);
thread::spawn(move || {
for _ in 0..10 {
thread_counter.inc();
}
})
.join()
.unwrap();
counter.reset();
assert_eq!(counter.get(), 0);
}
}
+150 -55
View File
@@ -1,16 +1,13 @@
use crate::common::sharded_counter::ShardedCounter;
use dashmap::DashMap; use dashmap::DashMap;
use hotpath::instant::Instant; use hotpath::instant::Instant;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::cell::UnsafeCell;
use std::fmt; use std::fmt;
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc;
use std::sync::{Arc, LazyLock};
use std::time::Duration; use std::time::Duration;
use tokio::time::interval; use tokio::time::interval;
use tokio_util::task::AbortOnDropHandle; use tokio_util::task::AbortOnDropHandle;
static START_INSTANT: LazyLock<Instant> = LazyLock::new(Instant::now);
/// Predefined metric names for type safety /// Predefined metric names for type safety
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum MetricName { pub enum MetricName {
@@ -378,10 +375,10 @@ impl Default for LabelSet {
} }
} }
/// High-performance counter backed by sharded thread-local accumulation. /// UnsafeCounter provides a high-performance counter using UnsafeCell
#[derive(Debug)] #[derive(Debug)]
pub struct UnsafeCounter { pub struct UnsafeCounter {
inner: ShardedCounter, value: UnsafeCell<u64>,
} }
impl Default for UnsafeCounter { impl Default for UnsafeCounter {
@@ -393,56 +390,121 @@ impl Default for UnsafeCounter {
impl UnsafeCounter { impl UnsafeCounter {
pub fn new() -> Self { pub fn new() -> Self {
Self { Self {
inner: ShardedCounter::new(), value: UnsafeCell::new(0),
} }
} }
pub fn add(&self, delta: u64) { pub fn new_with_value(initial: u64) -> Self {
self.inner.add(delta); Self {
value: UnsafeCell::new(initial),
}
} }
pub fn inc(&self) { /// Increment the counter by the given amount
self.inner.inc(); /// # 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 get(&self) -> u64 { /// Increment the counter by 1
self.inner.get() /// # 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 reset(&self) { /// Get the current value of the counter
self.inner.reset(); /// # 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 }
}
/// 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;
}
}
/// 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;
}
} }
} }
// 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 /// MetricData contains both the counter and last update timestamp
/// Uses UnsafeCell for lock-free access
#[derive(Debug)] #[derive(Debug)]
struct MetricData { struct MetricData {
counter: UnsafeCounter, counter: UnsafeCounter,
last_updated: AtomicU64, last_updated: UnsafeCell<Instant>,
}
pub(crate) fn now_monotonic_millis() -> u64 {
Instant::now().duration_since(*START_INSTANT).as_millis() as u64
} }
impl MetricData { impl MetricData {
fn new() -> Self { fn new() -> Self {
Self { Self {
counter: UnsafeCounter::new(), counter: UnsafeCounter::new(),
last_updated: AtomicU64::new(now_monotonic_millis()), last_updated: UnsafeCell::new(Instant::now()),
} }
} }
fn touch(&self) { fn new_with_value(initial: u64) -> Self {
self.last_updated Self {
.store(now_monotonic_millis(), Ordering::Relaxed); counter: UnsafeCounter::new_with_value(initial),
last_updated: UnsafeCell::new(Instant::now()),
}
} }
fn get_last_updated(&self) -> u64 { /// Update the last_updated timestamp
self.last_updated.load(Ordering::Relaxed) /// # 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();
}
}
/// 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 }
} }
} }
// 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 /// MetricKey uniquely identifies a metric with its name and labels
#[derive(Debug, Clone, PartialEq, Eq, Hash)] #[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct MetricKey { struct MetricKey {
@@ -483,23 +545,41 @@ impl CounterHandle {
} }
} }
/// Increment the counter by the given amount
pub fn add(&self, delta: u64) { pub fn add(&self, delta: u64) {
self.metric_data.counter.add(delta); unsafe {
self.metric_data.touch(); self.metric_data.counter.add(delta);
self.metric_data.touch();
}
} }
/// Increment the counter by 1
pub fn inc(&self) { pub fn inc(&self) {
self.metric_data.counter.inc(); unsafe {
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 { pub fn get(&self) -> u64 {
self.metric_data.counter.get() unsafe { self.metric_data.counter.get() }
} }
/// Reset the counter to zero
pub fn reset(&self) { pub fn reset(&self) {
self.metric_data.counter.reset(); unsafe {
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();
}
} }
} }
@@ -535,7 +615,9 @@ impl StatsManager {
loop { loop {
interval.tick().await; interval.tick().await;
let cutoff_millis = now_monotonic_millis().saturating_sub(180_000); let Some(cutoff_time) = Instant::now().checked_sub(Duration::from_secs(180)) else {
continue;
};
let Some(counters) = counters_clone.upgrade() else { let Some(counters) = counters_clone.upgrade() else {
break; break;
@@ -543,7 +625,7 @@ impl StatsManager {
counters.retain(|_, metric_data: &mut Arc<MetricData>| { counters.retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1 Arc::strong_count(metric_data) > 1
|| metric_data.get_last_updated() >= cutoff_millis || unsafe { metric_data.get_last_updated() > cutoff_time }
}); });
counters.shrink_to_fit(); counters.shrink_to_fit();
} }
@@ -581,7 +663,7 @@ impl StatsManager {
let key = entry.key(); let key = entry.key();
let metric_data = entry.value(); let metric_data = entry.value();
let value = metric_data.counter.get(); let value = unsafe { metric_data.counter.get() };
metrics.push(MetricSnapshot { metrics.push(MetricSnapshot {
name: key.name, name: key.name,
@@ -614,7 +696,7 @@ impl StatsManager {
let key = MetricKey::new(name, labels.clone()); let key = MetricKey::new(name, labels.clone());
if let Some(metric_data) = self.counters.get(&key) { if let Some(metric_data) = self.counters.get(&key) {
let value = metric_data.counter.get(); let value = unsafe { metric_data.counter.get() };
Some(MetricSnapshot { Some(MetricSnapshot {
name, name,
labels: labels.clone(), labels: labels.clone(),
@@ -715,9 +797,17 @@ mod tests {
async fn test_unsafe_counter() { async fn test_unsafe_counter() {
let counter = UnsafeCounter::new(); let counter = UnsafeCounter::new();
assert_eq!(counter.get(), 0); unsafe {
counter.add(256); assert_eq!(counter.get(), 0);
assert_eq!(counter.get(), 256); 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] #[tokio::test]
@@ -763,11 +853,11 @@ mod tests {
let stats = StatsManager::new(); let stats = StatsManager::new();
let counter1 = stats.get_simple_counter(MetricName::TrafficBytesTx); let counter1 = stats.get_simple_counter(MetricName::TrafficBytesTx);
counter1.add(100); counter1.set(100);
let labels = LabelSet::new().with_label("status", "success"); let labels = LabelSet::new().with_label("status", "success");
let counter2 = stats.get_counter(MetricName::PeerRpcClientTx, labels); let counter2 = stats.get_counter(MetricName::PeerRpcClientTx, labels);
counter2.add(50); counter2.set(50);
let traffic_labels = LabelSet::new() let traffic_labels = LabelSet::new()
.with_label_type(LabelType::NetworkName("default".to_string())) .with_label_type(LabelType::NetworkName("default".to_string()))
@@ -775,7 +865,7 @@ mod tests {
"87ede5a2-9c3d-492d-9bbe-989b9d07e742".to_string(), "87ede5a2-9c3d-492d-9bbe-989b9d07e742".to_string(),
)); ));
let counter3 = stats.get_counter(MetricName::TrafficBytesTxByInstance, traffic_labels); let counter3 = stats.get_counter(MetricName::TrafficBytesTxByInstance, traffic_labels);
counter3.add(25); counter3.set(25);
let prometheus_output = stats.export_prometheus(); let prometheus_output = stats.export_prometheus();
@@ -795,7 +885,7 @@ mod tests {
let labels = LabelSet::new().with_label("peer", "test"); let labels = LabelSet::new().with_label("peer", "test");
let counter = stats.get_counter(MetricName::PeerRpcClientTx, labels.clone()); let counter = stats.get_counter(MetricName::PeerRpcClientTx, labels.clone());
counter.add(42); counter.set(42);
let metric = stats let metric = stats
.get_metric(MetricName::PeerRpcClientTx, &labels) .get_metric(MetricName::PeerRpcClientTx, &labels)
@@ -812,11 +902,11 @@ mod tests {
stats stats
.get_simple_counter(MetricName::PeerRpcClientTx) .get_simple_counter(MetricName::PeerRpcClientTx)
.add(10); .set(10);
stats.get_simple_counter(MetricName::PeerRpcErrors).add(2); stats.get_simple_counter(MetricName::PeerRpcErrors).set(2);
stats stats
.get_simple_counter(MetricName::TrafficBytesTx) .get_simple_counter(MetricName::TrafficBytesTx)
.add(100); .set(100);
let rpc_metrics = stats.get_metrics_by_prefix("peer_rpc"); let rpc_metrics = stats.get_metrics_by_prefix("peer_rpc");
assert_eq!(rpc_metrics.len(), 2); assert_eq!(rpc_metrics.len(), 2);
@@ -831,15 +921,19 @@ mod tests {
// 创建一些计数器 // 创建一些计数器
let counter1 = stats.get_simple_counter(MetricName::PeerRpcClientTx); let counter1 = stats.get_simple_counter(MetricName::PeerRpcClientTx);
counter1.add(10); counter1.set(10);
let labels = LabelSet::new().with_label("test", "value"); let labels = LabelSet::new().with_label("test", "value");
let counter2 = stats.get_counter(MetricName::TrafficBytesTx, labels); let counter2 = stats.get_counter(MetricName::TrafficBytesTx, labels);
counter2.add(20); counter2.set(20);
// 验证计数器存在 // 验证计数器存在
assert_eq!(stats.metric_count(), 2); assert_eq!(stats.metric_count(), 2);
// 注意:实际的清理测试需要等待3分钟,这在单元测试中不现实
// 这里我们只验证清理机制的基本结构是否正确
// 清理逻辑在后台线程中运行,会自动删除超过3分钟未更新的条目
// 验证计数器仍然可以正常工作 // 验证计数器仍然可以正常工作
counter1.inc(); counter1.inc();
assert_eq!(counter1.get(), 11); assert_eq!(counter1.get(), 11);
@@ -852,14 +946,14 @@ mod tests {
async fn test_cleanup_keeps_metrics_with_live_handles() { async fn test_cleanup_keeps_metrics_with_live_handles() {
let stats = StatsManager::new(); let stats = StatsManager::new();
let counter = stats.get_simple_counter(MetricName::TrafficBytesForwarded); let counter = stats.get_simple_counter(MetricName::TrafficBytesForwarded);
counter.add(1); counter.set(1);
// Use a future cutoff so last_updated check always fails let cutoff_time = Instant::now().checked_add(Duration::from_secs(1)).unwrap();
let future_cutoff = now_monotonic_millis() + 1000;
stats stats
.counters .counters
.retain(|_, metric_data: &mut Arc<MetricData>| { .retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > future_cutoff Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
}); });
assert_eq!(stats.metric_count(), 1); assert_eq!(stats.metric_count(), 1);
@@ -869,7 +963,8 @@ mod tests {
stats stats
.counters .counters
.retain(|_, metric_data: &mut Arc<MetricData>| { .retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > future_cutoff Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
}); });
assert_eq!(stats.metric_count(), 0); assert_eq!(stats.metric_count(), 0);
} }
+6 -2
View File
@@ -998,9 +998,13 @@ mod tests {
s.sync_root_key(root_key, 2, 2, true); s.sync_root_key(root_key, 2, 2, true);
assert!(s.check_replay_for_test(2, 0, SecureDatagramDirection::AToB, now + 2)); assert!(s.check_replay_for_test(2, 0, SecureDatagramDirection::AToB, now + 2));
let expires_at = s.sync_rx_grace_expires_at_ms.load(Ordering::Relaxed);
assert!(!s.check_replay_for_test(0, 1, SecureDatagramDirection::AToB, expires_at + 1)); assert!(!s.check_replay_for_test(
0,
1,
SecureDatagramDirection::AToB,
now + SecureDatagramSession::SYNC_RX_GRACE_AFTER_MS + 3
));
} }
#[test] #[test]
+19 -72
View File
@@ -483,16 +483,8 @@ impl ZCPacket {
let payload_off = ret.packet_type.get_packet_offsets().payload_offset; let payload_off = ret.packet_type.get_packet_offsets().payload_offset;
let total_len = payload_off + payload.len(); let total_len = payload_off + payload.len();
ret.inner.reserve(total_len); ret.inner.reserve(total_len);
unsafe { ret.inner.set_len(total_len) };
// SAFETY: `reserve` guarantees capacity >= total_len. ret.mut_payload()[..payload.len()].copy_from_slice(payload);
// We zero the header region and copy payload before advancing length,
// so every byte in [0..total_len) is initialized before any read.
unsafe {
let ptr = ret.inner.as_mut_ptr();
std::ptr::write_bytes(ptr, 0, payload_off);
std::ptr::copy_nonoverlapping(payload.as_ptr(), ptr.add(payload_off), payload.len());
ret.inner.set_len(total_len);
}
ret ret
} }
@@ -500,12 +492,12 @@ impl ZCPacket {
let mut ret = Self::new_nic_packet(); let mut ret = Self::new_nic_packet();
ret.inner.reserve(cap); ret.inner.reserve(cap);
let total_len = ret.packet_type.get_packet_offsets().payload_offset - packet_info_len; let total_len = ret.packet_type.get_packet_offsets().payload_offset - packet_info_len;
ret.inner.resize(total_len, 0); unsafe { ret.inner.set_len(total_len) };
ret ret
} }
pub fn new_for_foreign_network( pub fn new_for_foreign_network(
network_name: &str, network_name: &String,
dst_peer_id: u32, dst_peer_id: u32,
foreign_zc_packet: &ZCPacket, foreign_zc_packet: &ZCPacket,
) -> Self { ) -> Self {
@@ -514,71 +506,26 @@ impl ZCPacket {
foreign_network_hdr.get_header_len() + foreign_zc_packet.tunnel_payload().len(); foreign_network_hdr.get_header_len() + foreign_zc_packet.tunnel_payload().len();
let mut ret = Self::new_nic_packet(); let mut ret = Self::new_nic_packet();
let offsets = ret.packet_type.get_packet_offsets(); let payload_off = ret.packet_type.get_packet_offsets().payload_offset;
let payload_off = offsets.payload_offset; ret.inner.reserve(payload_off + total_payload_len);
let pm_hdr_off = offsets.peer_manager_header_offset; unsafe { ret.inner.set_len(payload_off + total_payload_len) };
let total_len = payload_off + total_payload_len;
ret.inner.reserve(total_len);
let fixed_hdr_len = std::mem::size_of::<ForeignNetworkPacketHeader>(); let fixed_hdr_len = std::mem::size_of::<ForeignNetworkPacketHeader>();
ret.mut_payload()[..fixed_hdr_len].copy_from_slice(foreign_network_hdr.as_bytes());
let name_offset = foreign_network_hdr.network_name_offset.get() as usize; let name_offset = foreign_network_hdr.network_name_offset.get() as usize;
let name_len = foreign_network_hdr.network_name_len.get() as usize; let name_len = foreign_network_hdr.network_name_len.get() as usize;
let foreign_payload = foreign_zc_packet.tunnel_payload(); ret.mut_payload()[name_offset..name_offset + name_len]
.copy_from_slice(network_name.as_bytes());
// Construct the PeerManagerHeader on the stack so we can write it ret.mut_payload()[foreign_network_hdr.get_header_len()..]
// directly into the buffer, avoiding a separate mut_peer_manager_header() .copy_from_slice(foreign_zc_packet.tunnel_payload());
// call after set_len.
let pm_hdr = PeerManagerHeader {
from_peer_id: 0.into(),
to_peer_id: 0.into(),
packet_type: PacketType::ForeignNetworkPacket as u8,
flags: 0,
forward_counter: 0,
reserved: 0,
len: U32::new(total_payload_len as u32),
};
// SAFETY: `reserve` guarantees capacity >= total_len. let hdr = ret.mut_peer_manager_header().unwrap();
// We zero only the tunnel-header reserved space [0..pm_hdr_off], write hdr.from_peer_id = 0.into();
// the PeerManagerHeader directly at pm_hdr_off, then copy the foreign hdr.to_peer_id = 0.into();
// network header, network name, and payload. Every byte in [0..total_len) hdr.packet_type = PacketType::ForeignNetworkPacket as u8;
// is initialized before set_len. hdr.len.set(total_payload_len as u32);
unsafe {
let ptr = ret.inner.as_mut_ptr();
// Zero the tunnel header reserved space only (not the PM header region)
std::ptr::write_bytes(ptr, 0, pm_hdr_off);
// Write PeerManagerHeader directly
std::ptr::copy_nonoverlapping(
pm_hdr.as_bytes().as_ptr(),
ptr.add(pm_hdr_off),
std::mem::size_of::<PeerManagerHeader>(),
);
// Copy foreign network fixed header
std::ptr::copy_nonoverlapping(
foreign_network_hdr.as_bytes().as_ptr(),
ptr.add(payload_off),
fixed_hdr_len,
);
// Copy network name
std::ptr::copy_nonoverlapping(
network_name.as_ptr(),
ptr.add(payload_off + name_offset),
name_len,
);
// Copy foreign payload
std::ptr::copy_nonoverlapping(
foreign_payload.as_ptr(),
ptr.add(payload_off + foreign_network_hdr.get_header_len()),
foreign_payload.len(),
);
ret.inner.set_len(total_len);
}
ret ret
} }
@@ -748,7 +695,7 @@ impl ZCPacket {
.get_packet_offsets() .get_packet_offsets()
.peer_manager_header_offset; .peer_manager_header_offset;
let mut buf = BytesMut::with_capacity(new_pm_offset + tunnel_payload.len()); let mut buf = BytesMut::with_capacity(new_pm_offset + tunnel_payload.len());
buf.resize(new_pm_offset, 0); unsafe { buf.set_len(new_pm_offset) };
buf.extend_from_slice(tunnel_payload); buf.extend_from_slice(tunnel_payload);
return Self::new_from_buf(buf, target_packet_type); return Self::new_from_buf(buf, target_packet_type);
} }
+32 -21
View File
@@ -1,4 +1,7 @@
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering::Relaxed}; use std::{
cell::UnsafeCell,
sync::atomic::{AtomicU32, Ordering::Relaxed},
};
pub struct WindowLatency { pub struct WindowLatency {
latency_us_window: Vec<AtomicU32>, latency_us_window: Vec<AtomicU32>,
@@ -60,30 +63,34 @@ impl WindowLatency {
#[derive(Debug)] #[derive(Debug)]
pub struct Throughput { pub struct Throughput {
tx_bytes: AtomicU64, tx_bytes: UnsafeCell<u64>,
rx_bytes: AtomicU64, rx_bytes: UnsafeCell<u64>,
tx_packets: AtomicU64, tx_packets: UnsafeCell<u64>,
rx_packets: AtomicU64, rx_packets: UnsafeCell<u64>,
} }
impl Clone for Throughput { impl Clone for Throughput {
fn clone(&self) -> Self { fn clone(&self) -> Self {
Self { Self {
tx_bytes: AtomicU64::new(self.tx_bytes.load(Relaxed)), tx_bytes: UnsafeCell::new(unsafe { *self.tx_bytes.get() }),
rx_bytes: AtomicU64::new(self.rx_bytes.load(Relaxed)), rx_bytes: UnsafeCell::new(unsafe { *self.rx_bytes.get() }),
tx_packets: AtomicU64::new(self.tx_packets.load(Relaxed)), tx_packets: UnsafeCell::new(unsafe { *self.tx_packets.get() }),
rx_packets: AtomicU64::new(self.rx_packets.load(Relaxed)), rx_packets: UnsafeCell::new(unsafe { *self.rx_packets.get() }),
} }
} }
} }
// add sync::Send and sync::Sync traits to Throughput
unsafe impl Send for Throughput {}
unsafe impl Sync for Throughput {}
impl Default for Throughput { impl Default for Throughput {
fn default() -> Self { fn default() -> Self {
Self { Self {
tx_bytes: AtomicU64::new(0), tx_bytes: UnsafeCell::new(0),
rx_bytes: AtomicU64::new(0), rx_bytes: UnsafeCell::new(0),
tx_packets: AtomicU64::new(0), tx_packets: UnsafeCell::new(0),
rx_packets: AtomicU64::new(0), rx_packets: UnsafeCell::new(0),
} }
} }
} }
@@ -94,28 +101,32 @@ impl Throughput {
} }
pub fn tx_bytes(&self) -> u64 { pub fn tx_bytes(&self) -> u64 {
self.tx_bytes.load(Relaxed) unsafe { *self.tx_bytes.get() }
} }
pub fn rx_bytes(&self) -> u64 { pub fn rx_bytes(&self) -> u64 {
self.rx_bytes.load(Relaxed) unsafe { *self.rx_bytes.get() }
} }
pub fn tx_packets(&self) -> u64 { pub fn tx_packets(&self) -> u64 {
self.tx_packets.load(Relaxed) unsafe { *self.tx_packets.get() }
} }
pub fn rx_packets(&self) -> u64 { pub fn rx_packets(&self) -> u64 {
self.rx_packets.load(Relaxed) unsafe { *self.rx_packets.get() }
} }
pub fn record_tx_bytes(&self, bytes: u64) { pub fn record_tx_bytes(&self, bytes: u64) {
self.tx_bytes.fetch_add(bytes, Relaxed); unsafe {
self.tx_packets.fetch_add(1, Relaxed); *self.tx_bytes.get() += bytes;
*self.tx_packets.get() += 1;
}
} }
pub fn record_rx_bytes(&self, bytes: u64) { pub fn record_rx_bytes(&self, bytes: u64) {
self.rx_bytes.fetch_add(bytes, Relaxed); unsafe {
self.rx_packets.fetch_add(1, Relaxed); *self.rx_bytes.get() += bytes;
*self.rx_packets.get() += 1;
}
} }
} }