mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-09-02 09:09:17 +00:00
refactor(core): migrate packet processing from pnet to smoltcp (#2456)
Replace pnet_packet parsing and mutation across gateway packet paths with the existing smoltcp wire APIs. Preserve length validation, fragmentation classification, TCP flags, and checksum behavior while removing the core pnet_packet feature dependency. Reject stale non-initiator OSPF sync sessions: only initiator requests may create missing sessions, and a rejection clears the old initiator role only when the remote session generation is unchanged. This fixes an unbounded RPC storm caused by a delayed route sync recreating a session after both peers relinquished the initiator role, with regression tests for session creation and response reordering.
This commit is contained in:
+22
-3
@@ -152,7 +152,14 @@ socket2 = { version = "0.5.10", features = ["all"] }
|
||||
rand = "0.8.5"
|
||||
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
pnet = { version = "0.35.0", features = ["serde"] }
|
||||
pnet_datalink = { version = "0.35.0", optional = true }
|
||||
smoltcp = { git = "https://github.com/smoltcp-rs/smoltcp.git", rev = "0a926767a68bc88d5512afefa7529c5ecdade4ea", optional = true, default-features = false, features = [
|
||||
"std",
|
||||
"medium-ethernet",
|
||||
"proto-ipv4",
|
||||
"proto-ipv6",
|
||||
"socket-raw",
|
||||
] }
|
||||
serde_json = "1"
|
||||
|
||||
clap = { version = "4.5.30", features = [
|
||||
@@ -167,7 +174,7 @@ clap_complete_nushell = { version = "4.5.10" }
|
||||
|
||||
async-recursion = "1.0.5"
|
||||
|
||||
network-interface = "2.0"
|
||||
network-interface = "2.0.5"
|
||||
|
||||
# for wireguard
|
||||
boringtun = { package = "boringtun-easytier", version = "0.6.1", optional = true }
|
||||
@@ -287,6 +294,13 @@ ctor = "0.8.0"
|
||||
stun_codec = "0.3.4"
|
||||
bytecodec = "0.4.15"
|
||||
x25519-dalek = { version = "2.0", features = ["static_secrets"] }
|
||||
smoltcp = { git = "https://github.com/smoltcp-rs/smoltcp.git", rev = "0a926767a68bc88d5512afefa7529c5ecdade4ea", default-features = false, features = [
|
||||
"std",
|
||||
"medium-ethernet",
|
||||
"proto-ipv4",
|
||||
"proto-ipv6",
|
||||
"socket-raw",
|
||||
] }
|
||||
|
||||
[target.'cfg(target_os = "linux")'.dev-dependencies]
|
||||
defguard_wireguard_rs = "0.4.2"
|
||||
@@ -375,7 +389,12 @@ magic-dns = [
|
||||
"easytier-core/proxy-packet",
|
||||
"easytier-proto/magic-dns",
|
||||
]
|
||||
faketcp = ["dep:flume", "easytier-proto/faketcp"]
|
||||
faketcp = [
|
||||
"dep:flume",
|
||||
"dep:pnet_datalink",
|
||||
"dep:smoltcp",
|
||||
"easytier-proto/faketcp",
|
||||
]
|
||||
zstd = ["easytier-core/zstd", "easytier-proto/zstd"]
|
||||
upnp = ["dep:igd-next", "dep:natpmp"]
|
||||
endpoint-discovery = [
|
||||
|
||||
@@ -11,6 +11,17 @@ use std::{
|
||||
os::fd::AsRawFd,
|
||||
};
|
||||
|
||||
use super::{
|
||||
Error, IfConfiguerTrait,
|
||||
netlink_wire::{
|
||||
AddressMessage, MessageBuilder, MessageIter, NLM_F_ACK, NLM_F_CREATE, NLM_F_DUMP,
|
||||
NLM_F_DUMP_INTR, NLM_F_EXCL, NLM_F_REQUEST, NLMSG_DONE, NLMSG_ERROR, NeighborMessage,
|
||||
NetlinkDecode, NetlinkEncode, RTM_DELADDR, RTM_DELNEIGH, RTM_DELROUTE, RTM_GETNEIGH,
|
||||
RTM_GETROUTE, RTM_NEWADDR, RTM_NEWNEIGH, RTM_NEWROUTE, RouteMessage, RouteMessageBuilder,
|
||||
RouteType, netlink_error_code,
|
||||
},
|
||||
};
|
||||
use crate::common::network::ip_mask_to_prefix;
|
||||
use anyhow::Context;
|
||||
use async_trait::async_trait;
|
||||
use cidr::{IpInet, Ipv4Inet, Ipv6Inet};
|
||||
@@ -23,18 +34,6 @@ use nix::{
|
||||
net::if_::InterfaceFlags,
|
||||
sys::socket::SockaddrLike as _,
|
||||
};
|
||||
use pnet::ipnetwork::ip_mask_to_prefix;
|
||||
|
||||
use super::{
|
||||
Error, IfConfiguerTrait,
|
||||
netlink_wire::{
|
||||
AddressMessage, MessageBuilder, MessageIter, NLM_F_ACK, NLM_F_CREATE, NLM_F_DUMP,
|
||||
NLM_F_DUMP_INTR, NLM_F_EXCL, NLM_F_REQUEST, NLMSG_DONE, NLMSG_ERROR, NeighborMessage,
|
||||
NetlinkDecode, NetlinkEncode, RTM_DELADDR, RTM_DELNEIGH, RTM_DELROUTE, RTM_GETNEIGH,
|
||||
RTM_GETROUTE, RTM_NEWADDR, RTM_NEWNEIGH, RTM_NEWROUTE, RouteMessage, RouteMessageBuilder,
|
||||
RouteType, netlink_error_code,
|
||||
},
|
||||
};
|
||||
|
||||
pub(crate) fn dummy_socket() -> Result<std::net::UdpSocket, Error> {
|
||||
Ok(std::net::UdpSocket::bind("0:0")?)
|
||||
|
||||
+199
-111
@@ -1,13 +1,6 @@
|
||||
#[cfg(target_os = "windows")]
|
||||
use std::net::IpAddr;
|
||||
use std::{collections::HashMap, net::IpAddr};
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
use network_interface::{
|
||||
Addr as SystemAddr, NetworkInterface as SystemNetworkInterface, NetworkInterfaceConfig,
|
||||
};
|
||||
use pnet::datalink::NetworkInterface;
|
||||
#[cfg(target_os = "windows")]
|
||||
use pnet::{ipnetwork::IpNetwork, util::MacAddr};
|
||||
use network_interface::{NetworkInterface, NetworkInterfaceConfig};
|
||||
#[cfg(all(target_os = "macos", not(feature = "macos-ne")))]
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
@@ -15,8 +8,113 @@ use crate::proto::peer_rpc::GetIpListResponse;
|
||||
|
||||
use super::netns::NetNS;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
struct InterfaceState {
|
||||
is_point_to_point: bool,
|
||||
is_loopback: bool,
|
||||
is_up: bool,
|
||||
#[cfg(target_os = "linux")]
|
||||
is_lower_up: bool,
|
||||
}
|
||||
|
||||
#[cfg(any(
|
||||
all(target_os = "linux", not(target_env = "ohos")),
|
||||
all(target_os = "macos", not(feature = "macos-ne")),
|
||||
target_os = "freebsd"
|
||||
))]
|
||||
fn collect_interface_states() -> HashMap<String, InterfaceState> {
|
||||
let mut states = HashMap::new();
|
||||
if let Ok(interfaces) = nix::ifaddrs::getifaddrs() {
|
||||
use nix::net::if_::InterfaceFlags;
|
||||
|
||||
for interface in interfaces {
|
||||
let flags = interface.flags;
|
||||
#[cfg(target_os = "linux")]
|
||||
let is_lower_up = flags.contains(InterfaceFlags::IFF_LOWER_UP);
|
||||
states.insert(
|
||||
interface.interface_name,
|
||||
InterfaceState {
|
||||
is_point_to_point: flags.contains(InterfaceFlags::IFF_POINTOPOINT),
|
||||
is_loopback: flags.contains(InterfaceFlags::IFF_LOOPBACK),
|
||||
is_up: flags.contains(InterfaceFlags::IFF_UP),
|
||||
#[cfg(target_os = "linux")]
|
||||
is_lower_up,
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
states
|
||||
}
|
||||
|
||||
#[cfg(not(any(
|
||||
all(target_os = "linux", not(target_env = "ohos")),
|
||||
all(target_os = "macos", not(feature = "macos-ne")),
|
||||
target_os = "freebsd"
|
||||
)))]
|
||||
fn collect_interface_states() -> HashMap<String, InterfaceState> {
|
||||
HashMap::new()
|
||||
}
|
||||
|
||||
#[cfg(any(target_os = "freebsd", target_os = "windows"))]
|
||||
fn has_nonzero_mac(iface: &NetworkInterface) -> bool {
|
||||
iface.mac_addr.as_deref().is_some_and(|mac| {
|
||||
let mut octets = mac.split([':', '-']);
|
||||
let mut nonzero = false;
|
||||
for _ in 0..6 {
|
||||
let Some(value) = octets
|
||||
.next()
|
||||
.and_then(|octet| u8::from_str_radix(octet, 16).ok())
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
nonzero |= value != 0;
|
||||
}
|
||||
octets.next().is_none() && nonzero
|
||||
})
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub(crate) fn ip_mask_to_prefix(mask: IpAddr) -> Result<u8, ()> {
|
||||
match mask {
|
||||
IpAddr::V4(mask) => {
|
||||
let raw = u32::from(mask);
|
||||
let prefix = raw.leading_ones() as u8;
|
||||
let expected = if prefix == 0 {
|
||||
0
|
||||
} else {
|
||||
u32::MAX << (32 - prefix)
|
||||
};
|
||||
(raw == expected).then_some(prefix).ok_or(())
|
||||
}
|
||||
IpAddr::V6(mask) => {
|
||||
let raw = u128::from(mask);
|
||||
let prefix = raw.leading_ones() as u8;
|
||||
let expected = if prefix == 0 {
|
||||
0
|
||||
} else {
|
||||
u128::MAX << (128 - prefix)
|
||||
};
|
||||
(raw == expected).then_some(prefix).ok_or(())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct InterfaceFilter {
|
||||
iface: NetworkInterface,
|
||||
state: InterfaceState,
|
||||
}
|
||||
|
||||
fn interface_state(
|
||||
iface: &NetworkInterface,
|
||||
states: &HashMap<String, InterfaceState>,
|
||||
) -> InterfaceState {
|
||||
states.get(&iface.name).copied().unwrap_or(InterfaceState {
|
||||
is_loopback: iface.internal,
|
||||
is_up: true,
|
||||
#[cfg(target_os = "linux")]
|
||||
is_lower_up: true,
|
||||
..Default::default()
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(any(
|
||||
@@ -40,7 +138,7 @@ impl InterfaceFilter {
|
||||
|
||||
async fn has_valid_ip(&self) -> bool {
|
||||
self.iface
|
||||
.ips
|
||||
.addr
|
||||
.iter()
|
||||
.map(|ip| ip.ip())
|
||||
.any(|ip| !ip.is_loopback() && !ip.is_unspecified() && !ip.is_multicast())
|
||||
@@ -50,18 +148,18 @@ impl InterfaceFilter {
|
||||
tracing::trace!(
|
||||
"filter linux iface: {:?}, is_point_to_point: {}, is_loopback: {}, is_up: {}, is_lower_up: {}, is_tun: {}, has_valid_ip: {}",
|
||||
self.iface,
|
||||
self.iface.is_point_to_point(),
|
||||
self.iface.is_loopback(),
|
||||
self.iface.is_up(),
|
||||
self.iface.is_lower_up(),
|
||||
self.state.is_point_to_point,
|
||||
self.state.is_loopback,
|
||||
self.state.is_up,
|
||||
self.state.is_lower_up,
|
||||
self.is_tun_tap_device().await,
|
||||
self.has_valid_ip().await
|
||||
);
|
||||
|
||||
!self.iface.is_point_to_point()
|
||||
&& !self.iface.is_loopback()
|
||||
&& self.iface.is_up()
|
||||
&& self.iface.is_lower_up()
|
||||
!self.state.is_point_to_point
|
||||
&& !self.state.is_loopback
|
||||
&& self.state.is_up
|
||||
&& self.state.is_lower_up
|
||||
&& !self.is_tun_tap_device().await
|
||||
&& self.has_valid_ip().await
|
||||
}
|
||||
@@ -134,13 +232,13 @@ impl InterfaceFilter {
|
||||
#[cfg(target_os = "freebsd")]
|
||||
async fn is_interface_physical(&self) -> bool {
|
||||
// if mac addr is not zero, then it's physical interface
|
||||
self.iface.mac.map(|mac| !mac.is_zero()).unwrap_or(false)
|
||||
has_nonzero_mac(&self.iface)
|
||||
}
|
||||
|
||||
async fn filter_iface(&self) -> bool {
|
||||
!self.iface.is_point_to_point()
|
||||
&& !self.iface.is_loopback()
|
||||
&& self.iface.is_up()
|
||||
!self.state.is_point_to_point
|
||||
&& !self.state.is_loopback
|
||||
&& self.state.is_up
|
||||
&& self.is_interface_physical().await
|
||||
}
|
||||
}
|
||||
@@ -151,19 +249,19 @@ impl InterfaceFilter {
|
||||
tracing::debug!(
|
||||
"iface_name: {:?}, p2p: {:?}, is_up: {:?}, iface: {:?}",
|
||||
self.iface.name,
|
||||
self.iface.is_point_to_point(),
|
||||
self.iface.is_up(),
|
||||
self.state.is_point_to_point,
|
||||
self.state.is_up,
|
||||
self.iface
|
||||
);
|
||||
!self.iface.is_point_to_point()
|
||||
&& !self.iface.is_loopback()
|
||||
!self.state.is_point_to_point
|
||||
&& !self.state.is_loopback
|
||||
&& self
|
||||
.iface
|
||||
.ips
|
||||
.addr
|
||||
.iter()
|
||||
.map(|ip| ip.ip())
|
||||
.any(|ip| !ip.is_loopback() && !ip.is_unspecified() && !ip.is_multicast())
|
||||
&& self.iface.mac.map(|mac| !mac.is_zero()).unwrap_or(false)
|
||||
&& has_nonzero_mac(&self.iface)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -211,15 +309,66 @@ pub(crate) async fn collect_interfaces(net_ns: NetNS, filter: bool) -> Vec<Netwo
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "faketcp")]
|
||||
fn convert_pnet_interface(iface: pnet_datalink::NetworkInterface) -> NetworkInterface {
|
||||
let internal = iface.is_loopback();
|
||||
let addr = iface
|
||||
.ips
|
||||
.into_iter()
|
||||
.filter_map(|network| match (network.ip(), network.mask()) {
|
||||
(IpAddr::V4(ip), IpAddr::V4(netmask)) => {
|
||||
Some(network_interface::Addr::V4(network_interface::V4IfAddr {
|
||||
ip,
|
||||
broadcast: None,
|
||||
netmask: Some(netmask),
|
||||
}))
|
||||
}
|
||||
(IpAddr::V6(ip), IpAddr::V6(netmask)) => {
|
||||
Some(network_interface::Addr::V6(network_interface::V6IfAddr {
|
||||
ip,
|
||||
broadcast: None,
|
||||
netmask: Some(netmask),
|
||||
}))
|
||||
}
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
NetworkInterface {
|
||||
name: iface.name,
|
||||
addr,
|
||||
mac_addr: iface.mac.map(|mac| mac.to_string()),
|
||||
index: iface.index,
|
||||
internal,
|
||||
}
|
||||
}
|
||||
|
||||
async fn collect_interfaces_in_current_namespace(filter: bool) -> Vec<NetworkInterface> {
|
||||
#[cfg(target_os = "windows")]
|
||||
let ifaces = collect_interfaces_windows();
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
let ifaces = pnet::datalink::interfaces();
|
||||
let ifaces = match NetworkInterface::show() {
|
||||
Ok(ifaces) => ifaces,
|
||||
Err(error) => {
|
||||
tracing::warn!(?error, "failed to enumerate network interfaces");
|
||||
#[cfg(feature = "faketcp")]
|
||||
{
|
||||
match std::panic::catch_unwind(pnet_datalink::interfaces) {
|
||||
Ok(ifaces) => ifaces.into_iter().map(convert_pnet_interface).collect(),
|
||||
Err(_) => {
|
||||
tracing::error!(
|
||||
"failed to enumerate network interfaces via network-interface and pnet"
|
||||
);
|
||||
return Vec::new();
|
||||
}
|
||||
}
|
||||
}
|
||||
#[cfg(not(feature = "faketcp"))]
|
||||
return Vec::new();
|
||||
}
|
||||
};
|
||||
let states = collect_interface_states();
|
||||
let mut ret = vec![];
|
||||
for iface in ifaces {
|
||||
let f = InterfaceFilter {
|
||||
iface: iface.clone(),
|
||||
state: interface_state(&iface, &states),
|
||||
};
|
||||
|
||||
if filter && !f.filter_iface().await {
|
||||
@@ -250,83 +399,6 @@ where
|
||||
.expect("namespace-local network operation panicked")
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
fn collect_interfaces_windows() -> Vec<NetworkInterface> {
|
||||
match SystemNetworkInterface::show() {
|
||||
Ok(ifaces) => ifaces.into_iter().map(convert_windows_interface).collect(),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
?e,
|
||||
"failed to enumerate interfaces via network-interface, falling back to pnet"
|
||||
);
|
||||
match std::panic::catch_unwind(pnet::datalink::interfaces) {
|
||||
Ok(ifaces) => ifaces,
|
||||
Err(_) => {
|
||||
tracing::error!(
|
||||
"failed to enumerate interfaces via both network-interface and pnet"
|
||||
);
|
||||
Vec::new()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
fn convert_windows_interface(iface: SystemNetworkInterface) -> NetworkInterface {
|
||||
let mac = iface.mac_addr.as_deref().and_then(|mac| {
|
||||
mac.parse::<MacAddr>()
|
||||
.map_err(
|
||||
|e| tracing::debug!(iface = %iface.name, mac, ?e, "failed to parse interface mac"),
|
||||
)
|
||||
.ok()
|
||||
});
|
||||
|
||||
let ips = iface
|
||||
.addr
|
||||
.into_iter()
|
||||
.filter_map(convert_windows_interface_addr)
|
||||
.collect();
|
||||
|
||||
NetworkInterface {
|
||||
name: iface.name,
|
||||
description: String::new(),
|
||||
index: iface.index,
|
||||
mac,
|
||||
ips,
|
||||
// pnet does not populate Windows flags either, so keep the existing semantics.
|
||||
flags: 0,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
fn convert_windows_interface_addr(addr: SystemAddr) -> Option<IpNetwork> {
|
||||
match addr {
|
||||
SystemAddr::V4(addr) => {
|
||||
let netmask = addr
|
||||
.netmask
|
||||
.map(IpAddr::V4)
|
||||
.unwrap_or(IpAddr::V4(std::net::Ipv4Addr::new(255, 255, 255, 255)));
|
||||
IpNetwork::with_netmask(IpAddr::V4(addr.ip), netmask)
|
||||
.map_err(
|
||||
|e| tracing::debug!(ip = %addr.ip, ?addr.netmask, ?e, "failed to convert ipv4"),
|
||||
)
|
||||
.ok()
|
||||
}
|
||||
SystemAddr::V6(addr) => {
|
||||
let netmask = addr
|
||||
.netmask
|
||||
.map(IpAddr::V6)
|
||||
.unwrap_or(IpAddr::V6(std::net::Ipv6Addr::from(u128::MAX)));
|
||||
IpNetwork::with_netmask(IpAddr::V6(addr.ip), netmask)
|
||||
.map_err(
|
||||
|e| tracing::debug!(ip = %addr.ip, ?addr.netmask, ?e, "failed to convert ipv6"),
|
||||
)
|
||||
.ok()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(skip(net_ns))]
|
||||
pub(crate) async fn collect_local_ip_addrs(net_ns: NetNS) -> GetIpListResponse {
|
||||
#[cfg(target_os = "linux")]
|
||||
@@ -349,7 +421,7 @@ async fn collect_local_ip_addrs_in_current_namespace() -> GetIpListResponse {
|
||||
|
||||
let ifaces = collect_interfaces_in_current_namespace(true).await;
|
||||
for iface in ifaces {
|
||||
for ip in iface.ips {
|
||||
for ip in iface.addr {
|
||||
let ip: std::net::IpAddr = ip.ip();
|
||||
if let std::net::IpAddr::V4(v4) = ip {
|
||||
if ip.is_loopback() || ip.is_multicast() {
|
||||
@@ -362,7 +434,7 @@ async fn collect_local_ip_addrs_in_current_namespace() -> GetIpListResponse {
|
||||
|
||||
let ifaces = collect_interfaces_in_current_namespace(false).await;
|
||||
for iface in ifaces {
|
||||
for ip in iface.ips {
|
||||
for ip in iface.addr {
|
||||
let ip: std::net::IpAddr = ip.ip();
|
||||
if let std::net::IpAddr::V6(v6) = ip {
|
||||
if v6.is_multicast() || v6.is_loopback() || v6.is_unicast_link_local() {
|
||||
@@ -394,6 +466,22 @@ async fn collect_local_ip_addrs_in_current_namespace() -> GetIpListResponse {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn converts_contiguous_ip_masks_to_prefixes() {
|
||||
assert_eq!(
|
||||
ip_mask_to_prefix(IpAddr::V4("255.255.254.0".parse().unwrap())),
|
||||
Ok(23)
|
||||
);
|
||||
assert_eq!(
|
||||
ip_mask_to_prefix(IpAddr::V6("ffff:ffff:ffff:ffff::".parse().unwrap())),
|
||||
Ok(64)
|
||||
);
|
||||
assert_eq!(
|
||||
ip_mask_to_prefix(IpAddr::V4("255.0.255.0".parse().unwrap())),
|
||||
Err(())
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
#[tokio::test]
|
||||
async fn namespace_operation_does_not_migrate_between_os_threads() {
|
||||
|
||||
@@ -142,7 +142,7 @@ impl ConnectorRuntime for NativeHostRuntime {
|
||||
.into_iter()
|
||||
.find(|interface| {
|
||||
interface
|
||||
.ips
|
||||
.addr
|
||||
.iter()
|
||||
.any(|local| matches!(local.ip(), IpAddr::V6(local_ip) if local_ip == ip))
|
||||
})
|
||||
|
||||
@@ -260,11 +260,7 @@ mod tests {
|
||||
WrappedTransportConnect, WrappedTransportEngine,
|
||||
};
|
||||
use easytier_core::listener::plan::ListenerRuntimeConfig;
|
||||
use pnet::packet::{
|
||||
ip::IpNextHeaderProtocols,
|
||||
ipv4::{self, MutableIpv4Packet},
|
||||
udp::{self, MutableUdpPacket},
|
||||
};
|
||||
use smoltcp::wire::{IpAddress, IpProtocol, Ipv4Packet, UdpPacket};
|
||||
#[cfg(feature = "kcp")]
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
|
||||
@@ -500,30 +496,26 @@ mod tests {
|
||||
let destination_ip = "10.250.0.2".parse().unwrap();
|
||||
let mut ip_packet = vec![0u8; 28];
|
||||
{
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut ip_packet).unwrap();
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut ip_packet);
|
||||
ipv4.set_version(4);
|
||||
ipv4.set_header_length(5);
|
||||
ipv4.set_total_length(28);
|
||||
ipv4.set_ttl(64);
|
||||
ipv4.set_next_level_protocol(IpNextHeaderProtocols::Udp);
|
||||
ipv4.set_source(source_ip);
|
||||
ipv4.set_destination(destination_ip);
|
||||
ipv4.set_header_len(20);
|
||||
ipv4.set_total_len(28);
|
||||
ipv4.set_hop_limit(64);
|
||||
ipv4.set_next_header(IpProtocol::Udp);
|
||||
ipv4.set_src_addr(source_ip);
|
||||
ipv4.set_dst_addr(destination_ip);
|
||||
}
|
||||
{
|
||||
let mut udp = MutableUdpPacket::new(&mut ip_packet[20..]).unwrap();
|
||||
udp.set_source(10000);
|
||||
udp.set_destination(10001);
|
||||
udp.set_length(8);
|
||||
udp.set_checksum(udp::ipv4_checksum(
|
||||
&udp.to_immutable(),
|
||||
&source_ip,
|
||||
&destination_ip,
|
||||
));
|
||||
}
|
||||
{
|
||||
let mut ipv4 = MutableIpv4Packet::new(&mut ip_packet).unwrap();
|
||||
ipv4.set_checksum(ipv4::checksum(&ipv4.to_immutable()));
|
||||
let mut udp = UdpPacket::new_unchecked(&mut ip_packet[20..]);
|
||||
udp.set_src_port(10000);
|
||||
udp.set_dst_port(10001);
|
||||
udp.set_len(8);
|
||||
udp.fill_checksum(
|
||||
&IpAddress::Ipv4(source_ip),
|
||||
&IpAddress::Ipv4(destination_ip),
|
||||
);
|
||||
}
|
||||
Ipv4Packet::new_unchecked(&mut ip_packet).fill_checksum();
|
||||
let received = tokio::time::timeout(std::time::Duration::from_secs(10), async {
|
||||
loop {
|
||||
instance_a
|
||||
|
||||
@@ -237,9 +237,9 @@ fn detect_default_route_ipv6_interfaces(
|
||||
routes: &[DetectedIpv6Route],
|
||||
max_prefix_len: u8,
|
||||
) -> Vec<DetectedDefaultRouteIpv6Interface> {
|
||||
use crate::common::network::ip_mask_to_prefix;
|
||||
use nix::ifaddrs::getifaddrs;
|
||||
use nix::sys::socket::SockaddrLike;
|
||||
use pnet::ipnetwork::ip_mask_to_prefix;
|
||||
|
||||
let wan_ifindices = default_route_ifindices(routes);
|
||||
if wan_ifindices.is_empty() {
|
||||
|
||||
@@ -4,7 +4,6 @@ mod stack;
|
||||
|
||||
use bytes::BytesMut;
|
||||
use network_interface::NetworkInterfaceConfig;
|
||||
use pnet::util::MacAddr;
|
||||
use std::{
|
||||
io,
|
||||
net::{IpAddr, Ipv4Addr, SocketAddr, UdpSocket},
|
||||
@@ -25,6 +24,7 @@ use easytier_core::{
|
||||
use crate::{common::netns::NetNS, tunnel::FromUrl};
|
||||
|
||||
use self::netfilter::create_tun;
|
||||
use self::packet::MacAddr;
|
||||
|
||||
use futures::Future;
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
@@ -35,6 +35,15 @@ struct IpToIfNameCache {
|
||||
ip_to_ifname: DashMap<IpAddr, (String, Option<MacAddr>)>,
|
||||
}
|
||||
|
||||
fn parse_mac_addr(value: &str) -> Option<MacAddr> {
|
||||
let mut bytes = [0; 6];
|
||||
let mut octets = value.split([':', '-']);
|
||||
for byte in &mut bytes {
|
||||
*byte = u8::from_str_radix(octets.next()?, 16).ok()?;
|
||||
}
|
||||
octets.next().is_none().then(|| MacAddr::from_bytes(&bytes))
|
||||
}
|
||||
|
||||
impl IpToIfNameCache {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
@@ -50,9 +59,10 @@ impl IpToIfNameCache {
|
||||
};
|
||||
for iface in interfaces {
|
||||
let mac = iface.mac_addr.as_deref().and_then(|mac| {
|
||||
mac.parse::<MacAddr>().map_err(|e| {
|
||||
tracing::debug!(iface = %iface.name, mac, ?e, "failed to parse interface mac")
|
||||
}).ok()
|
||||
parse_mac_addr(mac).or_else(|| {
|
||||
tracing::debug!(iface = %iface.name, mac, "failed to parse interface mac");
|
||||
None
|
||||
})
|
||||
});
|
||||
for ip in iface.addr.iter() {
|
||||
self.ip_to_ifname.insert(ip.ip(), (iface.name.clone(), mac));
|
||||
|
||||
@@ -630,11 +630,9 @@ impl stack::Tun for LinuxBpfTun {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
use crate::socket::fake_tcp::packet::build_tcp_packet;
|
||||
use crate::socket::fake_tcp::packet::{MacAddr, TCP_FLAG_SYN, build_tcp_packet};
|
||||
use crate::socket::fake_tcp::stack::Tun;
|
||||
use pnet::datalink;
|
||||
use pnet::packet::tcp::TcpFlags;
|
||||
use pnet::util::MacAddr;
|
||||
use pnet_datalink as datalink;
|
||||
use rand::Rng;
|
||||
use std::net::{IpAddr, Ipv4Addr};
|
||||
use tokio::time::{Duration, timeout};
|
||||
@@ -656,7 +654,7 @@ mod tests {
|
||||
IpAddr::V4(ip) => Some(ip),
|
||||
IpAddr::V6(_) => None,
|
||||
})?;
|
||||
return Some((iface.name, ipv4, mac));
|
||||
return Some((iface.name, ipv4, MacAddr::from_bytes(&mac.octets())));
|
||||
}
|
||||
None
|
||||
}
|
||||
@@ -741,7 +739,7 @@ mod tests {
|
||||
dst_addr,
|
||||
1,
|
||||
0,
|
||||
TcpFlags::SYN,
|
||||
TCP_FLAG_SYN,
|
||||
Some(b"ping"),
|
||||
);
|
||||
|
||||
@@ -794,7 +792,7 @@ mod tests {
|
||||
non_matching_dst,
|
||||
1,
|
||||
0,
|
||||
TcpFlags::SYN,
|
||||
TCP_FLAG_SYN,
|
||||
Some(b"nope"),
|
||||
);
|
||||
send_raw_frame(&ifname, &non_matching).unwrap();
|
||||
@@ -819,7 +817,7 @@ mod tests {
|
||||
dst_addr,
|
||||
2,
|
||||
0,
|
||||
TcpFlags::SYN,
|
||||
TCP_FLAG_SYN,
|
||||
Some(b"ok"),
|
||||
);
|
||||
send_raw_frame(&ifname, &matching).unwrap();
|
||||
|
||||
@@ -10,9 +10,9 @@ use std::{
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use dashmap::DashMap;
|
||||
use once_cell::sync::Lazy;
|
||||
use pnet::{
|
||||
datalink::{self, DataLinkSender, NetworkInterface},
|
||||
packet::{ethernet::EtherTypes, ip::IpNextHeaderProtocols, ipv6::Ipv6Packet},
|
||||
use pnet_datalink::{self as datalink, DataLinkSender, NetworkInterface};
|
||||
use smoltcp::wire::{
|
||||
EthernetFrame, EthernetProtocol, IpProtocol, Ipv4Packet, Ipv6Packet, TcpPacket,
|
||||
};
|
||||
#[cfg(target_os = "linux")]
|
||||
use std::os::unix::fs::MetadataExt;
|
||||
@@ -27,49 +27,38 @@ fn filter_tcp_packet(
|
||||
src_addr: Option<&SocketAddr>,
|
||||
dst_addr: Option<&SocketAddr>,
|
||||
) -> bool {
|
||||
use pnet::packet::Packet;
|
||||
use pnet::packet::ethernet::EthernetPacket;
|
||||
use pnet::packet::ipv4::Ipv4Packet;
|
||||
use pnet::packet::tcp::TcpPacket;
|
||||
|
||||
let ethernet = if let Some(ethernet) = EthernetPacket::new(packet) {
|
||||
ethernet
|
||||
} else {
|
||||
let Ok(ethernet) = EthernetFrame::new_checked(packet) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
match ethernet.get_ethertype() {
|
||||
EtherTypes::Ipv4 => {
|
||||
let ipv4 = if let Some(ipv4) = Ipv4Packet::new(ethernet.payload()) {
|
||||
ipv4
|
||||
} else {
|
||||
match ethernet.ethertype() {
|
||||
EthernetProtocol::Ipv4 => {
|
||||
let Ok(ipv4) = Ipv4Packet::new_checked(ethernet.payload()) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if ipv4.get_next_level_protocol() != IpNextHeaderProtocols::Tcp {
|
||||
if ipv4.next_header() != IpProtocol::Tcp {
|
||||
return false;
|
||||
}
|
||||
|
||||
let tcp = if let Some(tcp) = TcpPacket::new(ipv4.payload()) {
|
||||
tcp
|
||||
} else {
|
||||
let Ok(tcp) = TcpPacket::new_checked(ipv4.payload()) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if let Some(src_addr) = src_addr {
|
||||
if IpAddr::V4(ipv4.get_source()) != src_addr.ip() {
|
||||
if IpAddr::V4(ipv4.src_addr()) != src_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_source() != src_addr.port() {
|
||||
if tcp.src_port() != src_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(dst_addr) = dst_addr {
|
||||
if IpAddr::V4(ipv4.get_destination()) != dst_addr.ip() {
|
||||
if IpAddr::V4(ipv4.dst_addr()) != dst_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_destination() != dst_addr.port() {
|
||||
if tcp.dst_port() != dst_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -79,43 +68,39 @@ fn filter_tcp_packet(
|
||||
"FakeTcpSocketListener packet matched filter, dispatching, src_addr: {:?}, dst_addr: {:?}, packet_src_ip: {:?}, packet_dst_ip: {:?}, packet_src_port: {:?}, packet_dst_port: {:?}",
|
||||
src_addr,
|
||||
dst_addr,
|
||||
ipv4.get_source(),
|
||||
ipv4.get_destination(),
|
||||
tcp.get_source(),
|
||||
tcp.get_destination(),
|
||||
ipv4.src_addr(),
|
||||
ipv4.dst_addr(),
|
||||
tcp.src_port(),
|
||||
tcp.dst_port(),
|
||||
);
|
||||
}
|
||||
EtherTypes::Ipv6 => {
|
||||
let ipv6 = if let Some(ipv6) = Ipv6Packet::new(ethernet.payload()) {
|
||||
ipv6
|
||||
} else {
|
||||
EthernetProtocol::Ipv6 => {
|
||||
let Ok(ipv6) = Ipv6Packet::new_checked(ethernet.payload()) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if ipv6.get_next_header() != IpNextHeaderProtocols::Tcp {
|
||||
if ipv6.next_header() != IpProtocol::Tcp {
|
||||
return false;
|
||||
}
|
||||
|
||||
let tcp = if let Some(tcp) = TcpPacket::new(ipv6.payload()) {
|
||||
tcp
|
||||
} else {
|
||||
let Ok(tcp) = TcpPacket::new_checked(ipv6.payload()) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if let Some(src_addr) = src_addr {
|
||||
if IpAddr::V6(ipv6.get_source()) != src_addr.ip() {
|
||||
if IpAddr::V6(ipv6.src_addr()) != src_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_source() != src_addr.port() {
|
||||
if tcp.src_port() != src_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(dst_addr) = dst_addr {
|
||||
if IpAddr::V6(ipv6.get_destination()) != dst_addr.ip() {
|
||||
if IpAddr::V6(ipv6.dst_addr()) != dst_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_destination() != dst_addr.port() {
|
||||
if tcp.dst_port() != dst_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -150,7 +135,7 @@ struct InterfaceWorker {
|
||||
impl InterfaceWorker {
|
||||
fn new(interface: NetworkInterface) -> io::Result<Arc<Self>> {
|
||||
let (tx, mut rx) = match datalink::channel(&interface, Default::default()) {
|
||||
Ok(pnet::datalink::Channel::Ethernet(tx, rx)) => (tx, rx),
|
||||
Ok(datalink::Channel::Ethernet(tx, rx)) => (tx, rx),
|
||||
Ok(_) => return Err(io::Error::other("Unhandled channel type")),
|
||||
Err(e) => return Err(io::Error::other(e)),
|
||||
};
|
||||
|
||||
@@ -1,36 +1,67 @@
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use pnet::packet::ethernet::{EtherTypes, EthernetPacket, MutableEthernetPacket};
|
||||
use pnet::packet::{ip, ipv4, ipv6, tcp};
|
||||
use pnet::util::MacAddr;
|
||||
use std::convert::TryInto;
|
||||
use smoltcp::wire::{
|
||||
ETHERNET_HEADER_LEN, EthernetAddress, EthernetFrame, EthernetProtocol, IpAddress, IpProtocol,
|
||||
Ipv4Packet, Ipv6Packet, TCP_HEADER_LEN, TcpPacket, TcpSeqNumber,
|
||||
};
|
||||
use std::net::{IpAddr, SocketAddr};
|
||||
|
||||
const IPV4_HEADER_LEN: usize = 20;
|
||||
const IPV6_HEADER_LEN: usize = 40;
|
||||
const TCP_HEADER_LEN: usize = 20;
|
||||
use smoltcp::wire::{IPV4_HEADER_LEN, IPV6_HEADER_LEN};
|
||||
|
||||
pub type MacAddr = EthernetAddress;
|
||||
|
||||
pub const TCP_FLAG_FIN: u8 = 0x01;
|
||||
pub const TCP_FLAG_SYN: u8 = 0x02;
|
||||
pub const TCP_FLAG_RST: u8 = 0x04;
|
||||
pub const TCP_FLAG_PSH: u8 = 0x08;
|
||||
pub const TCP_FLAG_ACK: u8 = 0x10;
|
||||
pub const TCP_FLAG_URG: u8 = 0x20;
|
||||
pub const TCP_FLAG_ECE: u8 = 0x40;
|
||||
pub const TCP_FLAG_CWR: u8 = 0x80;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum IPPacket<'p> {
|
||||
V4(ipv4::Ipv4Packet<'p>),
|
||||
V6(ipv6::Ipv6Packet<'p>),
|
||||
V4(Ipv4Packet<&'p [u8]>),
|
||||
V6(Ipv6Packet<&'p [u8]>),
|
||||
}
|
||||
|
||||
impl IPPacket<'_> {
|
||||
pub fn get_source(&self) -> IpAddr {
|
||||
match self {
|
||||
IPPacket::V4(p) => IpAddr::V4(p.get_source()),
|
||||
IPPacket::V6(p) => IpAddr::V6(p.get_source()),
|
||||
IPPacket::V4(p) => IpAddr::V4(p.src_addr()),
|
||||
IPPacket::V6(p) => IpAddr::V6(p.src_addr()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_destination(&self) -> IpAddr {
|
||||
match self {
|
||||
IPPacket::V4(p) => IpAddr::V4(p.get_destination()),
|
||||
IPPacket::V6(p) => IpAddr::V6(p.get_destination()),
|
||||
IPPacket::V4(p) => IpAddr::V4(p.dst_addr()),
|
||||
IPPacket::V6(p) => IpAddr::V6(p.dst_addr()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const ETH_HDR_LEN: usize = 14;
|
||||
fn set_tcp_flags<T: AsRef<[u8]> + AsMut<[u8]>>(tcp: &mut TcpPacket<T>, flags: u8) {
|
||||
tcp.set_fin(flags & TCP_FLAG_FIN != 0);
|
||||
tcp.set_syn(flags & TCP_FLAG_SYN != 0);
|
||||
tcp.set_rst(flags & TCP_FLAG_RST != 0);
|
||||
tcp.set_psh(flags & TCP_FLAG_PSH != 0);
|
||||
tcp.set_ack(flags & TCP_FLAG_ACK != 0);
|
||||
tcp.set_urg(flags & TCP_FLAG_URG != 0);
|
||||
tcp.set_ece(flags & TCP_FLAG_ECE != 0);
|
||||
tcp.set_cwr(flags & TCP_FLAG_CWR != 0);
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub fn tcp_flags<T: AsRef<[u8]>>(tcp: &TcpPacket<T>) -> u8 {
|
||||
u8::from(tcp.fin())
|
||||
| (u8::from(tcp.syn()) << 1)
|
||||
| (u8::from(tcp.rst()) << 2)
|
||||
| (u8::from(tcp.psh()) << 3)
|
||||
| (u8::from(tcp.ack()) << 4)
|
||||
| (u8::from(tcp.urg()) << 5)
|
||||
| (u8::from(tcp.ece()) << 6)
|
||||
| (u8::from(tcp.cwr()) << 7)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn build_tcp_packet(
|
||||
@@ -47,76 +78,72 @@ pub fn build_tcp_packet(
|
||||
SocketAddr::V4(_) => IPV4_HEADER_LEN,
|
||||
SocketAddr::V6(_) => IPV6_HEADER_LEN,
|
||||
};
|
||||
let wscale = (flags & tcp::TcpFlags::SYN) != 0;
|
||||
let wscale = flags & TCP_FLAG_SYN != 0;
|
||||
let tcp_header_len = TCP_HEADER_LEN + if wscale { 4 } else { 0 }; // nop + wscale
|
||||
let tcp_total_len = tcp_header_len + payload.map_or(0, |payload| payload.len());
|
||||
let total_len = ip_header_len + tcp_total_len;
|
||||
let mut buf = BytesMut::zeroed(ETH_HDR_LEN + total_len);
|
||||
let mut buf = BytesMut::zeroed(ETHERNET_HEADER_LEN + total_len);
|
||||
|
||||
let mut eth_buf = buf.split_to(ETH_HDR_LEN);
|
||||
let mut eth_buf = buf.split_to(ETHERNET_HEADER_LEN);
|
||||
let mut ip_buf = buf.split_to(ip_header_len);
|
||||
let mut tcp_buf = buf.split_to(tcp_total_len);
|
||||
assert_eq!(0, buf.len());
|
||||
|
||||
let mut tcp = tcp::MutableTcpPacket::new(&mut tcp_buf).unwrap();
|
||||
tcp.set_window(0xffff);
|
||||
tcp.set_source(local_addr.port());
|
||||
tcp.set_destination(remote_addr.port());
|
||||
tcp.set_sequence(seq);
|
||||
tcp.set_acknowledgement(ack);
|
||||
tcp.set_flags(flags);
|
||||
tcp.set_data_offset(TCP_HEADER_LEN as u8 / 4 + if wscale { 1 } else { 0 });
|
||||
let mut tcp = TcpPacket::new_unchecked(&mut tcp_buf);
|
||||
tcp.set_window_len(0xffff);
|
||||
tcp.set_src_port(local_addr.port());
|
||||
tcp.set_dst_port(remote_addr.port());
|
||||
tcp.set_seq_number(TcpSeqNumber(seq as i32));
|
||||
tcp.set_ack_number(TcpSeqNumber(ack as i32));
|
||||
set_tcp_flags(&mut tcp, flags);
|
||||
tcp.set_header_len(tcp_header_len as u8);
|
||||
if wscale {
|
||||
let wscale = tcp::TcpOption::wscale(14);
|
||||
tcp.set_options(&[tcp::TcpOption::nop(), wscale]);
|
||||
tcp.options_mut().copy_from_slice(&[1, 3, 3, 14]);
|
||||
}
|
||||
|
||||
if let Some(payload) = payload {
|
||||
tcp.set_payload(payload);
|
||||
tcp.payload_mut().copy_from_slice(payload);
|
||||
}
|
||||
|
||||
let mut ethernet = MutableEthernetPacket::new(&mut eth_buf).unwrap();
|
||||
ethernet.set_destination(dst_mac);
|
||||
ethernet.set_source(src_mac);
|
||||
let mut ethernet = EthernetFrame::new_unchecked(&mut eth_buf);
|
||||
ethernet.set_dst_addr(dst_mac);
|
||||
ethernet.set_src_addr(src_mac);
|
||||
ethernet.set_ethertype(match local_addr {
|
||||
SocketAddr::V4(_) => EtherTypes::Ipv4,
|
||||
SocketAddr::V6(_) => EtherTypes::Ipv6,
|
||||
SocketAddr::V4(_) => EthernetProtocol::Ipv4,
|
||||
SocketAddr::V6(_) => EthernetProtocol::Ipv6,
|
||||
});
|
||||
|
||||
match (local_addr, remote_addr) {
|
||||
(SocketAddr::V4(local), SocketAddr::V4(remote)) => {
|
||||
let mut v4 = ipv4::MutableIpv4Packet::new(&mut ip_buf).unwrap();
|
||||
let mut v4 = Ipv4Packet::new_unchecked(&mut ip_buf);
|
||||
v4.set_version(4);
|
||||
v4.set_header_length(IPV4_HEADER_LEN as u8 / 4);
|
||||
v4.set_next_level_protocol(ip::IpNextHeaderProtocols::Tcp);
|
||||
v4.set_ttl(64);
|
||||
v4.set_source(*local.ip());
|
||||
v4.set_destination(*remote.ip());
|
||||
v4.set_total_length(total_len.try_into().unwrap());
|
||||
v4.set_flags(ipv4::Ipv4Flags::DontFragment);
|
||||
v4.set_header_len(IPV4_HEADER_LEN as u8);
|
||||
v4.set_next_header(IpProtocol::Tcp);
|
||||
v4.set_hop_limit(64);
|
||||
v4.set_src_addr(*local.ip());
|
||||
v4.set_dst_addr(*remote.ip());
|
||||
v4.set_total_len(total_len.try_into().unwrap());
|
||||
v4.set_dont_frag(true);
|
||||
|
||||
tcp.set_checksum(tcp::ipv4_checksum(
|
||||
&tcp.to_immutable(),
|
||||
&v4.get_source(),
|
||||
&v4.get_destination(),
|
||||
));
|
||||
|
||||
v4.set_checksum(ipv4::checksum(&v4.to_immutable()));
|
||||
tcp.fill_checksum(
|
||||
&IpAddress::Ipv4(*local.ip()),
|
||||
&IpAddress::Ipv4(*remote.ip()),
|
||||
);
|
||||
v4.fill_checksum();
|
||||
}
|
||||
(SocketAddr::V6(local), SocketAddr::V6(remote)) => {
|
||||
let mut v6 = ipv6::MutableIpv6Packet::new(&mut ip_buf).unwrap();
|
||||
let mut v6 = Ipv6Packet::new_unchecked(&mut ip_buf);
|
||||
v6.set_version(6);
|
||||
v6.set_payload_length(tcp_total_len.try_into().unwrap());
|
||||
v6.set_next_header(ip::IpNextHeaderProtocols::Tcp);
|
||||
v6.set_payload_len(tcp_total_len.try_into().unwrap());
|
||||
v6.set_next_header(IpProtocol::Tcp);
|
||||
v6.set_hop_limit(64);
|
||||
v6.set_source(*local.ip());
|
||||
v6.set_destination(*remote.ip());
|
||||
v6.set_src_addr(*local.ip());
|
||||
v6.set_dst_addr(*remote.ip());
|
||||
|
||||
tcp.set_checksum(tcp::ipv6_checksum(
|
||||
&tcp.to_immutable(),
|
||||
&v6.get_source(),
|
||||
&v6.get_destination(),
|
||||
));
|
||||
tcp.fill_checksum(
|
||||
&IpAddress::Ipv6(*local.ip()),
|
||||
&IpAddress::Ipv6(*remote.ip()),
|
||||
);
|
||||
}
|
||||
_ => unreachable!(),
|
||||
};
|
||||
@@ -126,40 +153,34 @@ pub fn build_tcp_packet(
|
||||
eth_buf.freeze()
|
||||
}
|
||||
|
||||
pub fn parse_ip_packet(
|
||||
buf: &Bytes,
|
||||
) -> Option<(MacAddr, MacAddr, IPPacket<'_>, tcp::TcpPacket<'_>)> {
|
||||
let eth = EthernetPacket::new(buf.as_ref())?;
|
||||
let src_mac = eth.get_source();
|
||||
let dst_mac = eth.get_destination();
|
||||
let ethertype = eth.get_ethertype();
|
||||
pub fn parse_ip_packet(buf: &Bytes) -> Option<(MacAddr, MacAddr, IPPacket<'_>, TcpPacket<&[u8]>)> {
|
||||
let eth = EthernetFrame::new_checked(buf.as_ref()).ok()?;
|
||||
let src_mac = eth.src_addr();
|
||||
let dst_mac = eth.dst_addr();
|
||||
let ethertype = eth.ethertype();
|
||||
|
||||
tracing::trace!("Parsing IP packet: {:?}", eth);
|
||||
|
||||
let ip_payload = &buf[ETH_HDR_LEN..];
|
||||
let ip_payload = eth.payload();
|
||||
|
||||
match ethertype {
|
||||
EtherTypes::Ipv4 => {
|
||||
let v4 = ipv4::Ipv4Packet::new(ip_payload)?;
|
||||
if v4.get_next_level_protocol() != ip::IpNextHeaderProtocols::Tcp {
|
||||
EthernetProtocol::Ipv4 => {
|
||||
let v4 = Ipv4Packet::new_checked(ip_payload).ok()?;
|
||||
if usize::from(v4.header_len()) < IPV4_HEADER_LEN {
|
||||
return None;
|
||||
}
|
||||
|
||||
let tcp_offset = usize::from(v4.get_header_length()) * 4;
|
||||
if tcp_offset < IPV4_HEADER_LEN || tcp_offset > ip_payload.len() {
|
||||
if v4.next_header() != IpProtocol::Tcp {
|
||||
return None;
|
||||
}
|
||||
|
||||
let tcp = tcp::TcpPacket::new(&ip_payload[tcp_offset..])?;
|
||||
let tcp = TcpPacket::new_checked(v4.payload()).ok()?;
|
||||
Some((src_mac, dst_mac, IPPacket::V4(v4), tcp))
|
||||
}
|
||||
EtherTypes::Ipv6 => {
|
||||
let v6 = ipv6::Ipv6Packet::new(ip_payload)?;
|
||||
if v6.get_next_header() != ip::IpNextHeaderProtocols::Tcp {
|
||||
EthernetProtocol::Ipv6 => {
|
||||
let v6 = Ipv6Packet::new_checked(ip_payload).ok()?;
|
||||
if v6.next_header() != IpProtocol::Tcp {
|
||||
return None;
|
||||
}
|
||||
|
||||
let tcp = tcp::TcpPacket::new(&ip_payload[IPV6_HEADER_LEN..])?;
|
||||
let tcp = TcpPacket::new_checked(v6.payload()).ok()?;
|
||||
Some((src_mac, dst_mac, IPPacket::V6(v6), tcp))
|
||||
}
|
||||
_ => None,
|
||||
@@ -169,12 +190,11 @@ pub fn parse_ip_packet(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use pnet::packet::Packet as _;
|
||||
|
||||
#[test]
|
||||
fn parse_ipv4_packet_round_trip() {
|
||||
let src_mac = MacAddr::new(0x02, 0, 0, 0, 0, 1);
|
||||
let dst_mac = MacAddr::new(0x02, 0, 0, 0, 0, 2);
|
||||
let src_mac = MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 1]);
|
||||
let dst_mac = MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 2]);
|
||||
let local_addr: SocketAddr = "192.0.2.1:12345".parse().unwrap();
|
||||
let remote_addr: SocketAddr = "198.51.100.2:23456".parse().unwrap();
|
||||
let payload = b"hello fake tcp";
|
||||
@@ -186,7 +206,7 @@ mod tests {
|
||||
remote_addr,
|
||||
10,
|
||||
20,
|
||||
tcp::TcpFlags::ACK,
|
||||
TCP_FLAG_ACK,
|
||||
Some(payload),
|
||||
);
|
||||
|
||||
@@ -197,15 +217,15 @@ mod tests {
|
||||
assert_eq!(parsed_dst_mac, dst_mac);
|
||||
assert_eq!(ip_packet.get_source(), local_addr.ip());
|
||||
assert_eq!(ip_packet.get_destination(), remote_addr.ip());
|
||||
assert_eq!(tcp_packet.get_source(), local_addr.port());
|
||||
assert_eq!(tcp_packet.get_destination(), remote_addr.port());
|
||||
assert_eq!(tcp_packet.src_port(), local_addr.port());
|
||||
assert_eq!(tcp_packet.dst_port(), remote_addr.port());
|
||||
assert_eq!(tcp_packet.payload(), payload);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_and_parse_ipv6_packet_round_trip() {
|
||||
let src_mac = MacAddr::new(0x02, 0, 0, 0, 0, 3);
|
||||
let dst_mac = MacAddr::new(0x02, 0, 0, 0, 0, 4);
|
||||
let src_mac = MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 3]);
|
||||
let dst_mac = MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 4]);
|
||||
let local_addr: SocketAddr = "[2001:db8::1]:12345".parse().unwrap();
|
||||
let remote_addr: SocketAddr = "[2001:db8::2]:23456".parse().unwrap();
|
||||
let payload = b"ipv6 payload";
|
||||
@@ -217,12 +237,12 @@ mod tests {
|
||||
remote_addr,
|
||||
30,
|
||||
40,
|
||||
tcp::TcpFlags::ACK,
|
||||
TCP_FLAG_ACK,
|
||||
Some(payload),
|
||||
);
|
||||
|
||||
let ethernet = EthernetPacket::new(packet.as_ref()).unwrap();
|
||||
assert_eq!(ethernet.get_ethertype(), EtherTypes::Ipv6);
|
||||
let ethernet = EthernetFrame::new_checked(packet.as_ref()).unwrap();
|
||||
assert_eq!(ethernet.ethertype(), EthernetProtocol::Ipv6);
|
||||
|
||||
let (parsed_src_mac, parsed_dst_mac, ip_packet, tcp_packet) =
|
||||
parse_ip_packet(&packet).unwrap();
|
||||
@@ -231,48 +251,159 @@ mod tests {
|
||||
assert_eq!(parsed_dst_mac, dst_mac);
|
||||
assert_eq!(ip_packet.get_source(), local_addr.ip());
|
||||
assert_eq!(ip_packet.get_destination(), remote_addr.ip());
|
||||
assert_eq!(tcp_packet.get_source(), local_addr.port());
|
||||
assert_eq!(tcp_packet.get_destination(), remote_addr.port());
|
||||
assert_eq!(tcp_packet.src_port(), local_addr.port());
|
||||
assert_eq!(tcp_packet.dst_port(), remote_addr.port());
|
||||
assert_eq!(tcp_packet.payload(), payload);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rejects_short_ethernet_frame() {
|
||||
let packet = Bytes::from_static(&[0u8; ETH_HDR_LEN - 1]);
|
||||
let packet = Bytes::from_static(&[0u8; ETHERNET_HEADER_LEN - 1]);
|
||||
assert!(parse_ip_packet(&packet).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rejects_truncated_ipv4_tcp_packet() {
|
||||
let packet = build_tcp_packet(
|
||||
MacAddr::new(0x02, 0, 0, 0, 0, 5),
|
||||
MacAddr::new(0x02, 0, 0, 0, 0, 6),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 5]),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 6]),
|
||||
"192.0.2.10:1111".parse().unwrap(),
|
||||
"198.51.100.20:2222".parse().unwrap(),
|
||||
1,
|
||||
2,
|
||||
tcp::TcpFlags::ACK,
|
||||
TCP_FLAG_ACK,
|
||||
None,
|
||||
);
|
||||
let truncated = Bytes::copy_from_slice(&packet[..ETH_HDR_LEN + IPV4_HEADER_LEN + 10]);
|
||||
let truncated =
|
||||
Bytes::copy_from_slice(&packet[..ETHERNET_HEADER_LEN + IPV4_HEADER_LEN + 10]);
|
||||
|
||||
assert!(parse_ip_packet(&truncated).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rejects_ipv4_header_shorter_than_minimum() {
|
||||
let packet = build_tcp_packet(
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 5]),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 6]),
|
||||
"192.0.2.10:1111".parse().unwrap(),
|
||||
"198.51.100.20:2222".parse().unwrap(),
|
||||
1,
|
||||
0x5000_0000,
|
||||
TCP_FLAG_ACK,
|
||||
None,
|
||||
);
|
||||
let mut malformed = BytesMut::from(packet.as_ref());
|
||||
Ipv4Packet::new_unchecked(&mut malformed[ETHERNET_HEADER_LEN..])
|
||||
.set_header_len((IPV4_HEADER_LEN - 4) as u8);
|
||||
let malformed = malformed.freeze();
|
||||
|
||||
assert!(parse_ip_packet(&malformed).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rejects_truncated_ipv6_header() {
|
||||
let packet = build_tcp_packet(
|
||||
MacAddr::new(0x02, 0, 0, 0, 0, 7),
|
||||
MacAddr::new(0x02, 0, 0, 0, 0, 8),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 7]),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 8]),
|
||||
"[2001:db8::10]:1111".parse().unwrap(),
|
||||
"[2001:db8::20]:2222".parse().unwrap(),
|
||||
1,
|
||||
2,
|
||||
tcp::TcpFlags::ACK,
|
||||
TCP_FLAG_ACK,
|
||||
None,
|
||||
);
|
||||
let truncated = Bytes::copy_from_slice(&packet[..ETH_HDR_LEN + IPV6_HEADER_LEN - 1]);
|
||||
let truncated =
|
||||
Bytes::copy_from_slice(&packet[..ETHERNET_HEADER_LEN + IPV6_HEADER_LEN - 1]);
|
||||
|
||||
assert!(parse_ip_packet(&truncated).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn syn_packet_preserves_wire_format_and_unsigned_sequence() {
|
||||
let packet = build_tcp_packet(
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 1]),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 2]),
|
||||
"192.0.2.1:12345".parse().unwrap(),
|
||||
"198.51.100.2:23456".parse().unwrap(),
|
||||
0x8000_0001,
|
||||
0xffff_fffe,
|
||||
TCP_FLAG_SYN | TCP_FLAG_ACK,
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
&packet[..ETHERNET_HEADER_LEN],
|
||||
&[0x02, 0, 0, 0, 0, 2, 0x02, 0, 0, 0, 0, 1, 0x08, 0x00]
|
||||
);
|
||||
let (_, _, IPPacket::V4(ipv4), tcp) = parse_ip_packet(&packet).unwrap() else {
|
||||
panic!("expected IPv4 packet");
|
||||
};
|
||||
assert_eq!(ipv4.header_len(), IPV4_HEADER_LEN as u8);
|
||||
assert!(ipv4.dont_frag());
|
||||
assert_eq!(ipv4.hop_limit(), 64);
|
||||
assert!(ipv4.verify_checksum());
|
||||
assert_eq!(tcp.header_len(), (TCP_HEADER_LEN + 4) as u8);
|
||||
assert_eq!(tcp.options(), &[1, 3, 3, 14]);
|
||||
assert_eq!(tcp_flags(&tcp), TCP_FLAG_SYN | TCP_FLAG_ACK);
|
||||
assert_eq!(tcp.seq_number().0 as u32, 0x8000_0001);
|
||||
assert_eq!(tcp.ack_number().0 as u32, 0xffff_fffe);
|
||||
assert!(tcp.verify_checksum(
|
||||
&IpAddress::Ipv4(ipv4.src_addr()),
|
||||
&IpAddress::Ipv4(ipv4.dst_addr()),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ethernet_padding_is_not_tcp_payload() {
|
||||
let packet = build_tcp_packet(
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 1]),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 2]),
|
||||
"192.0.2.1:12345".parse().unwrap(),
|
||||
"198.51.100.2:23456".parse().unwrap(),
|
||||
1,
|
||||
2,
|
||||
TCP_FLAG_ACK,
|
||||
Some(b"payload"),
|
||||
);
|
||||
let mut padded = BytesMut::from(packet.as_ref());
|
||||
padded.extend_from_slice(&[0; 16]);
|
||||
let padded = padded.freeze();
|
||||
|
||||
let (_, _, _, tcp) = parse_ip_packet(&padded).unwrap();
|
||||
assert_eq!(tcp.payload(), b"payload");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ipv4_options_do_not_shift_tcp_payload() {
|
||||
let packet = build_tcp_packet(
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 1]),
|
||||
MacAddr::from_bytes(&[0x02, 0, 0, 0, 0, 2]),
|
||||
"192.0.2.1:12345".parse().unwrap(),
|
||||
"198.51.100.2:23456".parse().unwrap(),
|
||||
1,
|
||||
2,
|
||||
TCP_FLAG_ACK,
|
||||
Some(b"payload"),
|
||||
);
|
||||
let ip_start = ETHERNET_HEADER_LEN;
|
||||
let tcp_start = ip_start + IPV4_HEADER_LEN;
|
||||
let mut with_options = Vec::with_capacity(packet.len() + 4);
|
||||
with_options.extend_from_slice(&packet[..tcp_start]);
|
||||
with_options.extend_from_slice(&[1, 1, 1, 0]);
|
||||
with_options.extend_from_slice(&packet[tcp_start..]);
|
||||
{
|
||||
let total_len = with_options.len() - ip_start;
|
||||
let mut ipv4 = Ipv4Packet::new_unchecked(&mut with_options[ip_start..]);
|
||||
ipv4.set_header_len((IPV4_HEADER_LEN + 4) as u8);
|
||||
ipv4.set_total_len(total_len as u16);
|
||||
ipv4.fill_checksum();
|
||||
}
|
||||
let with_options = Bytes::from(with_options);
|
||||
|
||||
let (_, _, IPPacket::V4(ipv4), tcp) = parse_ip_packet(&with_options).unwrap() else {
|
||||
panic!("expected IPv4 packet");
|
||||
};
|
||||
assert_eq!(ipv4.header_len(), (IPV4_HEADER_LEN + 4) as u8);
|
||||
assert_eq!(tcp.payload(), b"payload");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -41,9 +41,6 @@
|
||||
use super::packet::*;
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use crossbeam::atomic::AtomicCell;
|
||||
use pnet::packet::tcp::TcpOptionNumbers;
|
||||
use pnet::packet::{Packet, tcp};
|
||||
use pnet::util::MacAddr;
|
||||
use std::collections::HashMap;
|
||||
use std::fmt;
|
||||
#[cfg(test)]
|
||||
@@ -60,6 +57,47 @@ use tracing::{error, info, trace, warn};
|
||||
|
||||
const TIMEOUT: time::Duration = time::Duration::from_secs(1);
|
||||
const MPMC_BUFFER_LEN: usize = 512;
|
||||
const TCP_OPTION_END: u8 = 0;
|
||||
const TCP_OPTION_NOP: u8 = 1;
|
||||
const TCP_OPTION_SACK: u8 = 5;
|
||||
|
||||
struct TcpOptionIter<'a> {
|
||||
remaining: &'a [u8],
|
||||
}
|
||||
|
||||
impl<'a> TcpOptionIter<'a> {
|
||||
fn new(options: &'a [u8]) -> Self {
|
||||
Self { remaining: options }
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> Iterator for TcpOptionIter<'a> {
|
||||
type Item = (u8, &'a [u8]);
|
||||
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
let kind = *self.remaining.first()?;
|
||||
if kind == TCP_OPTION_END {
|
||||
self.remaining = &[];
|
||||
return None;
|
||||
}
|
||||
if kind == TCP_OPTION_NOP {
|
||||
self.remaining = &self.remaining[1..];
|
||||
return Some((kind, &[]));
|
||||
}
|
||||
let Some(&length) = self.remaining.get(1) else {
|
||||
self.remaining = &[];
|
||||
return None;
|
||||
};
|
||||
let length = usize::from(length);
|
||||
if length < 2 || length > self.remaining.len() {
|
||||
self.remaining = &[];
|
||||
return None;
|
||||
}
|
||||
let payload = &self.remaining[2..length];
|
||||
self.remaining = &self.remaining[length..];
|
||||
Some((kind, payload))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
pub trait Tun: Send + Sync + 'static {
|
||||
@@ -181,7 +219,7 @@ impl Socket {
|
||||
|
||||
build_tcp_packet(
|
||||
self.local_mac,
|
||||
self.remote_mac.load().unwrap_or(MacAddr::zero()),
|
||||
self.remote_mac.load().unwrap_or_default(),
|
||||
self.local_addr,
|
||||
self.remote_addr,
|
||||
self.seq.load(Ordering::Relaxed),
|
||||
@@ -201,7 +239,7 @@ impl Socket {
|
||||
pub fn try_send(&self, payload: &[u8]) -> Option<()> {
|
||||
match self.state.load() {
|
||||
State::Established => {
|
||||
let buf = self.build_tcp_packet(tcp::TcpFlags::ACK, Some(payload));
|
||||
let buf = self.build_tcp_packet(TCP_FLAG_ACK, Some(payload));
|
||||
self.seq.fetch_add(payload.len() as u32, Ordering::Relaxed);
|
||||
self.tun.try_send(&buf).ok().and(Some(()))
|
||||
}
|
||||
@@ -211,7 +249,7 @@ impl Socket {
|
||||
|
||||
pub fn close(&self) {
|
||||
if self.state.load() != State::Idle {
|
||||
let buf = self.build_tcp_packet(tcp::TcpFlags::RST, None);
|
||||
let buf = self.build_tcp_packet(TCP_FLAG_RST, None);
|
||||
let _ = self.tun.try_send(&buf);
|
||||
self.state.store(State::Idle);
|
||||
}
|
||||
@@ -256,32 +294,30 @@ impl Socket {
|
||||
|
||||
self.remote_mac.store(Some(src_mac));
|
||||
|
||||
if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 {
|
||||
if tcp_packet.rst() {
|
||||
info!("Connection {} reset by peer", self);
|
||||
return None;
|
||||
}
|
||||
|
||||
if (tcp_packet.get_flags() & tcp::TcpFlags::ACK) != 0
|
||||
&& tcp_packet.payload().is_empty()
|
||||
{
|
||||
if tcp_packet.ack() && tcp_packet.payload().is_empty() {
|
||||
self.seq
|
||||
.store(tcp_packet.get_acknowledgement(), Ordering::Relaxed);
|
||||
.store(tcp_packet.ack_number().0 as u32, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
let payload = tcp_packet.payload();
|
||||
|
||||
let new_ack = tcp_packet.get_sequence().wrapping_add(payload.len() as u32);
|
||||
let new_ack =
|
||||
(tcp_packet.seq_number().0 as u32).wrapping_add(payload.len() as u32);
|
||||
self.ack.store(new_ack, Ordering::Relaxed);
|
||||
|
||||
for opt in tcp_packet.get_options_iter() {
|
||||
if opt.get_number() == TcpOptionNumbers::SACK {
|
||||
for (kind, option_payload) in TcpOptionIter::new(tcp_packet.options()) {
|
||||
if kind == TCP_OPTION_SACK {
|
||||
// SACK 选项类型为 5
|
||||
let payload = opt.payload();
|
||||
for chunk in payload.chunks(8) {
|
||||
for chunk in option_payload.chunks(8) {
|
||||
if chunk.len() != 8 {
|
||||
continue;
|
||||
}
|
||||
let left = tcp_packet.get_acknowledgement();
|
||||
let left = tcp_packet.ack_number().0 as u32;
|
||||
let right = u32::from_be_bytes(chunk[0..4].try_into().unwrap());
|
||||
let len = right.wrapping_sub(left);
|
||||
|
||||
@@ -295,12 +331,12 @@ impl Socket {
|
||||
|
||||
let buf = build_tcp_packet(
|
||||
self.local_mac,
|
||||
self.remote_mac.load().unwrap_or(MacAddr::zero()),
|
||||
self.remote_mac.load().unwrap_or_default(),
|
||||
self.local_addr,
|
||||
self.remote_addr,
|
||||
left,
|
||||
self.ack.load(Ordering::Relaxed),
|
||||
tcp::TcpFlags::ACK,
|
||||
TCP_FLAG_ACK,
|
||||
Some(&data),
|
||||
);
|
||||
|
||||
@@ -332,18 +368,19 @@ impl Socket {
|
||||
continue;
|
||||
};
|
||||
|
||||
if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 {
|
||||
if tcp_packet.rst() {
|
||||
tracing::trace!("Connection {} reset by peer", self);
|
||||
return None;
|
||||
}
|
||||
|
||||
let expected_flag = tcp::TcpFlags::SYN | tcp::TcpFlags::ACK;
|
||||
if (tcp_packet.get_flags() & expected_flag) == expected_flag {
|
||||
if tcp_packet.syn() && tcp_packet.ack() {
|
||||
// found our SYN + ACK
|
||||
self.seq
|
||||
.store(tcp_packet.get_acknowledgement(), Ordering::Relaxed);
|
||||
self.ack
|
||||
.store(tcp_packet.get_sequence() + 1, Ordering::Relaxed);
|
||||
.store(tcp_packet.ack_number().0 as u32, Ordering::Relaxed);
|
||||
self.ack.store(
|
||||
(tcp_packet.seq_number().0 as u32).wrapping_add(1),
|
||||
Ordering::Relaxed,
|
||||
);
|
||||
self.remote_mac.store(Some(src_mac));
|
||||
self.state.store(State::Established);
|
||||
return Some(0);
|
||||
@@ -385,12 +422,12 @@ impl Drop for Socket {
|
||||
|
||||
let buf = build_tcp_packet(
|
||||
self.local_mac,
|
||||
self.remote_mac.load().unwrap_or(MacAddr::zero()),
|
||||
self.remote_mac.load().unwrap_or_default(),
|
||||
self.local_addr,
|
||||
self.remote_addr,
|
||||
self.seq.load(Ordering::Relaxed),
|
||||
0,
|
||||
tcp::TcpFlags::RST,
|
||||
TCP_FLAG_RST,
|
||||
None,
|
||||
);
|
||||
if let Err(e) = self.tun.try_send(&buf) {
|
||||
@@ -434,7 +471,7 @@ impl Stack {
|
||||
|
||||
Stack {
|
||||
shared,
|
||||
local_mac: local_mac.unwrap_or(MacAddr::zero()),
|
||||
local_mac: local_mac.unwrap_or_default(),
|
||||
reader_task: AbortOnDropHandle::new(t),
|
||||
}
|
||||
}
|
||||
@@ -514,11 +551,11 @@ impl Stack {
|
||||
Some((_src_mac, _dst_mac, ip_packet, tcp_packet)) => {
|
||||
let local_addr = SocketAddr::new(
|
||||
ip_packet.get_destination(),
|
||||
tcp_packet.get_destination(),
|
||||
tcp_packet.dst_port(),
|
||||
);
|
||||
let remote_addr = SocketAddr::new(
|
||||
ip_packet.get_source(),
|
||||
tcp_packet.get_source(),
|
||||
tcp_packet.src_port(),
|
||||
);
|
||||
|
||||
let tuple = AddrTuple::new(local_addr, remote_addr);
|
||||
@@ -548,7 +585,7 @@ impl Stack {
|
||||
}
|
||||
}
|
||||
|
||||
if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 {
|
||||
if tcp_packet.rst() {
|
||||
info!("Unknown RST TCP packet from {}, ignoring", remote_addr);
|
||||
continue;
|
||||
} else {
|
||||
@@ -604,6 +641,43 @@ mod tests {
|
||||
time::{Duration, timeout},
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn tcp_option_iterator_preserves_all_sack_blocks() {
|
||||
for block_count in 1..=4_u32 {
|
||||
let mut options = vec![TCP_OPTION_NOP, TCP_OPTION_SACK, (2 + block_count * 8) as u8];
|
||||
for block in 0..block_count {
|
||||
options.extend_from_slice(&(100 + block * 10).to_be_bytes());
|
||||
options.extend_from_slice(&(110 + block * 10).to_be_bytes());
|
||||
}
|
||||
options.push(TCP_OPTION_END);
|
||||
|
||||
let parsed = TcpOptionIter::new(&options).collect::<Vec<_>>();
|
||||
assert_eq!(parsed[0], (TCP_OPTION_NOP, &[][..]));
|
||||
assert_eq!(parsed[1].0, TCP_OPTION_SACK);
|
||||
assert_eq!(parsed[1].1.len(), block_count as usize * 8);
|
||||
let last = parsed[1].1.len() - 8;
|
||||
assert_eq!(
|
||||
u32::from_be_bytes(parsed[1].1[last..last + 4].try_into().unwrap()),
|
||||
100 + (block_count - 1) * 10
|
||||
);
|
||||
assert_eq!(
|
||||
u32::from_be_bytes(parsed[1].1[last + 4..last + 8].try_into().unwrap()),
|
||||
110 + (block_count - 1) * 10
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tcp_option_iterator_stops_at_end_or_malformed_length() {
|
||||
assert_eq!(
|
||||
TcpOptionIter::new(&[TCP_OPTION_END, TCP_OPTION_SACK, 2]).count(),
|
||||
0
|
||||
);
|
||||
assert_eq!(TcpOptionIter::new(&[TCP_OPTION_SACK]).count(), 0);
|
||||
assert_eq!(TcpOptionIter::new(&[TCP_OPTION_SACK, 1]).count(), 0);
|
||||
assert_eq!(TcpOptionIter::new(&[TCP_OPTION_SACK, 10, 0, 0]).count(), 0);
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct FailingTun {
|
||||
fail: Notify,
|
||||
|
||||
Reference in New Issue
Block a user