Compare commits

..
Author SHA1 Message Date
fanyang b8d1d6b32c perf(smoltcp): unify device packet type on BytesMut for zero-copy TX
The smoltcp -> ZCPacket TX path still memcpy'd each outbound packet:
the device emitted Vec<u8>, and `Vec<u8> -> BytesMut` has no zero-copy
conversion in the bytes crate, so socks5/tcp_proxy paid a
`BytesMut::from(Bytes::from(data))` copy per packet.

Switch tokio_smoltcp's packet type to BytesMut end to end so smoltcp
writes into the headroom-reserved buffer that ZCPacket wraps directly.

- device: `Packet = BytesMut`; BufferTxToken allocates via
  BytesMut::with_capacity + resize; BufferRxToken derefs to &[u8].
- channel_device: Stream/Sink/channel carry BytesMut instead of Vec<u8>.
- socks5 / tcp_proxy: TX wraps the BytesMut via new_from_buf directly,
  removing the per-packet memcpy; inbound keeps an equivalent copy via
  BytesMut::from(payload).
- reactor: adapts through the Packet alias, no code change.
2026-06-27 12:30:35 +08:00
fanyang b092e78523 refactor(smoltcp): address review on NetConfig visibility and zcpacket bench
- tokio_smoltcp::NetConfig: narrow `packet_tx_headroom` to `pub(crate)`.
  The struct is already `#[non_exhaustive]` and the field is only read
  internally (BufferDevice creation); the `with_packet_tx_headroom`
  builder remains public.
- packet_def: harden the zcpacket benchmark by asserting copy/zerocopy
  payload equivalence outside the timed section, and black_box the
  constructed packet so the compiler cannot elide the construction.
2026-06-27 11:14:38 +08:00
fanyang 626a5fb4f1 perf: avoid smoltcp packet copy
Reserve NIC packet headroom in BufferDevice's TxToken so smoltcp writes
the IP packet directly into a buf already carrying the ZCPacket NIC
offset. socks5/tcp_proxy then wrap that buf zero-copy via
ZCPacket::new_from_buf instead of copying through new_with_payload.

Benchmark: tunnel::packet_def::tests::smoltcp_zcpacket_construct_bench
  copy (new_with_payload)    1280B: 13.5M pps
  zerocopy (new_from_buf)    1280B: 28.3M pps  (2.09x)
  copy (new_with_payload)    4096B: 10.7M pps
  zerocopy (new_from_buf)    4096B: 19.1M pps  (1.79x)

Environment: AMD Ryzen 9 9955HX, rustc 1.95.0, --release, median of 3 runs
2026-06-26 21:09:29 +08:00
fanyang 130d89a057 test: add smoltcp zcpacket construction benchmark
Benchmark the two ZCPacket construction paths used around the smoltcp
gateway: copy via new_with_payload (pre-f5ce0848) vs zero-copy via
new_from_buf with NIC headroom reserved (f5ce0848).
2026-06-26 21:09:29 +08:00
18 changed files with 2310 additions and 1771 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())
+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 {
+13 -6
View File
@@ -28,7 +28,7 @@ use crate::{
ip_reassembler::IpReassembler,
tokio_smoltcp::{BufferSize, Net, NetConfig, channel_device},
},
tunnel::packet_def::{PacketType, ZCPacket},
tunnel::packet_def::{PacketType, ZCPacket, ZCPacketType},
};
use anyhow::Context;
use dashmap::DashMap;
@@ -363,7 +363,10 @@ impl Socks5ServerNet {
let mut smoltcp_stack_receiver = packet_recv.lock().await;
while let Some(packet) = smoltcp_stack_receiver.recv().await {
tracing::trace!(?packet, "receive from peer send to smoltcp packet");
if let Err(e) = stack_sink.send(Ok(packet.payload().to_vec())).await {
if let Err(e) = stack_sink
.send(Ok(bytes::BytesMut::from(packet.payload())))
.await
{
tracing::error!("send to smoltcp stack failed: {:?}", e);
}
}
@@ -377,13 +380,16 @@ impl Socks5ServerNet {
"receive from smoltcp stack and send to peer mgr packet, len = {}",
data.len()
);
let Some(ipv4) = Ipv4Packet::new(&data) else {
tracing::error!(?data, "smoltcp stack stream get non ipv4 packet");
let packet = ZCPacket::new_from_buf(data, ZCPacketType::NIC);
let Some(ipv4) = Ipv4Packet::new(packet.payload()) else {
tracing::error!(
payload_len = packet.payload_len(),
"smoltcp stack stream get non ipv4 packet"
);
continue;
};
let dst = ipv4.get_destination();
let packet = ZCPacket::new_with_payload(&data);
let Some(peer_manager) = peer_manager.upgrade() else {
tracing::warn!("peer manager is gone, smoltcp sender exited");
return;
@@ -412,7 +418,8 @@ impl Socks5ServerNet {
tcp_tx_size: 1024 * 128,
..Default::default()
}),
),
)
.with_packet_tx_headroom(ZCPacketType::NIC.get_packet_offsets().payload_offset),
);
let forward_tasks = Arc::new(std::sync::Mutex::new(forward_tasks));
+14 -5
View File
@@ -39,6 +39,8 @@ use super::CidrSet;
#[cfg(feature = "smoltcp")]
use super::tokio_smoltcp::{self, Net, NetConfig, channel_device};
#[cfg(feature = "smoltcp")]
use crate::tunnel::packet_def::ZCPacketType;
#[async_trait::async_trait]
pub(crate) trait NatDstConnector: Send + Sync + Clone + 'static {
@@ -561,7 +563,10 @@ impl<C: NatDstConnector> TcpProxy<C> {
self.tasks.lock().unwrap().spawn(async move {
while let Some(packet) = smoltcp_stack_receiver.recv().await {
tracing::trace!(?packet, "receive from peer send to smoltcp packet");
if let Err(e) = stack_sink.send(Ok(packet.payload().to_vec())).await {
if let Err(e) = stack_sink
.send(Ok(bytes::BytesMut::from(packet.payload())))
.await
{
tracing::error!("send to smoltcp stack failed: {:?}", e);
}
}
@@ -575,13 +580,16 @@ impl<C: NatDstConnector> TcpProxy<C> {
?data,
"receive from smoltcp stack and send to peer mgr packet"
);
let Some(ipv4) = Ipv4Packet::new(&data) else {
tracing::error!(?data, "smoltcp stack stream get non ipv4 packet");
let packet = ZCPacket::new_from_buf(data, ZCPacketType::NIC);
let Some(ipv4) = Ipv4Packet::new(packet.payload()) else {
tracing::error!(
payload_len = packet.payload_len(),
"smoltcp stack stream get non ipv4 packet"
);
continue;
};
let dst = ipv4.get_destination();
let packet = ZCPacket::new_with_payload(&data);
let Some(peer_mgr) = peer_mgr.upgrade() else {
tracing::warn!("peer manager is gone, smoltcp sender exited");
return;
@@ -610,7 +618,8 @@ impl<C: NatDstConnector> TcpProxy<C> {
tcp_tx_size: 1024 * 16,
..Default::default()
}),
),
)
.with_packet_tx_headroom(ZCPacketType::NIC.get_packet_offsets().payload_offset),
);
net.set_any_ip(true);
self.smoltcp_net.lock().await.replace(net);
@@ -1,3 +1,4 @@
use bytes::BytesMut;
use futures::{Sink, Stream};
use smoltcp::phy::DeviceCapabilities;
use std::{
@@ -12,15 +13,15 @@ use super::device::AsyncDevice;
/// A device that send and receive packets using a channel.
pub struct ChannelDevice {
recv: Receiver<io::Result<Vec<u8>>>,
send: PollSender<Vec<u8>>,
recv: Receiver<io::Result<BytesMut>>,
send: PollSender<BytesMut>,
caps: DeviceCapabilities,
}
pub type ChannelDeviceNewRet = (
ChannelDevice,
Sender<io::Result<Vec<u8>>>,
Receiver<Vec<u8>>,
Sender<io::Result<BytesMut>>,
Receiver<BytesMut>,
);
impl ChannelDevice {
@@ -43,25 +44,25 @@ impl ChannelDevice {
}
impl Stream for ChannelDevice {
type Item = io::Result<Vec<u8>>;
type Item = io::Result<BytesMut>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.recv.poll_recv(cx)
}
}
fn map_err(e: PollSendError<Vec<u8>>) -> io::Error {
fn map_err(e: PollSendError<BytesMut>) -> io::Error {
io::Error::other(e)
}
impl Sink<Vec<u8>> for ChannelDevice {
impl Sink<BytesMut> for ChannelDevice {
type Error = io::Error;
fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.send.poll_reserve(cx).map_err(map_err)
}
fn start_send(mut self: Pin<&mut Self>, item: Vec<u8>) -> Result<(), Self::Error> {
fn start_send(mut self: Pin<&mut Self>, item: BytesMut) -> Result<(), Self::Error> {
self.send.send_item(item).map_err(map_err)
}
+35 -8
View File
@@ -1,3 +1,4 @@
use bytes::BytesMut;
use futures::{Sink, Stream};
pub use smoltcp::phy::DeviceCapabilities;
use smoltcp::{
@@ -10,7 +11,7 @@ use std::{collections::VecDeque, io};
pub const DEFAULT_MAX_BURST_SIZE: usize = 100;
/// A packet used in `AsyncDevice`.
pub type Packet = Vec<u8>;
pub type Packet = BytesMut;
/// A device that send and receive packets asynchronously.
pub trait AsyncDevice:
@@ -33,6 +34,7 @@ where
pub struct BufferDevice {
caps: DeviceCapabilities,
max_burst_size: usize,
tx_headroom: usize,
recv_queue: VecDeque<Packet>,
send_queue: VecDeque<Packet>,
}
@@ -41,13 +43,11 @@ pub struct BufferDevice {
pub struct BufferRxToken(Packet);
impl RxToken for BufferRxToken {
fn consume<R, F>(mut self, f: F) -> R
fn consume<R, F>(self, f: F) -> R
where
F: FnOnce(&[u8]) -> R,
{
let p = &mut self.0;
f(p)
f(&self.0[..])
}
}
@@ -59,8 +59,10 @@ impl<'d> TxToken for BufferTxToken<'d> {
where
F: FnOnce(&mut [u8]) -> R,
{
let mut buffer = vec![0u8; len];
let result = f(&mut buffer);
let tx_headroom = self.0.tx_headroom;
let mut buffer = BytesMut::with_capacity(tx_headroom + len);
buffer.resize(tx_headroom + len, 0);
let result = f(&mut buffer[tx_headroom..]);
self.0.send_queue.push_back(buffer);
@@ -98,11 +100,12 @@ impl Device for BufferDevice {
}
impl BufferDevice {
pub(crate) fn new(caps: DeviceCapabilities) -> BufferDevice {
pub(crate) fn new(caps: DeviceCapabilities, tx_headroom: usize) -> BufferDevice {
let max_burst_size = caps.max_burst_size.unwrap_or(DEFAULT_MAX_BURST_SIZE);
BufferDevice {
caps,
max_burst_size,
tx_headroom,
recv_queue: VecDeque::with_capacity(max_burst_size),
send_queue: VecDeque::with_capacity(max_burst_size),
}
@@ -123,3 +126,27 @@ impl BufferDevice {
self.recv_queue.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn buffer_device_reserves_tx_headroom() {
let mut caps = DeviceCapabilities::default();
caps.max_burst_size = Some(1);
let mut device = BufferDevice::new(caps, 16);
let token = device.transmit(Instant::now()).unwrap();
token.consume(4, |buf| {
assert_eq!(buf.len(), 4);
buf.copy_from_slice(&[1, 2, 3, 4]);
});
let mut queue = device.take_send_queue();
let packet = queue.pop_front().unwrap();
assert_eq!(packet.len(), 20);
assert_eq!(&packet[..16], &[0; 16]);
assert_eq!(&packet[16..], &[1, 2, 3, 4]);
}
}
+9 -1
View File
@@ -51,6 +51,7 @@ pub struct NetConfig {
pub ip_addr: IpCidr,
pub gateway: Vec<IpAddress>,
pub buffer_size: BufferSize,
pub(crate) packet_tx_headroom: usize,
}
impl NetConfig {
@@ -65,8 +66,14 @@ impl NetConfig {
ip_addr,
gateway,
buffer_size: buffer_size.unwrap_or_default(),
packet_tx_headroom: 0,
}
}
pub fn with_packet_tx_headroom(mut self, packet_tx_headroom: usize) -> Self {
self.packet_tx_headroom = packet_tx_headroom;
self
}
}
/// `Net` is the main interface to the network stack.
@@ -97,7 +104,8 @@ impl Net {
}
fn new2<D: device::AsyncDevice + 'static>(device: D, config: NetConfig) -> Net {
let mut buffer_device = BufferDevice::new(device.capabilities().clone());
let mut buffer_device =
BufferDevice::new(device.capabilities().clone(), config.packet_tx_headroom);
let mut iface = Interface::new(config.interface_config, &mut buffer_device, Instant::now());
let ip_addr = config.ip_addr;
iface.update_ip_addrs(|ip_addrs| {
+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]
+65
View File
@@ -764,6 +764,7 @@ impl ZCPacket {
#[cfg(test)]
mod tests {
use super::*;
use std::{hint::black_box, time::Instant};
#[test]
fn test_zc_packet() {
@@ -809,4 +810,68 @@ mod tests {
assert!(packet.mut_wg_tunnel_header().is_none());
}
fn bench_smoltcp_zcpacket_construct(payload_len: usize, iterations: usize) {
let nic_offset = ZCPacketType::NIC.get_packet_offsets().payload_offset;
// Correctness check (outside the timed section): both construction paths
// must yield equivalent payloads for the perf comparison to be meaningful.
{
let data = vec![7u8; payload_len];
let p_copy = ZCPacket::new_with_payload(&data);
let mut buf = BytesMut::with_capacity(nic_offset + payload_len);
buf.resize(nic_offset + payload_len, 0);
buf[nic_offset..].fill(7);
let p_zero = ZCPacket::new_from_buf(buf, ZCPacketType::NIC);
assert_eq!(p_copy.payload(), p_zero.payload());
}
// copy path: smoltcp emits a bare payload buf; socks5/tcp_proxy copy it
// via ZCPacket::new_with_payload (pre-f5ce0848 behavior).
let now = Instant::now();
let mut checksum = 0usize;
for _ in 0..iterations {
let data = vec![7u8; payload_len];
let p = ZCPacket::new_with_payload(black_box(&data));
// black_box forces the side-effect-free construction to be emitted;
// payload_len is stable per run so it cannot skew the numbers.
checksum = checksum.wrapping_add(black_box(&p).payload_len());
}
let copy_elapsed = now.elapsed().as_secs_f64();
// zerocopy path: device reserves NIC headroom in the buf; socks5/tcp_proxy
// wrap it zero-copy via ZCPacket::new_from_buf (f5ce0848 behavior).
let now = Instant::now();
let mut checksum2 = 0usize;
for _ in 0..iterations {
let mut buf = BytesMut::with_capacity(nic_offset + payload_len);
buf.resize(nic_offset + payload_len, 0);
buf[nic_offset..].fill(7);
let p = ZCPacket::new_from_buf(black_box(buf), ZCPacketType::NIC);
checksum2 = checksum2.wrapping_add(black_box(&p).payload_len());
}
let zerocopy_elapsed = now.elapsed().as_secs_f64();
println!(
"smoltcp_zcpacket payload_len={} iterations={} copy_pps={:.0} copy_bytes_per_sec={:.0} zerocopy_pps={:.0} zerocopy_bytes_per_sec={:.0} speedup={:.2}x checksums={}/{}",
payload_len,
iterations,
iterations as f64 / copy_elapsed,
(payload_len * iterations) as f64 / copy_elapsed,
iterations as f64 / zerocopy_elapsed,
(payload_len * iterations) as f64 / zerocopy_elapsed,
copy_elapsed / zerocopy_elapsed,
checksum,
checksum2
);
}
#[test]
#[ignore = "benchmark helper; run with --ignored --nocapture"]
fn smoltcp_zcpacket_construct_bench() {
bench_smoltcp_zcpacket_construct(1280, 1_000_000);
bench_smoltcp_zcpacket_construct(4096, 500_000);
}
}