mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-09-04 10:05:42 +00:00
refactor(core): separate portable core from native runtime (#2451)
Create easytier-core as the portable owner of configuration, connectivity, tunnels, peer and routing state, gateways, management, the data plane, and instance lifecycle. Keep operating-system integration, native protocol engines, process startup, and presentation in easytier behind explicit Host capability adapters. Create easytier-proto to own schemas, generated RPC types, descriptors, and feature-scoped protocol slices. Remove runtime protobuf reflection from core while preserving unknown route-peer fields across forwarding. Normalize instance construction through CoreInstance, CoreHostAdapters, CoreProcessRuntime, and InstanceManager. Make the runtime config store the only authoritative mutable configuration after startup. Move the portable TCP/UDP data plane into core and extract a generic OperationBroker for completion, cancellation, disposal, and capacity accounting. Expose the session-based FFI v2 completion API and keep the WASI guest ABI, wire schemas, and adapters with core. Migrate CLI, GUI, web, FFI, Android JNI, OHOS, uptime, and mobile consumers to the shared manager and core state. Add explicit user/web config ownership and revision-aware web reconciliation. Preserve configuration, wire, and management behavior while fixing regressions discovered by the full platform and integration matrix: - inherit advertised relay capabilities in foreign networks; - refresh OSPF peer state immediately after runtime config changes; - restore CLI GlobalCtx event output without forcing GUI logging; - retain legacy encryption names and standalone RPC tunnel metadata; - restore ICMP host composition and fragmented UDP handling; - use portable 64-bit atomics on 32-bit MIPS targets; and - retain discarded operations until late cancellation completes. Validate the refactor across 45 GitHub checks, including Linux, macOS, Windows, FreeBSD, web, GUI, Android, OHOS, feature profiles, and three-node and subnet-proxy integration tests. BREAKING CHANGE: internal Rust module paths are not preserved. Legacy native data-plane APIs are replaced by the session-based FFI v2 API. The dedicated Android data-plane wrapper is removed.
This commit is contained in:
@@ -1,92 +0,0 @@
|
||||
use std::collections::VecDeque;
|
||||
use std::io::IoSlice;
|
||||
|
||||
use bytes::{Buf, BufMut, Bytes, BytesMut};
|
||||
|
||||
pub(crate) struct BufList<T> {
|
||||
bufs: VecDeque<T>,
|
||||
}
|
||||
|
||||
impl<T: Buf> BufList<T> {
|
||||
pub(crate) fn new() -> BufList<T> {
|
||||
BufList {
|
||||
bufs: VecDeque::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub(crate) fn push(&mut self, buf: T) {
|
||||
debug_assert!(buf.has_remaining());
|
||||
self.bufs.push_back(buf);
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub(crate) fn bufs_cnt(&self) -> usize {
|
||||
self.bufs.len()
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Buf> Buf for BufList<T> {
|
||||
#[inline]
|
||||
fn remaining(&self) -> usize {
|
||||
self.bufs.iter().map(|buf| buf.remaining()).sum()
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn chunk(&self) -> &[u8] {
|
||||
self.bufs.front().map(Buf::chunk).unwrap_or_default()
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn advance(&mut self, mut cnt: usize) {
|
||||
while cnt > 0 {
|
||||
{
|
||||
let front = &mut self.bufs[0];
|
||||
let rem = front.remaining();
|
||||
if rem > cnt {
|
||||
front.advance(cnt);
|
||||
return;
|
||||
} else {
|
||||
front.advance(rem);
|
||||
cnt -= rem;
|
||||
}
|
||||
}
|
||||
self.bufs.pop_front();
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn chunks_vectored<'t>(&'t self, dst: &mut [IoSlice<'t>]) -> usize {
|
||||
if dst.is_empty() {
|
||||
return 0;
|
||||
}
|
||||
let mut vecs = 0;
|
||||
for buf in &self.bufs {
|
||||
vecs += buf.chunks_vectored(&mut dst[vecs..]);
|
||||
if vecs == dst.len() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
vecs
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn copy_to_bytes(&mut self, len: usize) -> Bytes {
|
||||
// Our inner buffer may have an optimized version of copy_to_bytes, and if the whole
|
||||
// request can be fulfilled by the front buffer, we can take advantage.
|
||||
match self.bufs.front_mut() {
|
||||
Some(front) if front.remaining() == len => {
|
||||
let b = front.copy_to_bytes(len);
|
||||
self.bufs.pop_front();
|
||||
b
|
||||
}
|
||||
Some(front) if front.remaining() > len => front.copy_to_bytes(len),
|
||||
_ => {
|
||||
assert!(len <= self.remaining(), "`len` greater than remaining");
|
||||
let mut bm = BytesMut::with_capacity(len);
|
||||
bm.put(self.take(len));
|
||||
bm.freeze()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+69
-427
@@ -1,336 +1,10 @@
|
||||
use bon::builder;
|
||||
use futures::{Future, Sink, Stream, stream::FuturesUnordered};
|
||||
use network_interface::NetworkInterfaceConfig as _;
|
||||
use pin_project_lite::pin_project;
|
||||
use std::{
|
||||
any::Any,
|
||||
net::{IpAddr, SocketAddr},
|
||||
pin::Pin,
|
||||
sync::{Arc, Mutex},
|
||||
task::{Poll, ready},
|
||||
};
|
||||
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
|
||||
use std::net::{IpAddr, SocketAddr};
|
||||
|
||||
use super::TunnelInfo;
|
||||
use super::{
|
||||
SinkItem, StreamItem, Tunnel, TunnelError, ZCPacketSink, ZCPacketStream,
|
||||
buf::BufList,
|
||||
packet_def::{TCP_TUNNEL_HEADER_SIZE, TCPTunnelHeader, ZCPacketType},
|
||||
};
|
||||
use crate::common::netns::NetNS;
|
||||
use crate::tunnel::packet_def::{PEER_MANAGER_HEADER_SIZE, ZCPacket};
|
||||
use bytes::{Buf, BufMut, Bytes, BytesMut};
|
||||
use easytier_core::tunnel::TunnelError;
|
||||
use tokio::net::{TcpListener, TcpSocket, UdpSocket};
|
||||
use tokio_stream::StreamExt;
|
||||
use tokio_util::io::poll_write_buf;
|
||||
use zerocopy::FromBytes as _;
|
||||
|
||||
pub struct TunnelWrapper<R, W> {
|
||||
reader: Arc<Mutex<Option<R>>>,
|
||||
writer: Arc<Mutex<Option<W>>>,
|
||||
info: Option<TunnelInfo>,
|
||||
associate_data: Option<Box<dyn Any + Send + 'static>>,
|
||||
}
|
||||
|
||||
impl<R, W> TunnelWrapper<R, W> {
|
||||
pub fn new(reader: R, writer: W, info: Option<TunnelInfo>) -> Self {
|
||||
Self::new_with_associate_data(reader, writer, info, None)
|
||||
}
|
||||
|
||||
pub fn new_with_associate_data(
|
||||
reader: R,
|
||||
writer: W,
|
||||
info: Option<TunnelInfo>,
|
||||
associate_data: Option<Box<dyn Any + Send + 'static>>,
|
||||
) -> Self {
|
||||
TunnelWrapper {
|
||||
reader: Arc::new(Mutex::new(Some(reader))),
|
||||
writer: Arc::new(Mutex::new(Some(writer))),
|
||||
info,
|
||||
associate_data,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<R, W> Tunnel for TunnelWrapper<R, W>
|
||||
where
|
||||
R: ZCPacketStream + Send + 'static,
|
||||
W: ZCPacketSink + Send + 'static,
|
||||
{
|
||||
fn split(&self) -> (Pin<Box<dyn ZCPacketStream>>, Pin<Box<dyn ZCPacketSink>>) {
|
||||
let reader = self.reader.lock().unwrap().take().unwrap();
|
||||
let writer = self.writer.lock().unwrap().take().unwrap();
|
||||
(Box::pin(reader), Box::pin(writer))
|
||||
}
|
||||
|
||||
fn info(&self) -> Option<TunnelInfo> {
|
||||
self.info.clone()
|
||||
}
|
||||
}
|
||||
|
||||
// a length delimited codec for async reader
|
||||
pin_project! {
|
||||
pub struct FramedReader<R> {
|
||||
#[pin]
|
||||
reader: R,
|
||||
buf: BytesMut,
|
||||
max_packet_size: usize,
|
||||
associate_data: Option<Box<dyn Any + Send + 'static>>,
|
||||
error: Option<TunnelError>,
|
||||
}
|
||||
}
|
||||
|
||||
impl<R> FramedReader<R> {
|
||||
pub fn new(reader: R, max_packet_size: usize) -> Self {
|
||||
Self::new_with_associate_data(reader, max_packet_size, None)
|
||||
}
|
||||
|
||||
pub fn new_with_associate_data(
|
||||
reader: R,
|
||||
max_packet_size: usize,
|
||||
associate_data: Option<Box<dyn Any + Send + 'static>>,
|
||||
) -> Self {
|
||||
FramedReader {
|
||||
reader,
|
||||
buf: BytesMut::with_capacity(max_packet_size),
|
||||
max_packet_size,
|
||||
associate_data,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn extract_one_packet(
|
||||
buf: &mut BytesMut,
|
||||
max_packet_size: usize,
|
||||
) -> Option<Result<ZCPacket, TunnelError>> {
|
||||
if buf.len() < TCP_TUNNEL_HEADER_SIZE {
|
||||
// header is not complete
|
||||
return None;
|
||||
}
|
||||
|
||||
let header = TCPTunnelHeader::ref_from_prefix(&buf[..]).unwrap();
|
||||
let body_len = header.len.get() as usize;
|
||||
if body_len > max_packet_size {
|
||||
// body is too long
|
||||
return Some(Err(TunnelError::InvalidPacket("body too long".to_string())));
|
||||
}
|
||||
|
||||
if body_len < PEER_MANAGER_HEADER_SIZE {
|
||||
return Some(Err(TunnelError::InvalidPacket(
|
||||
"body too short".to_string(),
|
||||
)));
|
||||
}
|
||||
|
||||
if buf.len() < TCP_TUNNEL_HEADER_SIZE + body_len {
|
||||
// body is not complete
|
||||
return None;
|
||||
}
|
||||
|
||||
// extract one packet
|
||||
let packet_buf = buf.split_to(TCP_TUNNEL_HEADER_SIZE + body_len);
|
||||
Some(Ok(ZCPacket::new_from_buf(packet_buf, ZCPacketType::TCP)))
|
||||
}
|
||||
}
|
||||
|
||||
impl<R> Stream for FramedReader<R>
|
||||
where
|
||||
R: AsyncRead + Send + 'static + Unpin,
|
||||
{
|
||||
type Item = StreamItem;
|
||||
|
||||
fn poll_next(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Option<Self::Item>> {
|
||||
let mut self_mut = self.project();
|
||||
|
||||
loop {
|
||||
if let Some(e) = self_mut.error.as_ref() {
|
||||
tracing::warn!("poll_next on a failed FramedReader, {:?}", e);
|
||||
return Poll::Ready(None);
|
||||
}
|
||||
|
||||
if let Some(packet) = Self::extract_one_packet(self_mut.buf, *self_mut.max_packet_size)
|
||||
{
|
||||
if let Err(TunnelError::InvalidPacket(msg)) = packet.as_ref() {
|
||||
self_mut
|
||||
.error
|
||||
.replace(TunnelError::InvalidPacket(msg.clone()));
|
||||
}
|
||||
return Poll::Ready(Some(packet));
|
||||
}
|
||||
|
||||
reserve_buf(
|
||||
self_mut.buf,
|
||||
*self_mut.max_packet_size,
|
||||
*self_mut.max_packet_size * 2,
|
||||
);
|
||||
|
||||
let cap = self_mut.buf.capacity() - self_mut.buf.len();
|
||||
let buf = self_mut.buf.chunk_mut().as_mut_ptr();
|
||||
let buf = unsafe { std::slice::from_raw_parts_mut(buf, cap) };
|
||||
let mut buf = ReadBuf::new(buf);
|
||||
|
||||
let ret = ready!(self_mut.reader.as_mut().poll_read(cx, &mut buf));
|
||||
let len = buf.filled().len();
|
||||
unsafe { self_mut.buf.advance_mut(len) };
|
||||
|
||||
match ret {
|
||||
Ok(_) => {
|
||||
if len == 0 {
|
||||
return Poll::Ready(None);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
return Poll::Ready(Some(Err(TunnelError::IOError(e))));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait ZCPacketToBytes {
|
||||
fn zcpacket_into_bytes(&self, zc_packet: ZCPacket) -> Result<Bytes, TunnelError>;
|
||||
}
|
||||
|
||||
pub struct TcpZCPacketToBytes;
|
||||
impl ZCPacketToBytes for TcpZCPacketToBytes {
|
||||
fn zcpacket_into_bytes(&self, item: ZCPacket) -> Result<Bytes, TunnelError> {
|
||||
let mut item = item.convert_type(ZCPacketType::TCP);
|
||||
|
||||
let tcp_len = PEER_MANAGER_HEADER_SIZE + item.payload_len();
|
||||
let Some(header) = item.mut_tcp_tunnel_header() else {
|
||||
return Err(TunnelError::InvalidPacket("packet too short".to_string()));
|
||||
};
|
||||
header.len.set(tcp_len.try_into().unwrap());
|
||||
|
||||
Ok(item.into_bytes())
|
||||
}
|
||||
}
|
||||
|
||||
pin_project! {
|
||||
pub struct FramedWriter<W, C> {
|
||||
#[pin]
|
||||
writer: W,
|
||||
sending_bufs: BufList<Bytes>,
|
||||
associate_data: Option<Box<dyn Any + Send + 'static>>,
|
||||
|
||||
converter: C,
|
||||
}
|
||||
}
|
||||
|
||||
impl<W, C> FramedWriter<W, C> {
|
||||
fn max_buffer_count(&self) -> usize {
|
||||
64
|
||||
}
|
||||
}
|
||||
|
||||
impl<W> FramedWriter<W, TcpZCPacketToBytes> {
|
||||
pub fn new(writer: W) -> Self {
|
||||
Self::new_with_associate_data(writer, None)
|
||||
}
|
||||
|
||||
pub fn new_with_associate_data(
|
||||
writer: W,
|
||||
associate_data: Option<Box<dyn Any + Send + 'static>>,
|
||||
) -> Self {
|
||||
FramedWriter {
|
||||
writer,
|
||||
sending_bufs: BufList::new(),
|
||||
associate_data,
|
||||
converter: TcpZCPacketToBytes {},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<W, C: ZCPacketToBytes + Send + 'static> FramedWriter<W, C> {
|
||||
pub fn new_with_converter(writer: W, converter: C) -> Self {
|
||||
Self::new_with_converter_and_associate_data(writer, converter, None)
|
||||
}
|
||||
|
||||
pub fn new_with_converter_and_associate_data(
|
||||
writer: W,
|
||||
converter: C,
|
||||
associate_data: Option<Box<dyn Any + Send + 'static>>,
|
||||
) -> Self {
|
||||
FramedWriter {
|
||||
writer,
|
||||
sending_bufs: BufList::new(),
|
||||
associate_data,
|
||||
converter,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<W, C> Sink<SinkItem> for FramedWriter<W, C>
|
||||
where
|
||||
W: AsyncWrite + Send + 'static,
|
||||
C: ZCPacketToBytes + Send + 'static,
|
||||
{
|
||||
type Error = TunnelError;
|
||||
|
||||
fn poll_ready(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Result<(), Self::Error>> {
|
||||
let max_buffer_count = self.max_buffer_count();
|
||||
if self.sending_bufs.bufs_cnt() >= max_buffer_count {
|
||||
self.as_mut().poll_flush(cx)
|
||||
} else {
|
||||
tracing::trace!(bufs_cnt = self.sending_bufs.bufs_cnt(), "ready to send");
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
fn start_send(self: Pin<&mut Self>, item: ZCPacket) -> Result<(), Self::Error> {
|
||||
let pinned = self.project();
|
||||
pinned
|
||||
.sending_bufs
|
||||
.push(pinned.converter.zcpacket_into_bytes(item)?);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn poll_flush(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
let mut pinned = self.project();
|
||||
let mut remaining = pinned.sending_bufs.remaining();
|
||||
while remaining != 0 {
|
||||
let n = ready!(poll_write_buf(
|
||||
pinned.writer.as_mut(),
|
||||
cx,
|
||||
pinned.sending_bufs
|
||||
))?;
|
||||
if n == 0 {
|
||||
return Poll::Ready(Err(TunnelError::IOError(std::io::Error::new(
|
||||
std::io::ErrorKind::WriteZero,
|
||||
"failed to \
|
||||
write frame to transport",
|
||||
))));
|
||||
}
|
||||
remaining -= n;
|
||||
}
|
||||
|
||||
tracing::trace!(?remaining, "flushed");
|
||||
|
||||
// Try flushing the underlying IO
|
||||
ready!(pinned.writer.poll_flush(cx))?;
|
||||
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn poll_close(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
ready!(self.as_mut().poll_flush(cx))?;
|
||||
ready!(self.project().writer.poll_shutdown(cx))?;
|
||||
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn get_interface_name_by_ip(local_ip: &IpAddr) -> Option<String> {
|
||||
if local_ip.is_unspecified() || local_ip.is_multicast() {
|
||||
@@ -349,27 +23,6 @@ pub(crate) fn get_interface_name_by_ip(local_ip: &IpAddr) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
pub(crate) async fn wait_for_connect_futures<Fut, Ret, E>(
|
||||
mut futures: FuturesUnordered<Fut>,
|
||||
) -> Result<Ret, TunnelError>
|
||||
where
|
||||
Fut: Future<Output = Result<Ret, E>> + Send,
|
||||
E: std::error::Error + Into<TunnelError> + Send + 'static,
|
||||
{
|
||||
// return last error
|
||||
let mut last_err = None;
|
||||
|
||||
while let Some(ret) = futures.next().await {
|
||||
if let Err(e) = ret {
|
||||
last_err = Some(e.into());
|
||||
} else {
|
||||
return ret.map_err(|e| e.into());
|
||||
}
|
||||
}
|
||||
|
||||
Err(last_err.unwrap_or(TunnelError::Shutdown))
|
||||
}
|
||||
|
||||
// region bind
|
||||
|
||||
pub trait Bindable: Sized {
|
||||
@@ -417,6 +70,8 @@ fn setup_socket2_ext(
|
||||
bind_addr: &SocketAddr,
|
||||
#[allow(unused_variables)] bind_dev: Option<String>,
|
||||
only_v6: bool,
|
||||
reuse_addr: bool,
|
||||
reuse_port: bool,
|
||||
socket_mark: Option<u32>,
|
||||
) -> Result<(), TunnelError> {
|
||||
#[cfg(target_os = "windows")]
|
||||
@@ -430,7 +85,15 @@ fn setup_socket2_ext(
|
||||
}
|
||||
|
||||
socket2_socket.set_nonblocking(true)?;
|
||||
socket2_socket.set_reuse_address(!cfg!(target_os = "windows"))?;
|
||||
socket2_socket.set_reuse_address(reuse_addr)?;
|
||||
#[cfg(all(unix, not(target_os = "solaris"), not(target_os = "illumos")))]
|
||||
if reuse_port {
|
||||
socket2_socket.set_reuse_port(true)?;
|
||||
}
|
||||
#[cfg(not(all(unix, not(target_os = "solaris"), not(target_os = "illumos"))))]
|
||||
{
|
||||
let _ = reuse_port;
|
||||
}
|
||||
|
||||
// SO_MARK must be set before bind() so the kernel applies the mark to
|
||||
// any source-address selection bind() triggers on unspecified binds.
|
||||
@@ -445,28 +108,27 @@ fn setup_socket2_ext(
|
||||
}
|
||||
}
|
||||
|
||||
// #[cfg(all(unix, not(target_os = "solaris"), not(target_os = "illumos")))]
|
||||
// socket2_socket.set_reuse_port(true)?;
|
||||
|
||||
if bind_addr.ip().is_unspecified() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// linux/mac does not use interface of bind_addr to send packet, so we need to bind device
|
||||
// win can handle this with bind correctly
|
||||
#[cfg(any(target_os = "ios", target_os = "macos"))]
|
||||
if let Some(dev_name) = bind_dev {
|
||||
// use IP_BOUND_IF to bind device
|
||||
unsafe {
|
||||
let dev_idx = nix::libc::if_nametoindex(dev_name.as_str().as_ptr() as *const i8);
|
||||
tracing::warn!(?dev_idx, ?dev_name, "bind device");
|
||||
if bind_addr.is_ipv4() {
|
||||
socket2_socket.bind_device_by_index_v4(std::num::NonZeroU32::new(dev_idx))?;
|
||||
} else {
|
||||
socket2_socket.bind_device_by_index_v6(std::num::NonZeroU32::new(dev_idx))?;
|
||||
}
|
||||
tracing::warn!(?dev_idx, ?dev_name, "bind device doen");
|
||||
let c_dev_name = std::ffi::CString::new(dev_name.clone()).map_err(|err| {
|
||||
TunnelError::InvalidAddr(format!("invalid interface name {dev_name}: {err}"))
|
||||
})?;
|
||||
let dev_idx = unsafe { nix::libc::if_nametoindex(c_dev_name.as_ptr()) };
|
||||
let Some(dev_idx) = std::num::NonZeroU32::new(dev_idx) else {
|
||||
return Err(TunnelError::InvalidAddr(format!(
|
||||
"network interface not found: {dev_name}"
|
||||
)));
|
||||
};
|
||||
tracing::warn!(?dev_idx, ?dev_name, "bind device");
|
||||
if bind_addr.is_ipv4() {
|
||||
socket2_socket.bind_device_by_index_v4(Some(dev_idx))?;
|
||||
} else {
|
||||
socket2_socket.bind_device_by_index_v6(Some(dev_idx))?;
|
||||
}
|
||||
tracing::warn!(?dev_idx, ?dev_name, "bind device done");
|
||||
}
|
||||
|
||||
#[cfg(any(
|
||||
@@ -559,6 +221,8 @@ pub fn bind<B: Bindable>(
|
||||
#[builder(default, into)] dev: BindDev,
|
||||
net_ns: Option<NetNS>,
|
||||
#[builder(default)] only_v6: bool,
|
||||
#[builder(default = !cfg!(target_os = "windows"))] reuse_addr: bool,
|
||||
#[builder(default)] reuse_port: bool,
|
||||
/// Linux SO_MARK (fwmark) to apply to the socket. `None` leaves SO_MARK
|
||||
/// untouched; `Some(mark)` applies that exact value, including `Some(0)`.
|
||||
socket_mark: Option<u32>,
|
||||
@@ -570,35 +234,34 @@ pub fn bind<B: Bindable>(
|
||||
BindDev::Custom(s) => Some(s),
|
||||
};
|
||||
let socket = socket2::Socket::new(socket2::Domain::for_address(addr), B::TYPE, B::PROTOCOL)?;
|
||||
setup_socket2_ext(&socket, &addr, dev, only_v6, socket_mark)?;
|
||||
setup_socket2_ext(
|
||||
&socket,
|
||||
&addr,
|
||||
dev,
|
||||
only_v6,
|
||||
reuse_addr,
|
||||
reuse_port,
|
||||
socket_mark,
|
||||
)?;
|
||||
B::finalize(socket)
|
||||
}
|
||||
|
||||
pub fn reserve_buf(buf: &mut BytesMut, min_size: usize, max_size: usize) {
|
||||
if buf.capacity() < min_size {
|
||||
buf.reserve(max_size);
|
||||
}
|
||||
}
|
||||
|
||||
// endregion
|
||||
|
||||
pub mod tests {
|
||||
#[cfg(test)]
|
||||
pub(crate) mod tests {
|
||||
use atomic_shim::AtomicU64;
|
||||
use std::{sync::Arc, time::Instant};
|
||||
|
||||
use futures::{Future, SinkExt, StreamExt};
|
||||
use tokio_util::bytes::{BufMut, Bytes, BytesMut};
|
||||
|
||||
use crate::{
|
||||
common::netns::NetNS,
|
||||
tunnel::{TunnelConnector, TunnelListener, packet_def::ZCPacket},
|
||||
use easytier_core::{
|
||||
connectivity::protocol::raw::TunnelDialer, packet::ZCPacket, socket::SocketListener,
|
||||
tunnel::Tunnel,
|
||||
};
|
||||
|
||||
#[cfg(test)]
|
||||
use crate::tunnel::{
|
||||
TunnelError,
|
||||
packet_def::{PEER_MANAGER_HEADER_SIZE, TCP_TUNNEL_HEADER_SIZE},
|
||||
};
|
||||
use crate::common::netns::NetNS;
|
||||
|
||||
#[cfg(any(target_os = "android", target_os = "fuchsia", target_os = "linux"))]
|
||||
#[test]
|
||||
@@ -653,21 +316,26 @@ pub mod tests {
|
||||
assert_eq!(read_so_mark(&socket), 0);
|
||||
}
|
||||
|
||||
#[cfg(any(
|
||||
target_os = "android",
|
||||
target_os = "fuchsia",
|
||||
target_os = "linux",
|
||||
target_env = "ohos"
|
||||
))]
|
||||
#[test]
|
||||
fn framed_reader_rejects_short_peer_manager_body() {
|
||||
let mut buf = BytesMut::new();
|
||||
buf.put_u32_le((PEER_MANAGER_HEADER_SIZE - 1) as u32);
|
||||
buf.resize(TCP_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE - 1, 0);
|
||||
fn bind_custom_device_is_applied_for_unspecified_addr() {
|
||||
use std::net::SocketAddr;
|
||||
use tokio::net::UdpSocket;
|
||||
|
||||
let ret = super::FramedReader::<tokio::io::Empty>::extract_one_packet(&mut buf, 2000);
|
||||
|
||||
assert!(matches!(
|
||||
ret,
|
||||
Some(Err(TunnelError::InvalidPacket(msg))) if msg == "body too short"
|
||||
));
|
||||
let addr: SocketAddr = "0.0.0.0:0".parse().unwrap();
|
||||
let _err = super::bind::<UdpSocket>()
|
||||
.addr(addr)
|
||||
.dev("et/invalid-device-name")
|
||||
.call()
|
||||
.expect_err("custom device must not be skipped for unspecified bind addr");
|
||||
}
|
||||
|
||||
pub async fn _tunnel_echo_server(tunnel: Box<dyn super::Tunnel>, once: bool) {
|
||||
pub async fn _tunnel_echo_server(tunnel: Box<dyn Tunnel>, once: bool) {
|
||||
let (mut recv, mut send) = tunnel.split();
|
||||
|
||||
if !once {
|
||||
@@ -700,33 +368,15 @@ pub mod tests {
|
||||
tracing::warn!("echo server exit...");
|
||||
}
|
||||
|
||||
pub(crate) async fn _tunnel_pingpong<L, C>(listener: L, connector: C)
|
||||
where
|
||||
L: TunnelListener + Send + Sync + 'static,
|
||||
C: TunnelConnector + Send + Sync + 'static,
|
||||
{
|
||||
_tunnel_pingpong_netns_with_timeout(
|
||||
listener,
|
||||
connector,
|
||||
NetNS::new(None),
|
||||
NetNS::new(None),
|
||||
"12345678abcdefg".as_bytes().to_vec(),
|
||||
// only used by tunnel test, so set a long timeout
|
||||
tokio::time::Duration::from_secs(5),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
async fn _tunnel_pingpong_netns<L, C>(
|
||||
mut listener: L,
|
||||
mut connector: C,
|
||||
connector: C,
|
||||
l_netns: NetNS,
|
||||
c_netns: NetNS,
|
||||
buf: Vec<u8>,
|
||||
) where
|
||||
L: TunnelListener + Send + Sync + 'static,
|
||||
C: TunnelConnector + Send + Sync + 'static,
|
||||
L: SocketListener<Accepted = Box<dyn Tunnel>> + Sync + 'static,
|
||||
C: TunnelDialer,
|
||||
{
|
||||
l_netns
|
||||
.run_async(|| async {
|
||||
@@ -791,8 +441,8 @@ pub mod tests {
|
||||
timeout: std::time::Duration,
|
||||
) -> Result<(), anyhow::Error>
|
||||
where
|
||||
L: TunnelListener + Send + Sync + 'static,
|
||||
C: TunnelConnector + Send + Sync + 'static,
|
||||
L: SocketListener<Accepted = Box<dyn Tunnel>> + Sync + 'static,
|
||||
C: TunnelDialer,
|
||||
{
|
||||
let handle = tokio::spawn(async move {
|
||||
_tunnel_pingpong_netns(listener, connector, l_netns, c_netns, buf).await;
|
||||
@@ -821,23 +471,15 @@ pub mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn _tunnel_bench<L, C>(listener: L, connector: C)
|
||||
where
|
||||
L: TunnelListener + Send + Sync + 'static,
|
||||
C: TunnelConnector + Send + Sync + 'static,
|
||||
{
|
||||
_tunnel_bench_netns(listener, connector, NetNS::new(None), NetNS::new(None)).await;
|
||||
}
|
||||
|
||||
pub(crate) async fn _tunnel_bench_netns<L, C>(
|
||||
mut listener: L,
|
||||
mut connector: C,
|
||||
connector: C,
|
||||
netns_l: NetNS,
|
||||
netns_c: NetNS,
|
||||
) -> usize
|
||||
where
|
||||
L: TunnelListener + Send + Sync + 'static,
|
||||
C: TunnelConnector + Send + Sync + 'static,
|
||||
L: SocketListener<Accepted = Box<dyn Tunnel>> + Sync + 'static,
|
||||
C: TunnelDialer,
|
||||
{
|
||||
{
|
||||
let _g = netns_l.guard();
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
Copyright 2021-2025 Datong Sun dndx@idndx.com
|
||||
|
||||
Licensed under the Apache License, Version 2.0 <LICENSE-APACHE or
|
||||
https://www.apache.org/licenses/LICENSE-2.0> or the MIT license <LICENSE-MIT or
|
||||
https://opensource.org/licenses/MIT>, at your option. Files in the project may
|
||||
not be copied, modified, or distributed except according to those terms.
|
||||
@@ -1,594 +0,0 @@
|
||||
mod netfilter;
|
||||
mod packet;
|
||||
mod stack;
|
||||
|
||||
use bytes::BytesMut;
|
||||
use futures::{Sink, Stream};
|
||||
use network_interface::NetworkInterfaceConfig;
|
||||
use pnet::util::MacAddr;
|
||||
use std::{
|
||||
net::{IpAddr, Ipv4Addr, SocketAddr, UdpSocket},
|
||||
pin::Pin,
|
||||
sync::Arc,
|
||||
task::{Context as TaskContext, Poll},
|
||||
};
|
||||
use tokio::{io::AsyncReadExt, net::TcpStream};
|
||||
|
||||
use crate::tunnel::{
|
||||
FromUrl, IpVersion, SinkError, SinkItem, StreamItem, Tunnel, TunnelConnector, TunnelError,
|
||||
TunnelInfo, TunnelListener,
|
||||
common::TunnelWrapper,
|
||||
fake_tcp::netfilter::create_tun,
|
||||
packet_def::{PEER_MANAGER_HEADER_SIZE, TCP_TUNNEL_HEADER_SIZE, ZCPacket, ZCPacketType},
|
||||
};
|
||||
|
||||
use futures::Future;
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
|
||||
use dashmap::DashMap;
|
||||
|
||||
struct IpToIfNameCache {
|
||||
ip_to_ifname: DashMap<IpAddr, (String, Option<MacAddr>)>,
|
||||
}
|
||||
|
||||
impl IpToIfNameCache {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
ip_to_ifname: DashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn reload_ip_to_ifname(&self) {
|
||||
self.ip_to_ifname.clear();
|
||||
let Ok(interfaces) = network_interface::NetworkInterface::show() else {
|
||||
tracing::warn!("failed to enumerate interfaces when reloading faketcp ip cache");
|
||||
return;
|
||||
};
|
||||
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()
|
||||
});
|
||||
for ip in iface.addr.iter() {
|
||||
self.ip_to_ifname.insert(ip.ip(), (iface.name.clone(), mac));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn get_ifname(&self, ip: &IpAddr) -> Option<(String, Option<MacAddr>)> {
|
||||
if let Some(ifname) = self.ip_to_ifname.get(ip) {
|
||||
Some(ifname.clone())
|
||||
} else {
|
||||
self.reload_ip_to_ifname();
|
||||
self.ip_to_ifname.get(ip).map(|s| s.clone())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn get_faketcp_tunnel_type_str(driver_type: &str) -> String {
|
||||
format!("faketcp_{}", driver_type)
|
||||
}
|
||||
|
||||
async fn create_tun_off_runtime(
|
||||
interface_name: String,
|
||||
src_addr: Option<SocketAddr>,
|
||||
dst_addr: SocketAddr,
|
||||
) -> Result<Arc<dyn stack::Tun>, TunnelError> {
|
||||
tokio::task::spawn_blocking(move || create_tun(&interface_name, src_addr, dst_addr))
|
||||
.await
|
||||
.map_err(|e| TunnelError::InternalError(format!("faketcp create_tun task failed: {e}")))?
|
||||
.map_err(Into::into)
|
||||
}
|
||||
|
||||
pub struct FakeTcpTunnelListener {
|
||||
addr: url::Url,
|
||||
os_listener: Option<tokio::net::TcpListener>,
|
||||
// interface_name -> fake tcp stack
|
||||
stack_map: DashMap<String, Arc<stack::Stack>>,
|
||||
// a cache from ip addr to interface name
|
||||
ip_to_ifname: IpToIfNameCache,
|
||||
}
|
||||
|
||||
impl FakeTcpTunnelListener {
|
||||
pub fn new(addr: url::Url) -> Self {
|
||||
// Define filter: Capture all packets (or refine this if needed)
|
||||
// For FakeTCP, we probably want to capture packets destined to us?
|
||||
// But `stack::Stack` handles IP/TCP logic.
|
||||
// Maybe we just capture everything for now as a raw tunnel?
|
||||
// Or better, filter based on some criteria?
|
||||
// The user said "satisfy filter function".
|
||||
// Let's create a filter that accepts everything for now, or maybe only IP packets?
|
||||
FakeTcpTunnelListener {
|
||||
addr,
|
||||
os_listener: None,
|
||||
stack_map: DashMap::new(),
|
||||
ip_to_ifname: IpToIfNameCache::new(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn do_accept(&mut self) -> Result<AcceptResult, TunnelError> {
|
||||
loop {
|
||||
match self.os_listener.as_mut().unwrap().accept().await {
|
||||
Ok((s, remote_addr)) => {
|
||||
let Ok(local_addr) = s.local_addr() else {
|
||||
tracing::warn!("accept fail with local_addr error");
|
||||
continue;
|
||||
};
|
||||
let Some((interface_name, mac)) =
|
||||
self.ip_to_ifname.get_ifname(&local_addr.ip())
|
||||
else {
|
||||
tracing::warn!("accept fail with interface_name error");
|
||||
continue;
|
||||
};
|
||||
return Ok(AcceptResult {
|
||||
socket: s,
|
||||
local_addr,
|
||||
remote_addr,
|
||||
interface_name,
|
||||
mac,
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
use std::io::ErrorKind::*;
|
||||
if matches!(
|
||||
e.kind(),
|
||||
NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset
|
||||
) {
|
||||
tracing::warn!(?e, "accept fail with retryable error: {:?}", e);
|
||||
continue;
|
||||
}
|
||||
tracing::warn!(?e, "accept fail");
|
||||
return Err(e.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_stack(
|
||||
&self,
|
||||
accept_result: &AcceptResult,
|
||||
) -> Result<Arc<stack::Stack>, TunnelError> {
|
||||
let local_socket_addr = accept_result.local_addr;
|
||||
|
||||
let interface_name = &accept_result.interface_name;
|
||||
|
||||
let (local_ip, local_ip6) = match local_socket_addr.ip() {
|
||||
IpAddr::V4(ip) => (Some(ip), None),
|
||||
IpAddr::V6(ip) => (None, Some(ip)),
|
||||
};
|
||||
|
||||
if let Some(entry) = self.stack_map.get(interface_name) {
|
||||
let stack = entry.clone();
|
||||
drop(entry);
|
||||
|
||||
if !stack.is_closed() {
|
||||
return Ok(stack);
|
||||
}
|
||||
|
||||
tracing::warn!(
|
||||
interface_name,
|
||||
"fake_tcp stack reader_task finished, recreating stack"
|
||||
);
|
||||
self.stack_map.remove(interface_name);
|
||||
}
|
||||
|
||||
let tun =
|
||||
create_tun_off_runtime(interface_name.to_string(), None, local_socket_addr).await?;
|
||||
tracing::info!(
|
||||
?local_socket_addr,
|
||||
"create new stack with interface_name: {:?}",
|
||||
interface_name
|
||||
);
|
||||
let stack = Arc::new(stack::Stack::new(
|
||||
tun,
|
||||
local_ip.unwrap_or(Ipv4Addr::UNSPECIFIED),
|
||||
local_ip6,
|
||||
accept_result.mac,
|
||||
));
|
||||
self.stack_map
|
||||
.insert(interface_name.to_string(), stack.clone());
|
||||
|
||||
Ok(stack)
|
||||
}
|
||||
}
|
||||
|
||||
fn build_os_socket_reader_task(mut socket: TcpStream) -> AbortOnDropHandle<()> {
|
||||
AbortOnDropHandle::new(tokio::spawn(async move {
|
||||
// read the os socket until it's closed
|
||||
let mut buf = [0u8; 1024];
|
||||
while let Ok(size) = socket.read(&mut buf).await {
|
||||
tracing::trace!("read {} bytes from os socket", size);
|
||||
if size == 0 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
tracing::info!("FakeTcpTunnelListener os socket closed");
|
||||
}))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct AcceptResult {
|
||||
socket: TcpStream,
|
||||
local_addr: SocketAddr,
|
||||
remote_addr: SocketAddr,
|
||||
interface_name: String,
|
||||
mac: Option<MacAddr>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl TunnelListener for FakeTcpTunnelListener {
|
||||
async fn listen(&mut self) -> Result<(), TunnelError> {
|
||||
let port = self.addr.port().unwrap_or(0);
|
||||
let bind_addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?;
|
||||
let os_listener = tokio::net::TcpListener::bind(bind_addr).await?;
|
||||
tracing::info!(port, "FakeTcpTunnelListener listening");
|
||||
self.os_listener = Some(os_listener);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn accept(&mut self) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
tracing::debug!("FakeTcpTunnelListener waiting for accept");
|
||||
let (res, stack, socket) = loop {
|
||||
let res = self.do_accept().await?;
|
||||
let stack = self.get_stack(&res).await?;
|
||||
let socket = stack.try_alloc_established_socket(
|
||||
res.local_addr,
|
||||
res.remote_addr,
|
||||
stack::State::Established,
|
||||
);
|
||||
let Some(socket) = socket else {
|
||||
tracing::warn!(
|
||||
interface_name = res.interface_name,
|
||||
"fake_tcp stack closed while accepting connection, dropping accepted socket"
|
||||
);
|
||||
self.stack_map.remove(&res.interface_name);
|
||||
continue;
|
||||
};
|
||||
break (res, stack, socket);
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
?res,
|
||||
remote = socket.remote_addr().to_string(),
|
||||
"FakeTcpTunnelListener accepted connection"
|
||||
);
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: get_faketcp_tunnel_type_str(stack.driver_type()),
|
||||
local_addr: Some(self.local_url().into()),
|
||||
remote_addr: Some(
|
||||
crate::tunnel::build_url_from_socket_addr(
|
||||
&socket.remote_addr().to_string(),
|
||||
"faketcp",
|
||||
)
|
||||
.into(),
|
||||
),
|
||||
resolved_remote_addr: Some(
|
||||
crate::tunnel::build_url_from_socket_addr(
|
||||
&socket.remote_addr().to_string(),
|
||||
"faketcp",
|
||||
)
|
||||
.into(),
|
||||
),
|
||||
};
|
||||
|
||||
// We treat the fake tcp socket as a datagram tunnel directly
|
||||
// The reader/writer will interface with the socket using recv_bytes/send
|
||||
// We need to adapt the socket to ZCPacketStream and ZCPacketSink
|
||||
|
||||
// Since FakeTcpTunnel is a datagram tunnel, we don't need FramedReader/Writer (which are for stream based tunnels like TCP)
|
||||
// We should wrap the socket into something that produces/consumes ZCPacket directly.
|
||||
|
||||
let socket = Arc::new(socket);
|
||||
let reader = FakeTcpStream::new(socket.clone());
|
||||
let writer = FakeTcpSink::new(socket);
|
||||
|
||||
Ok(Box::new(TunnelWrapper::new_with_associate_data(
|
||||
reader,
|
||||
writer,
|
||||
Some(info),
|
||||
Some(Box::new(build_os_socket_reader_task(res.socket))),
|
||||
)))
|
||||
}
|
||||
|
||||
fn local_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
}
|
||||
|
||||
pub struct FakeTcpTunnelConnector {
|
||||
addr: url::Url,
|
||||
ip_to_if_name: IpToIfNameCache,
|
||||
resolved_addr: Option<SocketAddr>,
|
||||
socket_mark: Option<u32>,
|
||||
}
|
||||
|
||||
impl FakeTcpTunnelConnector {
|
||||
pub fn new(addr: url::Url) -> Self {
|
||||
FakeTcpTunnelConnector {
|
||||
addr,
|
||||
ip_to_if_name: IpToIfNameCache::new(),
|
||||
resolved_addr: None,
|
||||
socket_mark: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn get_local_ip_for_destination(destination: IpAddr) -> Option<IpAddr> {
|
||||
// 使用一个不可路由的、私有的、或回环地址创建一个临时的 socket,让内核自动选择源接口。
|
||||
// 对于 IPv4,使用 0.0.0.0; 对于 IPv6,使用 ::
|
||||
let bind_addr = if destination.is_ipv4() {
|
||||
IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0))
|
||||
} else {
|
||||
IpAddr::V6(std::net::Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0))
|
||||
};
|
||||
|
||||
// 绑定到一个临时端口 (0)
|
||||
let socket = UdpSocket::bind((bind_addr, 0)).ok()?;
|
||||
|
||||
// 尝试连接到目标地址。这不会真正发送数据包,只是让内核确定路由。
|
||||
socket.connect((destination, 80)).ok()?; // 使用一个常见的端口,例如 80
|
||||
|
||||
// 获取 socket 的本地地址信息
|
||||
socket.local_addr().map(|addr| addr.ip()).ok()
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl TunnelConnector for FakeTcpTunnelConnector {
|
||||
async fn connect(&mut self) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
let remote_addr = match self.resolved_addr {
|
||||
Some(addr) => addr,
|
||||
None => SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?,
|
||||
};
|
||||
let local_ip = get_local_ip_for_destination(remote_addr.ip())
|
||||
.ok_or(TunnelError::InternalError("Failed to get local ip".into()))?;
|
||||
|
||||
let os_socket = tokio::net::TcpSocket::new_v4()?;
|
||||
// SO_MARK applies only to the kernel-visible "decoy" socket below.
|
||||
// The actual FakeTCP payload travels via crafted segments written
|
||||
// straight to the TUN device, which the kernel doesn't tag with
|
||||
// SO_MARK. Operators relying on fwmark for FakeTCP must mark the
|
||||
// TUN device's traffic with a separate nftables/iptables rule.
|
||||
crate::tunnel::common::apply_socket_mark(
|
||||
&socket2::SockRef::from(&os_socket),
|
||||
self.socket_mark,
|
||||
)?;
|
||||
os_socket.bind("0.0.0.0:0".parse().unwrap())?;
|
||||
let local_port = os_socket.local_addr()?.port();
|
||||
let local_addr = SocketAddr::new(local_ip, local_port);
|
||||
|
||||
let (interface_name, mac) =
|
||||
self.ip_to_if_name
|
||||
.get_ifname(&local_ip)
|
||||
.ok_or(TunnelError::InternalError(
|
||||
"Failed to get interface name".into(),
|
||||
))?;
|
||||
|
||||
let (local_ip, local_ip6) = match local_ip {
|
||||
IpAddr::V4(ip) => (Some(ip), None),
|
||||
IpAddr::V6(ip) => (None, Some(ip)),
|
||||
};
|
||||
|
||||
let tun =
|
||||
create_tun_off_runtime(interface_name.clone(), Some(remote_addr), local_addr).await?;
|
||||
let local_ip = local_ip.unwrap_or("0.0.0.0".parse().unwrap());
|
||||
let stack = stack::Stack::new(tun, local_ip, local_ip6, mac);
|
||||
let driver_type = stack.driver_type();
|
||||
|
||||
let socket = stack
|
||||
.try_alloc_established_socket(local_addr, remote_addr, stack::State::SynSent)
|
||||
.ok_or(TunnelError::InternalError(
|
||||
"FakeTCP stack closed while allocating socket".into(),
|
||||
))?;
|
||||
|
||||
let os_stream = os_socket.connect(remote_addr).await?;
|
||||
|
||||
tracing::info!(?remote_addr, "FakeTcpTunnelConnector connecting");
|
||||
|
||||
let mut buf = BytesMut::new();
|
||||
socket
|
||||
.recv(&mut buf)
|
||||
.await
|
||||
.ok_or(TunnelError::InternalError(
|
||||
"Failed to recv bytes to establish connection".into(),
|
||||
))?;
|
||||
|
||||
tracing::info!(local_addr = ?socket.local_addr(), "FakeTcpTunnelConnector connected");
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: get_faketcp_tunnel_type_str(driver_type),
|
||||
local_addr: Some(
|
||||
crate::tunnel::build_url_from_socket_addr(
|
||||
&socket.local_addr().to_string(),
|
||||
"faketcp",
|
||||
)
|
||||
.into(),
|
||||
),
|
||||
remote_addr: Some(self.addr.clone().into()),
|
||||
resolved_remote_addr: Some(
|
||||
crate::tunnel::build_url_from_socket_addr(&remote_addr.to_string(), "faketcp")
|
||||
.into(),
|
||||
),
|
||||
};
|
||||
|
||||
let socket = Arc::new(socket);
|
||||
let reader = FakeTcpStream::new(socket.clone());
|
||||
let writer = FakeTcpSink::new(socket);
|
||||
|
||||
Ok(Box::new(TunnelWrapper::new_with_associate_data(
|
||||
reader,
|
||||
writer,
|
||||
Some(info),
|
||||
Some(Box::new((build_os_socket_reader_task(os_stream), stack))),
|
||||
)))
|
||||
}
|
||||
|
||||
fn remote_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
|
||||
fn set_resolved_addr(&mut self, addr: SocketAddr) {
|
||||
self.resolved_addr = Some(addr);
|
||||
}
|
||||
|
||||
fn set_socket_mark(&mut self, socket_mark: Option<u32>) {
|
||||
self.socket_mark = socket_mark;
|
||||
}
|
||||
}
|
||||
|
||||
type RecvFut = Pin<Box<dyn Future<Output = Option<(BytesMut, usize)>> + Send + Sync>>;
|
||||
|
||||
enum FakeTcpStreamState {
|
||||
ConsumingBuf(BytesMut),
|
||||
PollFuture(RecvFut),
|
||||
Closed,
|
||||
}
|
||||
|
||||
struct FakeTcpStream {
|
||||
socket: Arc<stack::Socket>,
|
||||
state: FakeTcpStreamState,
|
||||
}
|
||||
|
||||
impl FakeTcpStream {
|
||||
fn new(socket: Arc<stack::Socket>) -> Self {
|
||||
Self {
|
||||
socket,
|
||||
state: FakeTcpStreamState::ConsumingBuf(BytesMut::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Stream for FakeTcpStream {
|
||||
type Item = StreamItem;
|
||||
|
||||
fn poll_next(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll<Option<Self::Item>> {
|
||||
let s = self.get_mut();
|
||||
loop {
|
||||
let state = std::mem::replace(&mut s.state, FakeTcpStreamState::Closed);
|
||||
match state {
|
||||
FakeTcpStreamState::ConsumingBuf(buf) => {
|
||||
let buf_len = buf.len();
|
||||
// check peer manager header and split buf out
|
||||
let packet = ZCPacket::new_from_buf(buf, ZCPacketType::TCP);
|
||||
if let Some(tcp_hdr) = packet.tcp_tunnel_header() {
|
||||
let expected_payload_len = tcp_hdr.len.get() as usize;
|
||||
let min_packet_len = TCP_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE;
|
||||
if expected_payload_len < min_packet_len {
|
||||
tracing::warn!(
|
||||
"drop fake tcp packet with invalid length: expected_payload_len={}, min_required={}",
|
||||
expected_payload_len,
|
||||
min_packet_len
|
||||
);
|
||||
s.state = FakeTcpStreamState::Closed;
|
||||
return Poll::Ready(None);
|
||||
}
|
||||
|
||||
if expected_payload_len <= buf_len {
|
||||
let mut buf = packet.inner();
|
||||
let new_inner = buf.split_to(expected_payload_len);
|
||||
s.state = FakeTcpStreamState::ConsumingBuf(buf);
|
||||
return Poll::Ready(Some(Ok(ZCPacket::new_from_buf(
|
||||
new_inner,
|
||||
ZCPacketType::TCP,
|
||||
))));
|
||||
}
|
||||
}
|
||||
|
||||
let mut buf = packet.inner();
|
||||
buf.truncate(0);
|
||||
|
||||
let socket = s.socket.clone();
|
||||
s.state = FakeTcpStreamState::PollFuture(Box::pin(async move {
|
||||
let ret = socket.recv(&mut buf).await;
|
||||
ret.map(|s| (buf, s))
|
||||
}));
|
||||
}
|
||||
FakeTcpStreamState::PollFuture(mut fut) => match fut.as_mut().poll(cx) {
|
||||
Poll::Ready(Some((buf, _sz))) => {
|
||||
s.state = FakeTcpStreamState::ConsumingBuf(buf);
|
||||
}
|
||||
Poll::Ready(None) => {
|
||||
s.state = FakeTcpStreamState::Closed;
|
||||
}
|
||||
Poll::Pending => {
|
||||
s.state = FakeTcpStreamState::PollFuture(fut);
|
||||
return Poll::Pending;
|
||||
}
|
||||
},
|
||||
FakeTcpStreamState::Closed => {
|
||||
return Poll::Ready(None);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct FakeTcpSink {
|
||||
socket: Arc<stack::Socket>,
|
||||
}
|
||||
|
||||
impl FakeTcpSink {
|
||||
fn new(socket: Arc<stack::Socket>) -> Self {
|
||||
Self { socket }
|
||||
}
|
||||
}
|
||||
|
||||
impl Sink<SinkItem> for FakeTcpSink {
|
||||
type Error = SinkError;
|
||||
|
||||
fn poll_ready(
|
||||
self: Pin<&mut Self>,
|
||||
_cx: &mut TaskContext<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn start_send(self: Pin<&mut Self>, item: SinkItem) -> Result<(), Self::Error> {
|
||||
// We need to send the packet as bytes
|
||||
// The item is ZCPacket, which has into_bytes() method
|
||||
let mut packet = item.convert_type(ZCPacketType::TCP);
|
||||
let len = packet.buf_len();
|
||||
packet.mut_tcp_tunnel_header().unwrap().len.set(len as u32);
|
||||
self.socket.try_send(&packet.into_bytes());
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn poll_flush(
|
||||
self: Pin<&mut Self>,
|
||||
_cx: &mut TaskContext<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn poll_close(
|
||||
self: Pin<&mut Self>,
|
||||
_cx: &mut TaskContext<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
self.socket.close();
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tunnel::common::tests::_tunnel_pingpong;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn faketcp_pingpong() {
|
||||
#[cfg(target_family = "unix")]
|
||||
{
|
||||
if unsafe { nix::libc::geteuid() } != 0 {
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
let listener = FakeTcpTunnelListener::new("faketcp://0.0.0.0:31011".parse().unwrap());
|
||||
let connector = FakeTcpTunnelConnector::new("faketcp://127.0.0.1:31011".parse().unwrap());
|
||||
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
}
|
||||
@@ -1,839 +0,0 @@
|
||||
use bytes::Bytes;
|
||||
use bytes::BytesMut;
|
||||
use nix::libc;
|
||||
use std::ffi::CString;
|
||||
use std::io;
|
||||
use std::mem;
|
||||
use std::net::IpAddr;
|
||||
use std::net::SocketAddr;
|
||||
use std::os::fd::{AsRawFd, FromRawFd, OwnedFd};
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering};
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::tunnel::fake_tcp::stack;
|
||||
|
||||
const ETH_HDR_LEN: usize = 14;
|
||||
const ETH_TYPE_OFFSET: u32 = 12;
|
||||
const ETHERTYPE_IPV4: u32 = 0x0800;
|
||||
const ETHERTYPE_IPV6: u32 = 0x86DD;
|
||||
const IPPROTO_TCP_U32: u32 = 6;
|
||||
|
||||
const BPF_LD: u16 = 0x00;
|
||||
const BPF_LDX: u16 = 0x01;
|
||||
const BPF_JMP: u16 = 0x05;
|
||||
const BPF_RET: u16 = 0x06;
|
||||
|
||||
const BPF_W: u16 = 0x00;
|
||||
const BPF_H: u16 = 0x08;
|
||||
const BPF_B: u16 = 0x10;
|
||||
|
||||
const BPF_ABS: u16 = 0x20;
|
||||
const BPF_IND: u16 = 0x40;
|
||||
const BPF_MSH: u16 = 0xa0;
|
||||
|
||||
const BPF_JA: u16 = 0x00;
|
||||
const BPF_JEQ: u16 = 0x10;
|
||||
|
||||
const BPF_K: u16 = 0x00;
|
||||
|
||||
const SOL_PACKET: i32 = 263;
|
||||
const PACKET_STATISTICS: i32 = 6;
|
||||
|
||||
const DEFAULT_RCVBUF_BYTES: i32 = 32 * 1024 * 1024;
|
||||
|
||||
fn stmt(code: u16, k: u32) -> libc::sock_filter {
|
||||
libc::sock_filter {
|
||||
code,
|
||||
jt: 0,
|
||||
jf: 0,
|
||||
k,
|
||||
}
|
||||
}
|
||||
|
||||
fn jeq(k: u32, jt: u8, jf: u8) -> libc::sock_filter {
|
||||
libc::sock_filter {
|
||||
code: BPF_JMP | BPF_JEQ | BPF_K,
|
||||
jt,
|
||||
jf,
|
||||
k,
|
||||
}
|
||||
}
|
||||
|
||||
fn ja(k: u32) -> libc::sock_filter {
|
||||
libc::sock_filter {
|
||||
code: BPF_JMP | BPF_JA,
|
||||
jt: 0,
|
||||
jf: 0,
|
||||
k,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
struct Label(usize);
|
||||
|
||||
struct JeqPatch {
|
||||
idx: usize,
|
||||
t: Label,
|
||||
f: Label,
|
||||
}
|
||||
|
||||
struct JaPatch {
|
||||
idx: usize,
|
||||
target: Label,
|
||||
}
|
||||
|
||||
struct BpfBuilder {
|
||||
insns: Vec<libc::sock_filter>,
|
||||
labels: Vec<Option<usize>>,
|
||||
jeq_patches: Vec<JeqPatch>,
|
||||
ja_patches: Vec<JaPatch>,
|
||||
}
|
||||
|
||||
impl BpfBuilder {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
insns: Vec::new(),
|
||||
labels: Vec::new(),
|
||||
jeq_patches: Vec::new(),
|
||||
ja_patches: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn new_label(&mut self) -> Label {
|
||||
let idx = self.labels.len();
|
||||
self.labels.push(None);
|
||||
Label(idx)
|
||||
}
|
||||
|
||||
fn set_label(&mut self, label: Label) {
|
||||
self.labels[label.0] = Some(self.insns.len());
|
||||
}
|
||||
|
||||
fn push(&mut self, insn: libc::sock_filter) {
|
||||
self.insns.push(insn);
|
||||
}
|
||||
|
||||
fn push_jeq(&mut self, k: u32, t: Label, f: Label) {
|
||||
let idx = self.insns.len();
|
||||
self.insns.push(jeq(k, 0, 0));
|
||||
self.jeq_patches.push(JeqPatch { idx, t, f });
|
||||
}
|
||||
|
||||
fn push_ja(&mut self, target: Label) {
|
||||
let idx = self.insns.len();
|
||||
self.insns.push(ja(0));
|
||||
self.ja_patches.push(JaPatch { idx, target });
|
||||
}
|
||||
|
||||
fn finish(mut self) -> io::Result<Vec<libc::sock_filter>> {
|
||||
for patch in self.jeq_patches {
|
||||
let JeqPatch { idx, t, f } = patch;
|
||||
let cur = idx + 1;
|
||||
let t_pos =
|
||||
self.labels.get(t.0).and_then(|v| *v).ok_or_else(|| {
|
||||
io::Error::new(io::ErrorKind::InvalidInput, "unresolved label")
|
||||
})?;
|
||||
let f_pos =
|
||||
self.labels.get(f.0).and_then(|v| *v).ok_or_else(|| {
|
||||
io::Error::new(io::ErrorKind::InvalidInput, "unresolved label")
|
||||
})?;
|
||||
|
||||
if t_pos < cur || f_pos < cur {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"backward bpf jump",
|
||||
));
|
||||
}
|
||||
|
||||
let jt: u8 = (t_pos - cur)
|
||||
.try_into()
|
||||
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "bpf jump too far"))?;
|
||||
let jf: u8 = (f_pos - cur)
|
||||
.try_into()
|
||||
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "bpf jump too far"))?;
|
||||
|
||||
self.insns[idx].jt = jt;
|
||||
self.insns[idx].jf = jf;
|
||||
}
|
||||
|
||||
for patch in self.ja_patches {
|
||||
let JaPatch { idx, target } = patch;
|
||||
let cur = idx + 1;
|
||||
let t_pos =
|
||||
self.labels.get(target.0).and_then(|v| *v).ok_or_else(|| {
|
||||
io::Error::new(io::ErrorKind::InvalidInput, "unresolved label")
|
||||
})?;
|
||||
if t_pos < cur {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"backward bpf jump",
|
||||
));
|
||||
}
|
||||
self.insns[idx].k = (t_pos - cur) as u32;
|
||||
}
|
||||
|
||||
Ok(self.insns)
|
||||
}
|
||||
}
|
||||
|
||||
fn build_tcp_filter(
|
||||
src_addr: Option<SocketAddr>,
|
||||
dst_addr: SocketAddr,
|
||||
) -> io::Result<Vec<libc::sock_filter>> {
|
||||
if let Some(src) = src_addr
|
||||
&& src.is_ipv4() != dst_addr.is_ipv4()
|
||||
{
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"src/dst addr family mismatch",
|
||||
));
|
||||
}
|
||||
|
||||
let mut b = BpfBuilder::new();
|
||||
let l_check_ipv6 = b.new_label();
|
||||
let l_ipv4 = b.new_label();
|
||||
let l_ipv6 = b.new_label();
|
||||
let l_accept = b.new_label();
|
||||
let l_reject = b.new_label();
|
||||
|
||||
b.push(stmt(BPF_LD | BPF_H | BPF_ABS, ETH_TYPE_OFFSET));
|
||||
b.push_jeq(ETHERTYPE_IPV4, l_ipv4, l_check_ipv6);
|
||||
|
||||
b.set_label(l_check_ipv6);
|
||||
b.push_jeq(ETHERTYPE_IPV6, l_ipv6, l_reject);
|
||||
|
||||
if dst_addr.is_ipv4() {
|
||||
b.set_label(l_ipv4);
|
||||
let l_v4_proto_ok = b.new_label();
|
||||
b.push(stmt(BPF_LD | BPF_B | BPF_ABS, (ETH_HDR_LEN + 9) as u32));
|
||||
b.push_jeq(IPPROTO_TCP_U32, l_v4_proto_ok, l_reject);
|
||||
|
||||
b.set_label(l_v4_proto_ok);
|
||||
let dst_ip = match dst_addr.ip() {
|
||||
IpAddr::V4(ip) => u32::from(ip),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
let l_v4_dstip_ok = b.new_label();
|
||||
b.push(stmt(BPF_LD | BPF_W | BPF_ABS, (ETH_HDR_LEN + 16) as u32));
|
||||
b.push_jeq(dst_ip, l_v4_dstip_ok, l_reject);
|
||||
|
||||
b.set_label(l_v4_dstip_ok);
|
||||
if let Some(src) = src_addr {
|
||||
let src_ip = match src.ip() {
|
||||
IpAddr::V4(ip) => u32::from(ip),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
let l_v4_srcip_ok = b.new_label();
|
||||
b.push(stmt(BPF_LD | BPF_W | BPF_ABS, (ETH_HDR_LEN + 12) as u32));
|
||||
b.push_jeq(src_ip, l_v4_srcip_ok, l_reject);
|
||||
b.set_label(l_v4_srcip_ok);
|
||||
}
|
||||
|
||||
b.push(stmt(BPF_LDX | BPF_B | BPF_MSH, ETH_HDR_LEN as u32));
|
||||
|
||||
let l_v4_dstport_ok = b.new_label();
|
||||
b.push(stmt(BPF_LD | BPF_H | BPF_IND, (ETH_HDR_LEN + 2) as u32));
|
||||
b.push_jeq(dst_addr.port() as u32, l_v4_dstport_ok, l_reject);
|
||||
|
||||
b.set_label(l_v4_dstport_ok);
|
||||
if let Some(src) = src_addr {
|
||||
b.push(stmt(BPF_LD | BPF_H | BPF_IND, ETH_HDR_LEN as u32));
|
||||
b.push_jeq(src.port() as u32, l_accept, l_reject);
|
||||
} else {
|
||||
b.push_ja(l_accept);
|
||||
}
|
||||
} else {
|
||||
b.set_label(l_ipv6);
|
||||
let l_v6_proto_ok = b.new_label();
|
||||
b.push(stmt(BPF_LD | BPF_B | BPF_ABS, (ETH_HDR_LEN + 6) as u32));
|
||||
b.push_jeq(IPPROTO_TCP_U32, l_v6_proto_ok, l_reject);
|
||||
|
||||
b.set_label(l_v6_proto_ok);
|
||||
let dst_ip = match dst_addr.ip() {
|
||||
IpAddr::V6(ip) => ip.octets(),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
for (i, chunk) in dst_ip.chunks_exact(4).enumerate() {
|
||||
let off = ETH_HDR_LEN + 24 + (i * 4);
|
||||
let v = u32::from_be_bytes(chunk.try_into().unwrap());
|
||||
let l_v6_dstip_word_ok = b.new_label();
|
||||
b.push(stmt(BPF_LD | BPF_W | BPF_ABS, off as u32));
|
||||
b.push_jeq(v, l_v6_dstip_word_ok, l_reject);
|
||||
b.set_label(l_v6_dstip_word_ok);
|
||||
}
|
||||
|
||||
if let Some(src) = src_addr {
|
||||
let src_ip = match src.ip() {
|
||||
IpAddr::V6(ip) => ip.octets(),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
for (i, chunk) in src_ip.chunks_exact(4).enumerate() {
|
||||
let off = ETH_HDR_LEN + 8 + (i * 4);
|
||||
let v = u32::from_be_bytes(chunk.try_into().unwrap());
|
||||
let l_v6_srcip_word_ok = b.new_label();
|
||||
b.push(stmt(BPF_LD | BPF_W | BPF_ABS, off as u32));
|
||||
b.push_jeq(v, l_v6_srcip_word_ok, l_reject);
|
||||
b.set_label(l_v6_srcip_word_ok);
|
||||
}
|
||||
}
|
||||
|
||||
let l_v6_dstport_ok = b.new_label();
|
||||
b.push(stmt(
|
||||
BPF_LD | BPF_H | BPF_ABS,
|
||||
(ETH_HDR_LEN + 40 + 2) as u32,
|
||||
));
|
||||
b.push_jeq(dst_addr.port() as u32, l_v6_dstport_ok, l_reject);
|
||||
|
||||
b.set_label(l_v6_dstport_ok);
|
||||
if let Some(src) = src_addr {
|
||||
b.push(stmt(BPF_LD | BPF_H | BPF_ABS, (ETH_HDR_LEN + 40) as u32));
|
||||
b.push_jeq(src.port() as u32, l_accept, l_reject);
|
||||
} else {
|
||||
b.push_ja(l_accept);
|
||||
}
|
||||
}
|
||||
|
||||
b.set_label(l_accept);
|
||||
b.push(stmt(BPF_RET | BPF_K, 0xFFFF));
|
||||
|
||||
b.set_label(l_reject);
|
||||
if dst_addr.is_ipv4() {
|
||||
b.set_label(l_ipv6);
|
||||
} else {
|
||||
b.set_label(l_ipv4);
|
||||
}
|
||||
b.push(stmt(BPF_RET | BPF_K, 0));
|
||||
|
||||
b.finish()
|
||||
}
|
||||
|
||||
#[repr(C)]
|
||||
#[derive(Clone, Copy, Default)]
|
||||
struct PacketSocketStats {
|
||||
tp_packets: u32,
|
||||
tp_drops: u32,
|
||||
}
|
||||
|
||||
fn set_socket_rcvbuf(fd: i32, desired_bytes: i32) -> io::Result<i32> {
|
||||
let ret = unsafe {
|
||||
libc::setsockopt(
|
||||
fd,
|
||||
libc::SOL_SOCKET,
|
||||
libc::SO_RCVBUF,
|
||||
&desired_bytes as *const _ as *const libc::c_void,
|
||||
mem::size_of_val(&desired_bytes) as u32,
|
||||
)
|
||||
};
|
||||
if ret != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
|
||||
let mut actual: i32 = 0;
|
||||
let mut len = mem::size_of_val(&actual) as libc::socklen_t;
|
||||
let ret = unsafe {
|
||||
libc::getsockopt(
|
||||
fd,
|
||||
libc::SOL_SOCKET,
|
||||
libc::SO_RCVBUF,
|
||||
&mut actual as *mut _ as *mut libc::c_void,
|
||||
&mut len as *mut _,
|
||||
)
|
||||
};
|
||||
if ret != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
|
||||
Ok(actual)
|
||||
}
|
||||
|
||||
fn read_packet_socket_stats(fd: i32) -> io::Result<PacketSocketStats> {
|
||||
let mut stats = PacketSocketStats::default();
|
||||
let mut len = mem::size_of_val(&stats) as libc::socklen_t;
|
||||
let ret = unsafe {
|
||||
libc::getsockopt(
|
||||
fd,
|
||||
SOL_PACKET,
|
||||
PACKET_STATISTICS,
|
||||
&mut stats as *mut _ as *mut libc::c_void,
|
||||
&mut len as *mut _,
|
||||
)
|
||||
};
|
||||
if ret != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
Ok(stats)
|
||||
}
|
||||
|
||||
pub struct LinuxBpfTun {
|
||||
fd: Arc<OwnedFd>,
|
||||
ifindex: i32,
|
||||
stop: Arc<AtomicBool>,
|
||||
worker: Option<std::thread::JoinHandle<()>>,
|
||||
recv_queue: Mutex<tokio::sync::mpsc::Receiver<Vec<u8>>>,
|
||||
}
|
||||
|
||||
impl LinuxBpfTun {
|
||||
pub fn new(
|
||||
interface_name: &str,
|
||||
src_addr: Option<SocketAddr>,
|
||||
dst_addr: SocketAddr,
|
||||
) -> io::Result<Self> {
|
||||
let c_ifname = CString::new(interface_name)
|
||||
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "invalid interface name"))?;
|
||||
let ifindex = unsafe { libc::if_nametoindex(c_ifname.as_ptr()) as i32 };
|
||||
if ifindex <= 0 {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
"interface not found",
|
||||
));
|
||||
}
|
||||
|
||||
let proto: i32 = (libc::ETH_P_ALL as u16).to_be() as i32;
|
||||
let fd = unsafe { libc::socket(libc::AF_PACKET, libc::SOCK_RAW, proto) };
|
||||
if fd < 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
let fd = Arc::new(unsafe { OwnedFd::from_raw_fd(fd) });
|
||||
|
||||
let mut addr: libc::sockaddr_ll = unsafe { mem::zeroed() };
|
||||
addr.sll_family = libc::AF_PACKET as u16;
|
||||
addr.sll_protocol = (libc::ETH_P_ALL as u16).to_be();
|
||||
addr.sll_ifindex = ifindex;
|
||||
|
||||
let bind_ret = unsafe {
|
||||
libc::bind(
|
||||
fd.as_ref().as_raw_fd(),
|
||||
&addr as *const _ as *const libc::sockaddr,
|
||||
mem::size_of::<libc::sockaddr_ll>() as u32,
|
||||
)
|
||||
};
|
||||
if bind_ret != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
|
||||
let actual_rcvbuf = set_socket_rcvbuf(fd.as_ref().as_raw_fd(), DEFAULT_RCVBUF_BYTES)?;
|
||||
|
||||
let filter = build_tcp_filter(src_addr, dst_addr)?;
|
||||
let mut prog = libc::sock_fprog {
|
||||
len: filter
|
||||
.len()
|
||||
.try_into()
|
||||
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "bpf program too long"))?,
|
||||
filter: filter.as_ptr() as *mut libc::sock_filter,
|
||||
};
|
||||
let opt_ret = unsafe {
|
||||
libc::setsockopt(
|
||||
fd.as_ref().as_raw_fd(),
|
||||
libc::SOL_SOCKET,
|
||||
libc::SO_ATTACH_FILTER,
|
||||
&mut prog as *mut _ as *mut libc::c_void,
|
||||
mem::size_of::<libc::sock_fprog>() as u32,
|
||||
)
|
||||
};
|
||||
if opt_ret != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
|
||||
let timeout = libc::timeval {
|
||||
tv_sec: 0,
|
||||
tv_usec: 200_000,
|
||||
};
|
||||
let _ = unsafe {
|
||||
libc::setsockopt(
|
||||
fd.as_ref().as_raw_fd(),
|
||||
libc::SOL_SOCKET,
|
||||
libc::SO_RCVTIMEO,
|
||||
&timeout as *const _ as *const libc::c_void,
|
||||
mem::size_of::<libc::timeval>() as u32,
|
||||
)
|
||||
};
|
||||
|
||||
let stop = Arc::new(AtomicBool::new(false));
|
||||
let (tx, rx) = tokio::sync::mpsc::channel(1024);
|
||||
let stop_clone = stop.clone();
|
||||
let read_fd = fd.as_ref().as_raw_fd();
|
||||
let fd_guard = fd.clone();
|
||||
let interface_name_for_worker = interface_name.to_string();
|
||||
|
||||
let worker = std::thread::spawn(move || {
|
||||
// Keep the packet socket alive until the detached worker actually exits.
|
||||
let _fd_guard = fd_guard;
|
||||
let mut buf = vec![0u8; 65536];
|
||||
let mut stats_enabled = true;
|
||||
let mut total_packets: u64 = 0;
|
||||
let mut total_drops: u64 = 0;
|
||||
let mut total_bytes: u64 = 0;
|
||||
let mut dropped_by_queue_full: u64 = 0;
|
||||
let mut last_stats_log = Instant::now();
|
||||
while !stop_clone.load(AtomicOrdering::Relaxed) {
|
||||
let n = unsafe {
|
||||
libc::recv(read_fd, buf.as_mut_ptr() as *mut libc::c_void, buf.len(), 0)
|
||||
};
|
||||
if n < 0 {
|
||||
let err = io::Error::last_os_error();
|
||||
if matches!(
|
||||
err.kind(),
|
||||
io::ErrorKind::Interrupted | io::ErrorKind::WouldBlock
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
break;
|
||||
}
|
||||
if n == 0 {
|
||||
continue;
|
||||
}
|
||||
let data = buf[..(n as usize)].to_vec();
|
||||
total_bytes = total_bytes.wrapping_add(n as u64);
|
||||
match tx.try_send(data) {
|
||||
Ok(()) => {}
|
||||
Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
|
||||
dropped_by_queue_full = dropped_by_queue_full.wrapping_add(1);
|
||||
}
|
||||
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => break,
|
||||
}
|
||||
|
||||
if last_stats_log.elapsed() >= Duration::from_secs(1) {
|
||||
if stats_enabled {
|
||||
match read_packet_socket_stats(read_fd) {
|
||||
Ok(delta) => {
|
||||
total_packets = total_packets.wrapping_add(delta.tp_packets as u64);
|
||||
total_drops = total_drops.wrapping_add(delta.tp_drops as u64);
|
||||
|
||||
let denom =
|
||||
(delta.tp_packets as u64).saturating_add(delta.tp_drops as u64);
|
||||
let drop_rate = if denom == 0 {
|
||||
0.0
|
||||
} else {
|
||||
(delta.tp_drops as f64) / (denom as f64)
|
||||
};
|
||||
|
||||
tracing::debug!(
|
||||
"{}: delta_packets = {}, delta_drops = {}, delta_drop_rate = {}, total_packets = {}, total_drops = {}, total_bytes = {}, dropped_by_queue_full = {}",
|
||||
interface_name_for_worker,
|
||||
delta.tp_packets,
|
||||
delta.tp_drops,
|
||||
drop_rate,
|
||||
total_packets,
|
||||
total_drops,
|
||||
total_bytes,
|
||||
dropped_by_queue_full,
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
stats_enabled = false;
|
||||
tracing::warn!(
|
||||
?e,
|
||||
interface_name_for_worker,
|
||||
"LinuxBpfTun failed to read PACKET_STATISTICS, stats disabled"
|
||||
);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
tracing::debug!(
|
||||
"{}: total_bytes = {}, dropped_by_queue_full = {}",
|
||||
interface_name_for_worker,
|
||||
total_bytes,
|
||||
dropped_by_queue_full,
|
||||
);
|
||||
}
|
||||
last_stats_log = Instant::now();
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
tracing::info!(
|
||||
interface_name,
|
||||
ifindex,
|
||||
desired_rcvbuf = DEFAULT_RCVBUF_BYTES,
|
||||
actual_rcvbuf,
|
||||
"LinuxBpfTun created with filter {:?}",
|
||||
filter
|
||||
);
|
||||
|
||||
Ok(Self {
|
||||
fd,
|
||||
ifindex,
|
||||
stop,
|
||||
worker: Some(worker),
|
||||
recv_queue: Mutex::new(rx),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for LinuxBpfTun {
|
||||
fn drop(&mut self) {
|
||||
self.stop.store(true, AtomicOrdering::Relaxed);
|
||||
let _ = unsafe { libc::shutdown(self.fd.as_ref().as_raw_fd(), libc::SHUT_RD) };
|
||||
if let Some(worker) = self.worker.take() {
|
||||
// Dropping the JoinHandle detaches the worker. The worker holds its own Arc<OwnedFd>
|
||||
// clone, so the packet socket stays valid until recv wakes up and the thread exits.
|
||||
drop(worker);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl stack::Tun for LinuxBpfTun {
|
||||
async fn recv(&self, packet: &mut BytesMut) -> Result<usize, std::io::Error> {
|
||||
let mut rx = self.recv_queue.lock().await;
|
||||
match rx.recv().await {
|
||||
Some(data) => {
|
||||
packet.extend_from_slice(&data);
|
||||
Ok(data.len())
|
||||
}
|
||||
None => Err(std::io::Error::new(
|
||||
std::io::ErrorKind::BrokenPipe,
|
||||
"LinuxBpfTun channel closed",
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn try_send(&self, packet: &Bytes) -> Result<(), std::io::Error> {
|
||||
if packet.len() < 6 {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"packet too short",
|
||||
));
|
||||
}
|
||||
|
||||
let mut addr: libc::sockaddr_ll = unsafe { mem::zeroed() };
|
||||
addr.sll_family = libc::AF_PACKET as u16;
|
||||
addr.sll_protocol = (libc::ETH_P_ALL as u16).to_be();
|
||||
addr.sll_ifindex = self.ifindex;
|
||||
addr.sll_halen = 6;
|
||||
addr.sll_addr[..6].copy_from_slice(&packet[..6]);
|
||||
|
||||
let ret = unsafe {
|
||||
libc::sendto(
|
||||
self.fd.as_ref().as_raw_fd(),
|
||||
packet.as_ptr() as *const libc::c_void,
|
||||
packet.len(),
|
||||
0,
|
||||
&addr as *const _ as *const libc::sockaddr,
|
||||
mem::size_of::<libc::sockaddr_ll>() as u32,
|
||||
)
|
||||
};
|
||||
if ret < 0 {
|
||||
return Err(std::io::Error::last_os_error());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn driver_type(&self) -> &'static str {
|
||||
"linux_bpf"
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, target_os = "linux"))]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
use crate::tunnel::fake_tcp::packet::build_tcp_packet;
|
||||
use crate::tunnel::fake_tcp::stack::Tun;
|
||||
use pnet::datalink;
|
||||
use pnet::packet::tcp::TcpFlags;
|
||||
use pnet::util::MacAddr;
|
||||
use rand::Rng;
|
||||
use std::net::{IpAddr, Ipv4Addr};
|
||||
use tokio::time::{Duration, timeout};
|
||||
|
||||
fn is_root() -> bool {
|
||||
unsafe { libc::geteuid() == 0 }
|
||||
}
|
||||
|
||||
fn pick_interface_v4() -> Option<(String, Ipv4Addr, MacAddr)> {
|
||||
let interfaces = datalink::interfaces();
|
||||
for iface in interfaces {
|
||||
let Some(mac) = iface.mac else {
|
||||
continue;
|
||||
};
|
||||
if iface.is_loopback() {
|
||||
continue;
|
||||
}
|
||||
let ipv4 = iface.ips.iter().find_map(|n| match n.ip() {
|
||||
IpAddr::V4(ip) => Some(ip),
|
||||
IpAddr::V6(_) => None,
|
||||
})?;
|
||||
return Some((iface.name, ipv4, mac));
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn send_raw_frame(interface_name: &str, frame: &[u8]) -> io::Result<()> {
|
||||
if frame.len() < 6 {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"frame too short",
|
||||
));
|
||||
}
|
||||
|
||||
let c_ifname = CString::new(interface_name)
|
||||
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "invalid interface name"))?;
|
||||
let ifindex = unsafe { libc::if_nametoindex(c_ifname.as_ptr()) as i32 };
|
||||
if ifindex <= 0 {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
"interface not found",
|
||||
));
|
||||
}
|
||||
|
||||
let proto: i32 = (libc::ETH_P_ALL as u16).to_be() as i32;
|
||||
let fd = unsafe { libc::socket(libc::AF_PACKET, libc::SOCK_RAW, proto) };
|
||||
if fd < 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
let fd = unsafe { OwnedFd::from_raw_fd(fd) };
|
||||
|
||||
let mut addr: libc::sockaddr_ll = unsafe { mem::zeroed() };
|
||||
addr.sll_family = libc::AF_PACKET as u16;
|
||||
addr.sll_protocol = (libc::ETH_P_ALL as u16).to_be();
|
||||
addr.sll_ifindex = ifindex;
|
||||
addr.sll_halen = 6;
|
||||
addr.sll_addr[..6].copy_from_slice(&frame[..6]);
|
||||
|
||||
let ret = unsafe {
|
||||
libc::sendto(
|
||||
fd.as_raw_fd(),
|
||||
frame.as_ptr() as *const libc::c_void,
|
||||
frame.len(),
|
||||
0,
|
||||
&addr as *const _ as *const libc::sockaddr,
|
||||
mem::size_of::<libc::sockaddr_ll>() as u32,
|
||||
)
|
||||
};
|
||||
if ret < 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn linux_bpf_tun_receives_matching_ipv4_frame() {
|
||||
if !is_root() {
|
||||
eprintln!("linux_bpf_tun_receives_matching_ipv4_frame: skipped (not root)");
|
||||
return;
|
||||
}
|
||||
|
||||
let Some((ifname, dst_ip, mac)) = pick_interface_v4() else {
|
||||
eprintln!("linux_bpf_tun_receives_matching_ipv4_frame: skipped (no suitable iface)");
|
||||
return;
|
||||
};
|
||||
|
||||
let dst_port: u16 = rand::thread_rng().gen_range(40000..60000);
|
||||
let dst_addr = SocketAddr::new(IpAddr::V4(dst_ip), dst_port);
|
||||
eprintln!(
|
||||
"linux_bpf_tun_receives_matching_ipv4_frame: ifname={ifname} dst_addr={dst_addr} mac={mac}"
|
||||
);
|
||||
|
||||
let tun = LinuxBpfTun::new(&ifname, None, dst_addr).unwrap();
|
||||
|
||||
let src_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 123, 0, 1)), 12345);
|
||||
eprintln!(
|
||||
"linux_bpf_tun_receives_matching_ipv4_frame: sending frame src_addr={src_addr} -> dst_addr={dst_addr}"
|
||||
);
|
||||
let frame = build_tcp_packet(
|
||||
mac,
|
||||
mac,
|
||||
src_addr,
|
||||
dst_addr,
|
||||
1,
|
||||
0,
|
||||
TcpFlags::SYN,
|
||||
Some(b"ping"),
|
||||
);
|
||||
|
||||
send_raw_frame(&ifname, &frame).unwrap();
|
||||
|
||||
let mut received = BytesMut::new();
|
||||
let n = timeout(Duration::from_secs(2), tun.recv(&mut received))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
eprintln!(
|
||||
"linux_bpf_tun_receives_matching_ipv4_frame: received {} bytes",
|
||||
n
|
||||
);
|
||||
assert_eq!(n, frame.len());
|
||||
assert_eq!(&received[..], &frame[..]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn linux_bpf_tun_filters_out_non_matching_ipv4_frame() {
|
||||
if !is_root() {
|
||||
eprintln!("linux_bpf_tun_filters_out_non_matching_ipv4_frame: skipped (not root)");
|
||||
return;
|
||||
}
|
||||
|
||||
let Some((ifname, dst_ip, mac)) = pick_interface_v4() else {
|
||||
eprintln!(
|
||||
"linux_bpf_tun_filters_out_non_matching_ipv4_frame: skipped (no suitable iface)"
|
||||
);
|
||||
return;
|
||||
};
|
||||
|
||||
let dst_port: u16 = rand::thread_rng().gen_range(40000..60000);
|
||||
let dst_addr = SocketAddr::new(IpAddr::V4(dst_ip), dst_port);
|
||||
eprintln!(
|
||||
"linux_bpf_tun_filters_out_non_matching_ipv4_frame: ifname={ifname} dst_addr={dst_addr} mac={mac}"
|
||||
);
|
||||
|
||||
let tun = LinuxBpfTun::new(&ifname, None, dst_addr).unwrap();
|
||||
|
||||
let src_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 123, 0, 2)), 23456);
|
||||
let non_matching_dst = SocketAddr::new(IpAddr::V4(dst_ip), dst_port.wrapping_add(1));
|
||||
eprintln!(
|
||||
"linux_bpf_tun_filters_out_non_matching_ipv4_frame: sending non-matching src_addr={src_addr} -> dst_addr={non_matching_dst}"
|
||||
);
|
||||
let non_matching = build_tcp_packet(
|
||||
mac,
|
||||
mac,
|
||||
src_addr,
|
||||
non_matching_dst,
|
||||
1,
|
||||
0,
|
||||
TcpFlags::SYN,
|
||||
Some(b"nope"),
|
||||
);
|
||||
send_raw_frame(&ifname, &non_matching).unwrap();
|
||||
|
||||
let mut received = BytesMut::new();
|
||||
let non_matching_timeout = timeout(Duration::from_millis(400), tun.recv(&mut received))
|
||||
.await
|
||||
.is_err();
|
||||
eprintln!(
|
||||
"linux_bpf_tun_filters_out_non_matching_ipv4_frame: non-matching recv timeout={}",
|
||||
non_matching_timeout
|
||||
);
|
||||
assert!(non_matching_timeout);
|
||||
|
||||
eprintln!(
|
||||
"linux_bpf_tun_filters_out_non_matching_ipv4_frame: sending matching src_addr={src_addr} -> dst_addr={dst_addr}"
|
||||
);
|
||||
let matching = build_tcp_packet(
|
||||
mac,
|
||||
mac,
|
||||
src_addr,
|
||||
dst_addr,
|
||||
2,
|
||||
0,
|
||||
TcpFlags::SYN,
|
||||
Some(b"ok"),
|
||||
);
|
||||
send_raw_frame(&ifname, &matching).unwrap();
|
||||
|
||||
let mut received2 = BytesMut::new();
|
||||
let n = timeout(Duration::from_secs(2), tun.recv(&mut received2))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
eprintln!(
|
||||
"linux_bpf_tun_filters_out_non_matching_ipv4_frame: received {} bytes",
|
||||
n
|
||||
);
|
||||
assert_eq!(n, matching.len());
|
||||
assert_eq!(&received2[..], &matching[..]);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,93 +0,0 @@
|
||||
pub mod pnet;
|
||||
|
||||
use std::{io, net::SocketAddr, sync::Arc};
|
||||
|
||||
cfg_select! {
|
||||
target_os = "linux" => {
|
||||
pub mod linux_bpf;
|
||||
|
||||
pub fn create_tun(
|
||||
interface_name: &str,
|
||||
src_addr: Option<SocketAddr>,
|
||||
dst_addr: SocketAddr,
|
||||
) -> io::Result<Arc<dyn super::stack::Tun>> {
|
||||
match linux_bpf::LinuxBpfTun::new(interface_name, src_addr, dst_addr) {
|
||||
Ok(tun) => Ok(Arc::new(tun)),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
?e,
|
||||
interface_name,
|
||||
"LinuxBpfTun init failed, falling back to PnetTun"
|
||||
);
|
||||
Ok(Arc::new(pnet::PnetTun::new(
|
||||
interface_name,
|
||||
pnet::create_packet_filter(src_addr, dst_addr),
|
||||
)?))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
all(target_os = "macos", not(feature = "macos-ne")) => {
|
||||
pub mod macos_bpf;
|
||||
|
||||
pub fn create_tun(
|
||||
interface_name: &str,
|
||||
src_addr: Option<SocketAddr>,
|
||||
dst_addr: SocketAddr,
|
||||
) -> io::Result<Arc<dyn super::stack::Tun>> {
|
||||
match macos_bpf::MacosBpfTun::new(interface_name, src_addr, dst_addr) {
|
||||
Ok(tun) => Ok(Arc::new(tun)),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
?e,
|
||||
interface_name,
|
||||
"MacosBpfTun init failed, falling back to PnetTun"
|
||||
);
|
||||
Ok(Arc::new(pnet::PnetTun::new(
|
||||
interface_name,
|
||||
pnet::create_packet_filter(src_addr, dst_addr),
|
||||
)?))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
all(windows, any(target_arch = "x86_64", target_arch = "x86")) => {
|
||||
pub mod windivert;
|
||||
|
||||
pub fn create_tun(
|
||||
interface_name: &str,
|
||||
src_addr: Option<SocketAddr>,
|
||||
dst_addr: SocketAddr,
|
||||
) -> io::Result<Arc<dyn super::stack::Tun>> {
|
||||
match windivert::WinDivertTun::new(src_addr, dst_addr) {
|
||||
Ok(tun) => Ok(Arc::new(tun)),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
?e,
|
||||
?dst_addr,
|
||||
"WinDivertTun init failed, falling back to PnetTun"
|
||||
);
|
||||
Ok(Arc::new(pnet::PnetTun::new(
|
||||
interface_name,
|
||||
pnet::create_packet_filter(src_addr, dst_addr),
|
||||
)?))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
_ => {
|
||||
pub fn create_tun(
|
||||
interface_name: &str,
|
||||
src_addr: Option<SocketAddr>,
|
||||
dst_addr: SocketAddr,
|
||||
) -> io::Result<Arc<dyn super::stack::Tun>> {
|
||||
Ok(Arc::new(pnet::PnetTun::new(
|
||||
interface_name,
|
||||
pnet::create_packet_filter(src_addr, dst_addr),
|
||||
)?))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,308 +0,0 @@
|
||||
use std::{
|
||||
io,
|
||||
net::{IpAddr, SocketAddr},
|
||||
sync::{
|
||||
Arc, Weak,
|
||||
atomic::{AtomicU32, Ordering},
|
||||
},
|
||||
};
|
||||
|
||||
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 tokio::sync::Mutex;
|
||||
|
||||
use crate::tunnel::fake_tcp::stack;
|
||||
|
||||
type PacketFilter = Box<dyn Fn(&[u8]) -> bool + Send + Sync>;
|
||||
|
||||
fn filter_tcp_packet(
|
||||
packet: &[u8],
|
||||
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 {
|
||||
return false;
|
||||
};
|
||||
|
||||
match ethernet.get_ethertype() {
|
||||
EtherTypes::Ipv4 => {
|
||||
let ipv4 = if let Some(ipv4) = Ipv4Packet::new(ethernet.payload()) {
|
||||
ipv4
|
||||
} else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if ipv4.get_next_level_protocol() != IpNextHeaderProtocols::Tcp {
|
||||
return false;
|
||||
}
|
||||
|
||||
let tcp = if let Some(tcp) = TcpPacket::new(ipv4.payload()) {
|
||||
tcp
|
||||
} else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if let Some(src_addr) = src_addr {
|
||||
if IpAddr::V4(ipv4.get_source()) != src_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_source() != src_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(dst_addr) = dst_addr {
|
||||
if IpAddr::V4(ipv4.get_destination()) != dst_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_destination() != dst_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
tracing::trace!(
|
||||
?tcp,
|
||||
"FakeTcpTunnelListener 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(),
|
||||
);
|
||||
}
|
||||
EtherTypes::Ipv6 => {
|
||||
let ipv6 = if let Some(ipv6) = Ipv6Packet::new(ethernet.payload()) {
|
||||
ipv6
|
||||
} else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if ipv6.get_next_header() != IpNextHeaderProtocols::Tcp {
|
||||
return false;
|
||||
}
|
||||
|
||||
let tcp = if let Some(tcp) = TcpPacket::new(ipv6.payload()) {
|
||||
tcp
|
||||
} else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if let Some(src_addr) = src_addr {
|
||||
if IpAddr::V6(ipv6.get_source()) != src_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_source() != src_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(dst_addr) = dst_addr {
|
||||
if IpAddr::V6(ipv6.get_destination()) != dst_addr.ip() {
|
||||
return false;
|
||||
}
|
||||
if tcp.get_destination() != dst_addr.port() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
tracing::trace!(
|
||||
?tcp,
|
||||
"FakeTcpTunnelListener packet matched filter, dispatching"
|
||||
);
|
||||
}
|
||||
_ => return false,
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
pub fn create_packet_filter(src_addr: Option<SocketAddr>, dst_addr: SocketAddr) -> PacketFilter {
|
||||
Box::new(move |packet: &[u8]| -> bool {
|
||||
filter_tcp_packet(packet, src_addr.as_ref(), Some(&dst_addr))
|
||||
})
|
||||
}
|
||||
|
||||
struct Subscriber {
|
||||
filter: PacketFilter,
|
||||
sender: tokio::sync::mpsc::Sender<Vec<u8>>,
|
||||
}
|
||||
|
||||
struct InterfaceWorker {
|
||||
tx: Mutex<Box<dyn DataLinkSender>>,
|
||||
subscribers: Arc<DashMap<u32, Subscriber>>,
|
||||
}
|
||||
|
||||
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(_) => return Err(io::Error::other("Unhandled channel type")),
|
||||
Err(e) => return Err(io::Error::other(e)),
|
||||
};
|
||||
|
||||
let subscribers = Arc::new(DashMap::<u32, Subscriber>::new());
|
||||
let subscribers_clone = subscribers.clone();
|
||||
|
||||
std::thread::spawn(move || {
|
||||
loop {
|
||||
match rx.next() {
|
||||
Ok(packet) => {
|
||||
// Iterate over subscribers and send packet if filter matches
|
||||
// Note: DashMap iteration might be slow if many subscribers, but usually few per interface.
|
||||
// For high performance we might need a better structure or read-copy-update.
|
||||
for r in subscribers_clone.iter() {
|
||||
let subscriber = r.value();
|
||||
if (subscriber.filter)(packet) {
|
||||
tracing::trace!(
|
||||
?packet,
|
||||
"InterfaceWorker packet matched filter, dispatching"
|
||||
);
|
||||
// Try send, ignore errors (best effort)
|
||||
let _ = subscriber.sender.try_send(packet.to_vec());
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("InterfaceWorker read error: {}", e);
|
||||
// If interface goes down, we might need to handle it.
|
||||
// For now just break and maybe the worker is dead.
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
Ok(Arc::new(Self {
|
||||
tx: Mutex::new(tx),
|
||||
subscribers,
|
||||
}))
|
||||
}
|
||||
|
||||
fn subscribe(&self, filter: PacketFilter, sender: tokio::sync::mpsc::Sender<Vec<u8>>) -> u32 {
|
||||
static ID_GEN: AtomicU32 = AtomicU32::new(0);
|
||||
let id = ID_GEN.fetch_add(1, Ordering::Relaxed);
|
||||
self.subscribers.insert(id, Subscriber { filter, sender });
|
||||
id
|
||||
}
|
||||
|
||||
fn unsubscribe(&self, id: u32) {
|
||||
self.subscribers.remove(&id);
|
||||
}
|
||||
}
|
||||
|
||||
static INTERFACE_MANAGERS: Lazy<DashMap<String, Weak<InterfaceWorker>>> = Lazy::new(DashMap::new);
|
||||
|
||||
fn get_or_create_worker(interface_name: &str) -> io::Result<Arc<InterfaceWorker>> {
|
||||
// Check if we have an active worker
|
||||
if let Some(worker) = INTERFACE_MANAGERS
|
||||
.get(interface_name)
|
||||
.and_then(|w| w.upgrade())
|
||||
{
|
||||
return Ok(worker);
|
||||
}
|
||||
|
||||
// Need to create new worker.
|
||||
// Lock effectively by using entry API? DashMap entry API might not be enough for complex init.
|
||||
// Let's use a double-check locking style or just accept race condition (creating two workers and one wins).
|
||||
// DashMap doesn't support easy "compute_if_absent" with async or heavy logic without blocking the map shard.
|
||||
|
||||
// But creation is rare.
|
||||
// Let's find interface first.
|
||||
let interfaces = std::panic::catch_unwind(datalink::interfaces)
|
||||
.map_err(|_| io::Error::other("failed to enumerate network interfaces: pnet panicked"))?;
|
||||
let interface = interfaces
|
||||
.into_iter()
|
||||
.find(|iface| iface.name == interface_name)
|
||||
.ok_or_else(|| {
|
||||
io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
format!("Network interface '{}' not found", interface_name),
|
||||
)
|
||||
})?;
|
||||
|
||||
let worker = InterfaceWorker::new(interface)?;
|
||||
INTERFACE_MANAGERS.insert(interface_name.to_string(), Arc::downgrade(&worker));
|
||||
Ok(worker)
|
||||
}
|
||||
|
||||
pub struct PnetTun {
|
||||
worker: Arc<InterfaceWorker>,
|
||||
subscription_id: u32,
|
||||
recv_queue: Mutex<tokio::sync::mpsc::Receiver<Vec<u8>>>,
|
||||
}
|
||||
|
||||
impl PnetTun {
|
||||
pub fn new(interface_name: &str, filter: PacketFilter) -> io::Result<Self> {
|
||||
tracing::debug!(interface_name, "Creating new PnetTun");
|
||||
let worker = get_or_create_worker(interface_name)?;
|
||||
let (tx, rx) = tokio::sync::mpsc::channel(1024);
|
||||
let id = worker.subscribe(filter, tx);
|
||||
|
||||
Ok(Self {
|
||||
worker,
|
||||
subscription_id: id,
|
||||
recv_queue: Mutex::new(rx),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PnetTun {
|
||||
fn drop(&mut self) {
|
||||
tracing::debug!(subscription_id = self.subscription_id, "Dropping PnetTun");
|
||||
self.worker.unsubscribe(self.subscription_id);
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl stack::Tun for PnetTun {
|
||||
async fn recv(&self, packet: &mut BytesMut) -> Result<usize, std::io::Error> {
|
||||
let mut rx = self.recv_queue.lock().await;
|
||||
match rx.recv().await {
|
||||
Some(data) => {
|
||||
tracing::trace!(?data, "PnetTun received packet");
|
||||
packet.extend_from_slice(&data);
|
||||
Ok(data.len())
|
||||
}
|
||||
None => {
|
||||
tracing::warn!("PnetTun recv channel closed");
|
||||
Err(std::io::Error::new(
|
||||
std::io::ErrorKind::BrokenPipe,
|
||||
"PnetTun channel closed",
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn try_send(&self, packet: &Bytes) -> Result<(), std::io::Error> {
|
||||
tracing::trace!(len = packet.len(), "PnetTun try_sending packet");
|
||||
// We need async lock for tx.
|
||||
// try_send is sync. We can use try_lock if available or blocking lock.
|
||||
// tokio::sync::Mutex::try_lock is available.
|
||||
if let Ok(mut tx) = self.worker.tx.try_lock() {
|
||||
tx.send_to(packet, None)
|
||||
.ok_or(std::io::Error::other("send_to failed"))?
|
||||
} else {
|
||||
Err(std::io::Error::new(
|
||||
std::io::ErrorKind::WouldBlock,
|
||||
"PnetTun tx lock busy",
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
fn driver_type(&self) -> &'static str {
|
||||
"pnet"
|
||||
}
|
||||
}
|
||||
@@ -1,225 +0,0 @@
|
||||
use std::cell::UnsafeCell;
|
||||
use std::io;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
|
||||
use anyhow::Context as _;
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use tokio::sync::Mutex;
|
||||
use windivert::error::WinDivertError;
|
||||
use windivert::packet::WinDivertPacket;
|
||||
use windivert::prelude::{WinDivertFlags, WinDivertShutdownMode};
|
||||
use windivert::{WinDivert, layer};
|
||||
|
||||
use crate::tunnel::fake_tcp::stack;
|
||||
|
||||
struct WinDivertReader {
|
||||
inner: UnsafeCell<WinDivert<layer::NetworkLayer>>,
|
||||
}
|
||||
|
||||
unsafe impl Send for WinDivertReader {}
|
||||
unsafe impl Sync for WinDivertReader {}
|
||||
|
||||
impl WinDivertReader {
|
||||
fn new(inner: WinDivert<layer::NetworkLayer>) -> Self {
|
||||
Self {
|
||||
inner: UnsafeCell::new(inner),
|
||||
}
|
||||
}
|
||||
|
||||
fn recv<'a>(
|
||||
&self,
|
||||
buffer: Option<&'a mut [u8]>,
|
||||
) -> Result<WinDivertPacket<'a, layer::NetworkLayer>, WinDivertError> {
|
||||
let inner = unsafe { &*self.inner.get() };
|
||||
inner.recv(buffer)
|
||||
}
|
||||
|
||||
fn shutdown(&self) -> anyhow::Result<()> {
|
||||
let inner = unsafe { &mut *self.inner.get() };
|
||||
inner
|
||||
.shutdown(WinDivertShutdownMode::Recv)
|
||||
.with_context(|| "WinDivertReader shutdown failed")?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn close(&self) -> anyhow::Result<()> {
|
||||
let inner = unsafe { &mut *self.inner.get() };
|
||||
inner
|
||||
.close(windivert::CloseAction::Nothing)
|
||||
.with_context(|| "WinDivertReader close failed")?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for WinDivertReader {
|
||||
fn drop(&mut self) {
|
||||
if let Err(e) = self.close() {
|
||||
tracing::error!("WinDivertReader close failed: {:?}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct WinDivertTun {
|
||||
recv_queue: Mutex<tokio::sync::mpsc::Receiver<Vec<u8>>>,
|
||||
sender: Arc<std::sync::Mutex<WinDivert<layer::NetworkLayer>>>,
|
||||
reader: Arc<WinDivertReader>,
|
||||
}
|
||||
|
||||
impl Drop for WinDivertTun {
|
||||
fn drop(&mut self) {
|
||||
if let Ok(mut sender) = self.sender.lock()
|
||||
&& let Err(e) = sender.close(windivert::CloseAction::Nothing)
|
||||
{
|
||||
tracing::error!("WinDivertSender close failed: {:?}", e);
|
||||
}
|
||||
if let Err(e) = self.reader.shutdown() {
|
||||
tracing::error!("WinDivertReader shutdown failed: {:?}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl WinDivertTun {
|
||||
pub fn new(src_addr: Option<SocketAddr>, dst_addr: SocketAddr) -> io::Result<Self> {
|
||||
let (tx, rx) = tokio::sync::mpsc::channel(1024);
|
||||
|
||||
let filter = build_filter(src_addr, dst_addr)?;
|
||||
tracing::debug!(%filter, "WinDivertTun created with filter");
|
||||
|
||||
// Sniff mode: 1 (WINDIVERT_FLAG_SNIFF)
|
||||
// Layer: Network (0)
|
||||
// Priority: 0
|
||||
let flags = WinDivertFlags::default().set_sniff();
|
||||
let reader = WinDivert::network(&filter, 0, flags).map_err(io::Error::other)?;
|
||||
let reader = Arc::new(WinDivertReader::new(reader));
|
||||
let reader_clone = reader.clone();
|
||||
|
||||
std::thread::spawn(move || {
|
||||
let reader = reader_clone;
|
||||
let mut buffer = vec![0u8; 65536];
|
||||
loop {
|
||||
match reader.recv(Some(&mut buffer)) {
|
||||
Ok(packet) => {
|
||||
let data = &packet.data;
|
||||
|
||||
let mut eth_data = vec![0u8; 14 + data.len()];
|
||||
// Set EtherType
|
||||
if !data.is_empty() && data[0] >> 4 == 4 {
|
||||
eth_data[12] = 0x08;
|
||||
eth_data[13] = 0x00;
|
||||
} else {
|
||||
eth_data[12] = 0x86;
|
||||
eth_data[13] = 0xDD;
|
||||
}
|
||||
eth_data[14..].copy_from_slice(data);
|
||||
|
||||
if tx.blocking_send(eth_data).is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::error!(?error, "WinDivert recv failed");
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Sender: non-sniff, empty filter?
|
||||
// Use "false" to avoid capturing anything.
|
||||
// Flags: 0
|
||||
let sender =
|
||||
WinDivert::network("false", 0, WinDivertFlags::default()).map_err(io::Error::other)?;
|
||||
|
||||
Ok(Self {
|
||||
recv_queue: Mutex::new(rx),
|
||||
sender: Arc::new(std::sync::Mutex::new(sender)),
|
||||
reader,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn build_filter(src_addr: Option<SocketAddr>, dst_addr: SocketAddr) -> io::Result<String> {
|
||||
if let Some(src_addr) = src_addr
|
||||
&& src_addr.is_ipv4() != dst_addr.is_ipv4()
|
||||
{
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"src/dst addr family mismatch",
|
||||
));
|
||||
}
|
||||
|
||||
let mut filters = Vec::with_capacity(5);
|
||||
filters.push("tcp".to_owned());
|
||||
|
||||
match dst_addr {
|
||||
SocketAddr::V4(addr) => {
|
||||
filters.push(format!("ip.DstAddr == {}", addr.ip()));
|
||||
filters.push(format!("tcp.DstPort == {}", addr.port()));
|
||||
}
|
||||
SocketAddr::V6(addr) => {
|
||||
filters.push(format!("ipv6.DstAddr == {}", addr.ip()));
|
||||
filters.push(format!("tcp.DstPort == {}", addr.port()));
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(src_addr) = src_addr {
|
||||
match src_addr {
|
||||
SocketAddr::V4(addr) => {
|
||||
filters.push(format!("ip.SrcAddr == {}", addr.ip()));
|
||||
filters.push(format!("tcp.SrcPort == {}", addr.port()));
|
||||
}
|
||||
SocketAddr::V6(addr) => {
|
||||
filters.push(format!("ipv6.SrcAddr == {}", addr.ip()));
|
||||
filters.push(format!("tcp.SrcPort == {}", addr.port()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(filters.join(" and "))
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl stack::Tun for WinDivertTun {
|
||||
async fn recv(&self, packet: &mut BytesMut) -> Result<usize, std::io::Error> {
|
||||
let mut rx = self.recv_queue.lock().await;
|
||||
match rx.recv().await {
|
||||
Some(data) => {
|
||||
packet.extend_from_slice(&data);
|
||||
Ok(data.len())
|
||||
}
|
||||
None => Err(std::io::Error::new(
|
||||
std::io::ErrorKind::BrokenPipe,
|
||||
"Channel closed",
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn try_send(&self, packet: &Bytes) -> Result<(), std::io::Error> {
|
||||
// Strip ethernet header
|
||||
if packet.len() < 14 {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"Packet too short",
|
||||
));
|
||||
}
|
||||
let ip_data = &packet[14..];
|
||||
|
||||
let Ok(sender) = self.sender.try_lock() else {
|
||||
return Err(std::io::Error::other("WinDivert sender lock failed"));
|
||||
};
|
||||
|
||||
let mut pkt = unsafe { WinDivertPacket::<layer::NetworkLayer>::new(ip_data.to_vec()) };
|
||||
pkt.address.set_outbound(true);
|
||||
|
||||
sender
|
||||
.send(&pkt)
|
||||
.map_err(|e| std::io::Error::other(format!("WinDivert send failed: {}", e)))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn driver_type(&self) -> &'static str {
|
||||
"windivert"
|
||||
}
|
||||
}
|
||||
@@ -1,280 +0,0 @@
|
||||
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 std::net::{IpAddr, SocketAddr};
|
||||
|
||||
const IPV4_HEADER_LEN: usize = 20;
|
||||
const IPV6_HEADER_LEN: usize = 40;
|
||||
const TCP_HEADER_LEN: usize = 20;
|
||||
pub const MAX_PACKET_LEN: usize = 1500;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum IPPacket<'p> {
|
||||
V4(ipv4::Ipv4Packet<'p>),
|
||||
V6(ipv6::Ipv6Packet<'p>),
|
||||
}
|
||||
|
||||
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()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_destination(&self) -> IpAddr {
|
||||
match self {
|
||||
IPPacket::V4(p) => IpAddr::V4(p.get_destination()),
|
||||
IPPacket::V6(p) => IpAddr::V6(p.get_destination()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const ETH_HDR_LEN: usize = 14;
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn build_tcp_packet(
|
||||
src_mac: MacAddr,
|
||||
dst_mac: MacAddr,
|
||||
local_addr: SocketAddr,
|
||||
remote_addr: SocketAddr,
|
||||
seq: u32,
|
||||
ack: u32,
|
||||
flags: u8,
|
||||
payload: Option<&[u8]>,
|
||||
) -> Bytes {
|
||||
let ip_header_len = match local_addr {
|
||||
SocketAddr::V4(_) => IPV4_HEADER_LEN,
|
||||
SocketAddr::V6(_) => IPV6_HEADER_LEN,
|
||||
};
|
||||
let wscale = (flags & tcp::TcpFlags::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 eth_buf = buf.split_to(ETH_HDR_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 });
|
||||
if wscale {
|
||||
let wscale = tcp::TcpOption::wscale(14);
|
||||
tcp.set_options(&[tcp::TcpOption::nop(), wscale]);
|
||||
}
|
||||
|
||||
if let Some(payload) = payload {
|
||||
tcp.set_payload(payload);
|
||||
}
|
||||
|
||||
let mut ethernet = MutableEthernetPacket::new(&mut eth_buf).unwrap();
|
||||
ethernet.set_destination(dst_mac);
|
||||
ethernet.set_source(src_mac);
|
||||
ethernet.set_ethertype(match local_addr {
|
||||
SocketAddr::V4(_) => EtherTypes::Ipv4,
|
||||
SocketAddr::V6(_) => EtherTypes::Ipv6,
|
||||
});
|
||||
|
||||
match (local_addr, remote_addr) {
|
||||
(SocketAddr::V4(local), SocketAddr::V4(remote)) => {
|
||||
let mut v4 = ipv4::MutableIpv4Packet::new(&mut ip_buf).unwrap();
|
||||
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);
|
||||
|
||||
tcp.set_checksum(tcp::ipv4_checksum(
|
||||
&tcp.to_immutable(),
|
||||
&v4.get_source(),
|
||||
&v4.get_destination(),
|
||||
));
|
||||
|
||||
v4.set_checksum(ipv4::checksum(&v4.to_immutable()));
|
||||
}
|
||||
(SocketAddr::V6(local), SocketAddr::V6(remote)) => {
|
||||
let mut v6 = ipv6::MutableIpv6Packet::new(&mut ip_buf).unwrap();
|
||||
v6.set_version(6);
|
||||
v6.set_payload_length(tcp_total_len.try_into().unwrap());
|
||||
v6.set_next_header(ip::IpNextHeaderProtocols::Tcp);
|
||||
v6.set_hop_limit(64);
|
||||
v6.set_source(*local.ip());
|
||||
v6.set_destination(*remote.ip());
|
||||
|
||||
tcp.set_checksum(tcp::ipv6_checksum(
|
||||
&tcp.to_immutable(),
|
||||
&v6.get_source(),
|
||||
&v6.get_destination(),
|
||||
));
|
||||
}
|
||||
_ => unreachable!(),
|
||||
};
|
||||
|
||||
ip_buf.unsplit(tcp_buf);
|
||||
eth_buf.unsplit(ip_buf);
|
||||
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();
|
||||
|
||||
tracing::trace!("Parsing IP packet: {:?}", eth);
|
||||
|
||||
let ip_payload = &buf[ETH_HDR_LEN..];
|
||||
|
||||
match ethertype {
|
||||
EtherTypes::Ipv4 => {
|
||||
let v4 = ipv4::Ipv4Packet::new(ip_payload)?;
|
||||
if v4.get_next_level_protocol() != ip::IpNextHeaderProtocols::Tcp {
|
||||
return None;
|
||||
}
|
||||
|
||||
let tcp_offset = usize::from(v4.get_header_length()) * 4;
|
||||
if tcp_offset < IPV4_HEADER_LEN || tcp_offset > ip_payload.len() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let tcp = tcp::TcpPacket::new(&ip_payload[tcp_offset..])?;
|
||||
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 {
|
||||
return None;
|
||||
}
|
||||
|
||||
let tcp = tcp::TcpPacket::new(&ip_payload[IPV6_HEADER_LEN..])?;
|
||||
Some((src_mac, dst_mac, IPPacket::V6(v6), tcp))
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[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 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";
|
||||
|
||||
let packet = build_tcp_packet(
|
||||
src_mac,
|
||||
dst_mac,
|
||||
local_addr,
|
||||
remote_addr,
|
||||
10,
|
||||
20,
|
||||
tcp::TcpFlags::ACK,
|
||||
Some(payload),
|
||||
);
|
||||
|
||||
let (parsed_src_mac, parsed_dst_mac, ip_packet, tcp_packet) =
|
||||
parse_ip_packet(&packet).unwrap();
|
||||
|
||||
assert_eq!(parsed_src_mac, src_mac);
|
||||
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.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 local_addr: SocketAddr = "[2001:db8::1]:12345".parse().unwrap();
|
||||
let remote_addr: SocketAddr = "[2001:db8::2]:23456".parse().unwrap();
|
||||
let payload = b"ipv6 payload";
|
||||
|
||||
let packet = build_tcp_packet(
|
||||
src_mac,
|
||||
dst_mac,
|
||||
local_addr,
|
||||
remote_addr,
|
||||
30,
|
||||
40,
|
||||
tcp::TcpFlags::ACK,
|
||||
Some(payload),
|
||||
);
|
||||
|
||||
let ethernet = EthernetPacket::new(packet.as_ref()).unwrap();
|
||||
assert_eq!(ethernet.get_ethertype(), EtherTypes::Ipv6);
|
||||
|
||||
let (parsed_src_mac, parsed_dst_mac, ip_packet, tcp_packet) =
|
||||
parse_ip_packet(&packet).unwrap();
|
||||
|
||||
assert_eq!(parsed_src_mac, src_mac);
|
||||
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.payload(), payload);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_rejects_short_ethernet_frame() {
|
||||
let packet = Bytes::from_static(&[0u8; ETH_HDR_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),
|
||||
"192.0.2.10:1111".parse().unwrap(),
|
||||
"198.51.100.20:2222".parse().unwrap(),
|
||||
1,
|
||||
2,
|
||||
tcp::TcpFlags::ACK,
|
||||
None,
|
||||
);
|
||||
let truncated = Bytes::copy_from_slice(&packet[..ETH_HDR_LEN + IPV4_HEADER_LEN + 10]);
|
||||
|
||||
assert!(parse_ip_packet(&truncated).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),
|
||||
"[2001:db8::10]:1111".parse().unwrap(),
|
||||
"[2001:db8::20]:2222".parse().unwrap(),
|
||||
1,
|
||||
2,
|
||||
tcp::TcpFlags::ACK,
|
||||
None,
|
||||
);
|
||||
let truncated = Bytes::copy_from_slice(&packet[..ETH_HDR_LEN + IPV6_HEADER_LEN - 1]);
|
||||
|
||||
assert!(parse_ip_packet(&truncated).is_none());
|
||||
}
|
||||
}
|
||||
@@ -1,695 +0,0 @@
|
||||
//! A minimum, userspace TCP based datagram stack
|
||||
//!
|
||||
//! # Overview
|
||||
//!
|
||||
//! `fake-tcp` is a reusable library that implements a minimum TCP stack in
|
||||
//! user space using the Tun interface. It allows programs to send datagrams
|
||||
//! as if they are part of a TCP connection. `fake-tcp` has been tested to
|
||||
//! be able to pass through a variety of NAT and stateful firewalls while
|
||||
//! fully preserves certain desirable behavior such as out of order delivery
|
||||
//! and no congestion/flow controls.
|
||||
//!
|
||||
//! # Core Concepts
|
||||
//!
|
||||
//! The core of the `fake-tcp` crate compose of two structures. [`Stack`] and
|
||||
//! [`Socket`].
|
||||
//!
|
||||
//! ## [`Stack`]
|
||||
//!
|
||||
//! [`Stack`] represents a virtual TCP stack that operates at
|
||||
//! Layer 3. It is responsible for:
|
||||
//!
|
||||
//! * TCP active and passive open and handshake
|
||||
//! * `RST` handling
|
||||
//! * Interact with the Tun interface at Layer 3
|
||||
//! * Distribute incoming datagrams to corresponding [`Socket`]
|
||||
//!
|
||||
//! ## [`Socket`]
|
||||
//!
|
||||
//! [`Socket`] represents a TCP connection. It registers the identifying
|
||||
//! tuple `(src_ip, src_port, dest_ip, dest_port)` inside the [`Stack`] so
|
||||
//! so that incoming packets can be distributed to the right [`Socket`] with
|
||||
//! using a channel. It is also what the client should use for
|
||||
//! sending/receiving datagrams.
|
||||
//!
|
||||
//! # Examples
|
||||
//!
|
||||
//! Please see [`client.rs`](https://github.com/dndx/phantun/blob/main/phantun/src/bin/client.rs)
|
||||
//! and [`server.rs`](https://github.com/dndx/phantun/blob/main/phantun/src/bin/server.rs) files
|
||||
//! from the `phantun` crate for how to use this library in client/server mode, respectively.
|
||||
|
||||
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, HashSet};
|
||||
use std::fmt;
|
||||
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
|
||||
use std::sync::{
|
||||
Arc, RwLock,
|
||||
atomic::{AtomicU32, Ordering},
|
||||
};
|
||||
use tokio::sync::broadcast;
|
||||
use tokio::time;
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
use tracing::{error, info, trace, warn};
|
||||
|
||||
const TIMEOUT: time::Duration = time::Duration::from_secs(1);
|
||||
const RETRIES: usize = 6;
|
||||
const MPMC_BUFFER_LEN: usize = 512;
|
||||
const MAX_UNACKED_LEN: u32 = 128 * 1024 * 1024; // 128MB
|
||||
|
||||
#[async_trait::async_trait]
|
||||
pub trait Tun: Send + Sync + 'static {
|
||||
async fn recv(&self, packet: &mut BytesMut) -> Result<usize, std::io::Error>;
|
||||
fn try_send(&self, packet: &Bytes) -> Result<(), std::io::Error>;
|
||||
fn driver_type(&self) -> &'static str;
|
||||
}
|
||||
|
||||
#[derive(Hash, Eq, PartialEq, Clone, Debug)]
|
||||
struct AddrTuple {
|
||||
local_addr: SocketAddr,
|
||||
remote_addr: SocketAddr,
|
||||
}
|
||||
|
||||
impl AddrTuple {
|
||||
fn new(local_addr: SocketAddr, remote_addr: SocketAddr) -> AddrTuple {
|
||||
AddrTuple {
|
||||
local_addr,
|
||||
remote_addr,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct StackState {
|
||||
tuples: HashMap<AddrTuple, flume::Sender<Bytes>>,
|
||||
closed: bool,
|
||||
}
|
||||
|
||||
struct Shared {
|
||||
state: RwLock<StackState>,
|
||||
listening: RwLock<HashSet<u16>>,
|
||||
tun: Arc<dyn Tun>,
|
||||
tuples_purge: broadcast::Sender<AddrTuple>,
|
||||
}
|
||||
|
||||
impl Shared {
|
||||
fn is_closed(&self) -> bool {
|
||||
self.state.read().unwrap().closed
|
||||
}
|
||||
|
||||
fn mark_closed_and_clear_tuples(&self) -> usize {
|
||||
let mut state = self.state.write().unwrap();
|
||||
state.closed = true;
|
||||
let len = state.tuples.len();
|
||||
state.tuples.clear();
|
||||
len
|
||||
}
|
||||
}
|
||||
|
||||
pub struct Stack {
|
||||
shared: Arc<Shared>,
|
||||
local_ip: Ipv4Addr,
|
||||
local_ip6: Option<Ipv6Addr>,
|
||||
local_mac: MacAddr,
|
||||
reader_task: AbortOnDropHandle<()>,
|
||||
}
|
||||
|
||||
#[derive(Hash, Eq, PartialEq, Clone, Copy, Debug)]
|
||||
pub enum State {
|
||||
Idle,
|
||||
SynSent,
|
||||
SynReceived,
|
||||
Established,
|
||||
}
|
||||
|
||||
pub struct Socket {
|
||||
shared: Arc<Shared>,
|
||||
tun: Arc<dyn Tun>,
|
||||
incoming: flume::Receiver<Bytes>,
|
||||
local_addr: SocketAddr,
|
||||
remote_addr: SocketAddr,
|
||||
local_mac: MacAddr,
|
||||
remote_mac: AtomicCell<Option<MacAddr>>,
|
||||
seq: AtomicU32,
|
||||
ack: AtomicU32,
|
||||
last_ack: AtomicU32,
|
||||
state: AtomicCell<State>,
|
||||
}
|
||||
|
||||
/// A socket that represents a unique TCP connection between a server and client.
|
||||
///
|
||||
/// The `Socket` object itself satisfies `Sync` and `Send`, which means it can
|
||||
/// be safely called within an async future.
|
||||
///
|
||||
/// To close a TCP connection that is no longer needed, simply drop this object
|
||||
/// out of scope.
|
||||
impl Socket {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn new(
|
||||
shared: Arc<Shared>,
|
||||
tun: Arc<dyn Tun>,
|
||||
local_addr: SocketAddr,
|
||||
remote_addr: SocketAddr,
|
||||
local_mac: MacAddr,
|
||||
remote_mac: Option<MacAddr>,
|
||||
ack: Option<u32>,
|
||||
state: State,
|
||||
) -> (Socket, flume::Sender<Bytes>) {
|
||||
let (incoming_tx, incoming_rx) = flume::bounded(MPMC_BUFFER_LEN);
|
||||
|
||||
(
|
||||
Socket {
|
||||
shared,
|
||||
tun,
|
||||
incoming: incoming_rx,
|
||||
local_addr,
|
||||
remote_addr,
|
||||
local_mac,
|
||||
remote_mac: AtomicCell::new(remote_mac),
|
||||
seq: AtomicU32::new(0),
|
||||
ack: AtomicU32::new(ack.unwrap_or(0)),
|
||||
last_ack: AtomicU32::new(ack.unwrap_or(0)),
|
||||
state: AtomicCell::new(state),
|
||||
},
|
||||
incoming_tx,
|
||||
)
|
||||
}
|
||||
|
||||
fn build_tcp_packet(&self, flags: u8, payload: Option<&[u8]>) -> Bytes {
|
||||
let ack = self.ack.load(Ordering::Relaxed);
|
||||
self.last_ack.store(ack, Ordering::Relaxed);
|
||||
|
||||
build_tcp_packet(
|
||||
self.local_mac,
|
||||
self.remote_mac.load().unwrap_or(MacAddr::zero()),
|
||||
self.local_addr,
|
||||
self.remote_addr,
|
||||
self.seq.load(Ordering::Relaxed),
|
||||
ack,
|
||||
flags,
|
||||
payload,
|
||||
)
|
||||
}
|
||||
|
||||
/// Sends a datagram to the other end.
|
||||
///
|
||||
/// This method takes `&self`, and it can be called safely by multiple threads
|
||||
/// at the same time.
|
||||
///
|
||||
/// A return of `None` means the Tun socket returned an error
|
||||
/// and this socket must be closed.
|
||||
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));
|
||||
self.seq.fetch_add(payload.len() as u32, Ordering::Relaxed);
|
||||
self.tun.try_send(&buf).ok().and(Some(()))
|
||||
}
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn close(&self) {
|
||||
if self.state.load() != State::Idle {
|
||||
let buf = self.build_tcp_packet(tcp::TcpFlags::RST, None);
|
||||
let _ = self.tun.try_send(&buf);
|
||||
self.state.store(State::Idle);
|
||||
}
|
||||
}
|
||||
|
||||
/// Attempt to receive a datagram from the other end.
|
||||
///
|
||||
/// This method takes `&self`, and it can be called safely by multiple threads
|
||||
/// at the same time.
|
||||
///
|
||||
/// A return of `None` means the TCP connection is broken
|
||||
/// and this socket must be closed.
|
||||
pub async fn recv(&self, buf: &mut BytesMut) -> Option<usize> {
|
||||
tracing::trace!(
|
||||
"Socket recv called, local_addr: {:?}, remote_addr: {:?}",
|
||||
self.local_addr,
|
||||
self.remote_addr
|
||||
);
|
||||
loop {
|
||||
match self.state.load() {
|
||||
State::Established => {
|
||||
let Ok(raw_buf) = self.incoming.recv_async().await else {
|
||||
info!("Connection {} recv error", self);
|
||||
return None;
|
||||
};
|
||||
|
||||
let Some((src_mac, dst_mac, _v4_packet, tcp_packet)) =
|
||||
parse_ip_packet(&raw_buf)
|
||||
else {
|
||||
trace!("Dropping malformed fake tcp packet for established socket");
|
||||
continue;
|
||||
};
|
||||
|
||||
tracing::trace!(
|
||||
"Socket received TCP packet from {}({:?}) to {}({:?}): {:?}",
|
||||
self.remote_addr,
|
||||
src_mac,
|
||||
self.local_addr,
|
||||
dst_mac,
|
||||
tcp_packet
|
||||
);
|
||||
|
||||
self.remote_mac.store(Some(src_mac));
|
||||
|
||||
if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 {
|
||||
info!("Connection {} reset by peer", self);
|
||||
return None;
|
||||
}
|
||||
|
||||
if (tcp_packet.get_flags() & tcp::TcpFlags::ACK) != 0
|
||||
&& tcp_packet.payload().is_empty()
|
||||
{
|
||||
self.seq
|
||||
.store(tcp_packet.get_acknowledgement(), Ordering::Relaxed);
|
||||
}
|
||||
|
||||
let payload = tcp_packet.payload();
|
||||
|
||||
let new_ack = tcp_packet.get_sequence().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 {
|
||||
// SACK 选项类型为 5
|
||||
let payload = opt.payload();
|
||||
for chunk in payload.chunks(8) {
|
||||
if chunk.len() != 8 {
|
||||
continue;
|
||||
}
|
||||
let left = tcp_packet.get_acknowledgement();
|
||||
let right = u32::from_be_bytes(chunk[0..4].try_into().unwrap());
|
||||
let len = right.wrapping_sub(left);
|
||||
|
||||
let sack_end = u32::from_be_bytes(chunk[4..8].try_into().unwrap());
|
||||
if len == 0 || sack_end <= left {
|
||||
continue;
|
||||
}
|
||||
|
||||
let send_len = std::cmp::min(len, 1400) as usize;
|
||||
let data = vec![0u8; send_len];
|
||||
|
||||
let buf = build_tcp_packet(
|
||||
self.local_mac,
|
||||
self.remote_mac.load().unwrap_or(MacAddr::zero()),
|
||||
self.local_addr,
|
||||
self.remote_addr,
|
||||
left,
|
||||
self.ack.load(Ordering::Relaxed),
|
||||
tcp::TcpFlags::ACK,
|
||||
Some(&data),
|
||||
);
|
||||
|
||||
if let Err(e) = self.tun.try_send(&buf) {
|
||||
tracing::error!("Failed to send SACK response: {}", e);
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if payload.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
buf.extend_from_slice(payload);
|
||||
|
||||
return Some(payload.len());
|
||||
}
|
||||
State::SynSent => {
|
||||
let Ok(Ok(buf)) = time::timeout(TIMEOUT, self.incoming.recv_async()).await
|
||||
else {
|
||||
info!("Waiting for client SYN + ACK timed out");
|
||||
return None;
|
||||
};
|
||||
let Some((src_mac, _dst_mac, _v4_packet, tcp_packet)) = parse_ip_packet(&buf)
|
||||
else {
|
||||
trace!("Dropping malformed fake tcp packet during handshake");
|
||||
continue;
|
||||
};
|
||||
|
||||
if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 {
|
||||
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 {
|
||||
// found our SYN + ACK
|
||||
self.seq
|
||||
.store(tcp_packet.get_acknowledgement(), Ordering::Relaxed);
|
||||
self.ack
|
||||
.store(tcp_packet.get_sequence() + 1, Ordering::Relaxed);
|
||||
self.remote_mac.store(Some(src_mac));
|
||||
self.state.store(State::Established);
|
||||
return Some(0);
|
||||
}
|
||||
}
|
||||
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn local_addr(&self) -> SocketAddr {
|
||||
self.local_addr
|
||||
}
|
||||
|
||||
pub fn remote_addr(&self) -> SocketAddr {
|
||||
self.remote_addr
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Socket {
|
||||
/// Drop the socket and close the TCP connection
|
||||
fn drop(&mut self) {
|
||||
let tuple = AddrTuple::new(self.local_addr, self.remote_addr);
|
||||
// dissociates ourself from the dispatch map
|
||||
let (removed, closed) = {
|
||||
let mut state = self.shared.state.write().unwrap();
|
||||
(state.tuples.remove(&tuple).is_some(), state.closed)
|
||||
};
|
||||
if !removed {
|
||||
if closed {
|
||||
trace!(?tuple, "Fake TCP tuple already removed after stack closed");
|
||||
} else {
|
||||
warn!(?tuple, "Fake TCP tuple missing while dropping socket");
|
||||
}
|
||||
}
|
||||
// purge cache
|
||||
let _ = self.shared.tuples_purge.send(tuple);
|
||||
|
||||
let buf = build_tcp_packet(
|
||||
self.local_mac,
|
||||
self.remote_mac.load().unwrap_or(MacAddr::zero()),
|
||||
self.local_addr,
|
||||
self.remote_addr,
|
||||
self.seq.load(Ordering::Relaxed),
|
||||
0,
|
||||
tcp::TcpFlags::RST,
|
||||
None,
|
||||
);
|
||||
if let Err(e) = self.tun.try_send(&buf) {
|
||||
warn!("Unable to send RST to remote end: {}", e);
|
||||
}
|
||||
|
||||
info!("Fake TCP connection to {} closed", self);
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for Socket {
|
||||
/// User-friendly string representation of the socket
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(
|
||||
f,
|
||||
"(Fake TCP connection from {} to {})",
|
||||
self.local_addr, self.remote_addr
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// A userspace TCP state machine
|
||||
impl Stack {
|
||||
/// Create a new stack, `tun` is an array of [`Tun`](tokio_tun::Tun).
|
||||
/// When more than one [`Tun`](tokio_tun::Tun) object is passed in, same amount
|
||||
/// of reader will be spawned later. This allows user to utilize the performance
|
||||
/// benefit of Multiqueue Tun support on machines with SMP.
|
||||
pub fn new(
|
||||
tun: Arc<dyn Tun>,
|
||||
local_ip: Ipv4Addr,
|
||||
local_ip6: Option<Ipv6Addr>,
|
||||
local_mac: Option<MacAddr>,
|
||||
) -> Stack {
|
||||
let (tuples_purge_tx, _tuples_purge_rx) = broadcast::channel(16);
|
||||
let shared = Arc::new(Shared {
|
||||
state: RwLock::new(StackState::default()),
|
||||
tun: tun.clone(),
|
||||
listening: RwLock::new(HashSet::new()),
|
||||
tuples_purge: tuples_purge_tx.clone(),
|
||||
});
|
||||
|
||||
let t = tokio::spawn(Stack::reader_task(
|
||||
tun,
|
||||
shared.clone(),
|
||||
tuples_purge_tx.subscribe(),
|
||||
));
|
||||
|
||||
Stack {
|
||||
shared,
|
||||
local_ip,
|
||||
local_ip6,
|
||||
local_mac: local_mac.unwrap_or(MacAddr::zero()),
|
||||
reader_task: AbortOnDropHandle::new(t),
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns the driver type of the stack.
|
||||
pub fn driver_type(&self) -> &'static str {
|
||||
self.shared.tun.driver_type()
|
||||
}
|
||||
|
||||
pub fn is_closed(&self) -> bool {
|
||||
self.shared.is_closed() || self.reader_task.is_finished()
|
||||
}
|
||||
|
||||
/// Listens for incoming connections on the given `port`.
|
||||
pub fn listen(&mut self, port: u16) {
|
||||
assert!(self.shared.listening.write().unwrap().insert(port));
|
||||
}
|
||||
|
||||
pub fn try_alloc_established_socket(
|
||||
&self,
|
||||
local_addr: SocketAddr,
|
||||
remote_addr: SocketAddr,
|
||||
state: State,
|
||||
) -> Option<Socket> {
|
||||
let tuple = AddrTuple::new(local_addr, remote_addr);
|
||||
let mut stack_state = self.shared.state.write().unwrap();
|
||||
if stack_state.closed || self.reader_task.is_finished() {
|
||||
stack_state.closed = true;
|
||||
warn!(
|
||||
?tuple,
|
||||
"fake_tcp stack is closed, refusing to allocate socket"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
let (sock, incoming) = Socket::new(
|
||||
self.shared.clone(),
|
||||
// self.shared.tun.choose(&mut rng).unwrap().clone(),
|
||||
self.shared.tun.clone(), // Simplification: just use the first tun
|
||||
local_addr,
|
||||
remote_addr,
|
||||
self.local_mac,
|
||||
None,
|
||||
Some(0), // Initial ACK
|
||||
state,
|
||||
);
|
||||
assert!(stack_state.tuples.insert(tuple, incoming).is_none());
|
||||
Some(sock)
|
||||
}
|
||||
|
||||
async fn reader_task(
|
||||
tun: Arc<dyn Tun>,
|
||||
shared: Arc<Shared>,
|
||||
mut tuples_purge: broadcast::Receiver<AddrTuple>,
|
||||
) {
|
||||
let mut tuples: HashMap<AddrTuple, flume::Sender<Bytes>> = HashMap::new();
|
||||
|
||||
loop {
|
||||
let mut buf = BytesMut::new();
|
||||
|
||||
tokio::select! {
|
||||
size = tun.recv(&mut buf) => {
|
||||
let size = match size {
|
||||
Ok(size) => size,
|
||||
Err(e) => {
|
||||
let shared_tuple_count = shared.mark_closed_and_clear_tuples();
|
||||
let cached_tuple_count = tuples.len();
|
||||
tuples.clear();
|
||||
error!(
|
||||
?e,
|
||||
driver_type = tun.driver_type(),
|
||||
shared_tuple_count,
|
||||
cached_tuple_count,
|
||||
"fake_tcp tun recv failed, reader_task exiting"
|
||||
);
|
||||
break;
|
||||
}
|
||||
};
|
||||
tracing::trace!(len = size, ?buf, "PnetTun received packet");
|
||||
let buf = buf.split().freeze();
|
||||
|
||||
match parse_ip_packet(&buf) {
|
||||
Some((_src_mac, _dst_mac, ip_packet, tcp_packet)) => {
|
||||
let local_addr = SocketAddr::new(
|
||||
ip_packet.get_destination(),
|
||||
tcp_packet.get_destination(),
|
||||
);
|
||||
let remote_addr = SocketAddr::new(
|
||||
ip_packet.get_source(),
|
||||
tcp_packet.get_source(),
|
||||
);
|
||||
|
||||
let tuple = AddrTuple::new(local_addr, remote_addr);
|
||||
if let Some(c) = tuples.get(&tuple) {
|
||||
if c.send_async(buf).await.is_err() {
|
||||
trace!("Cache hit, but receiver already closed, dropping packet");
|
||||
}
|
||||
|
||||
continue;
|
||||
|
||||
// If not Ok, receiver has been closed and just fall through to the slow
|
||||
// path below
|
||||
} else {
|
||||
trace!("Cache miss, checking the shared tuples table for connection");
|
||||
let sender = {
|
||||
let state = shared.state.read().unwrap();
|
||||
state.tuples.get(&tuple).cloned()
|
||||
};
|
||||
|
||||
if let Some(c) = sender {
|
||||
trace!("Storing connection information into local tuples");
|
||||
tuples.insert(tuple, c.clone());
|
||||
if let Err(e) = c.send_async(buf).await {
|
||||
trace!("Error sending packet to connection: {:?}", e);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
if tcp_packet.get_flags() == tcp::TcpFlags::SYN
|
||||
&& shared
|
||||
.listening
|
||||
.read()
|
||||
.unwrap()
|
||||
.contains(&tcp_packet.get_destination())
|
||||
{
|
||||
trace!(?tcp_packet, "Received SYN packet for port {}, ignoring", tcp_packet.get_destination());
|
||||
continue;
|
||||
} else if (tcp_packet.get_flags() & tcp::TcpFlags::RST) != 0 {
|
||||
info!("Unknown RST TCP packet from {}, ignoring", remote_addr);
|
||||
continue;
|
||||
} else {
|
||||
trace!("Unknown TCP packet from {}, ignoring", remote_addr);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
None => {
|
||||
trace!("Dropping packet with no IP/TCP header");
|
||||
continue;
|
||||
}
|
||||
}
|
||||
},
|
||||
tuple = tuples_purge.recv() => {
|
||||
match tuple {
|
||||
Ok(tuple) => {
|
||||
tuples.remove(&tuple);
|
||||
trace!("Removed cached tuple: {:?}", tuple);
|
||||
}
|
||||
Err(broadcast::error::RecvError::Lagged(skipped)) => {
|
||||
let cached_tuple_count = tuples.len();
|
||||
tuples.clear();
|
||||
warn!(
|
||||
skipped,
|
||||
cached_tuple_count,
|
||||
"fake_tcp tuples purge receiver lagged, cleared local cache"
|
||||
);
|
||||
}
|
||||
Err(broadcast::error::RecvError::Closed) => {
|
||||
let shared_tuple_count = shared.mark_closed_and_clear_tuples();
|
||||
let cached_tuple_count = tuples.len();
|
||||
tuples.clear();
|
||||
warn!(
|
||||
shared_tuple_count,
|
||||
cached_tuple_count,
|
||||
"fake_tcp tuples purge channel closed, reader_task exiting"
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::io;
|
||||
use tokio::{
|
||||
sync::Notify,
|
||||
time::{Duration, timeout},
|
||||
};
|
||||
|
||||
#[derive(Default)]
|
||||
struct FailingTun {
|
||||
fail: Notify,
|
||||
}
|
||||
|
||||
impl FailingTun {
|
||||
fn fail(&self) {
|
||||
self.fail.notify_one();
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Tun for FailingTun {
|
||||
async fn recv(&self, _packet: &mut BytesMut) -> Result<usize, io::Error> {
|
||||
self.fail.notified().await;
|
||||
Err(io::Error::new(io::ErrorKind::BrokenPipe, "test tun closed"))
|
||||
}
|
||||
|
||||
fn try_send(&self, _packet: &Bytes) -> Result<(), io::Error> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn driver_type(&self) -> &'static str {
|
||||
"test"
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reader_task_closes_sockets_on_tun_recv_error() {
|
||||
let tun = Arc::new(FailingTun::default());
|
||||
let mut stack = Stack::new(tun.clone(), Ipv4Addr::LOCALHOST, None, None);
|
||||
let socket = stack
|
||||
.try_alloc_established_socket(
|
||||
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 10_000),
|
||||
SocketAddr::new(Ipv4Addr::new(192, 0, 2, 1).into(), 20_000),
|
||||
State::Established,
|
||||
)
|
||||
.expect("socket allocation should succeed before tun failure");
|
||||
|
||||
tun.fail();
|
||||
|
||||
let join_result = timeout(Duration::from_secs(1), &mut stack.reader_task)
|
||||
.await
|
||||
.expect("reader task should exit after tun recv error");
|
||||
assert!(join_result.is_ok());
|
||||
assert!(stack.is_closed());
|
||||
|
||||
let mut buf = BytesMut::new();
|
||||
let recv_result = timeout(Duration::from_secs(1), socket.recv(&mut buf))
|
||||
.await
|
||||
.expect("socket recv should not hang after reader task exits");
|
||||
assert_eq!(recv_result, None);
|
||||
|
||||
let new_socket = stack.try_alloc_established_socket(
|
||||
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 10_001),
|
||||
SocketAddr::new(Ipv4Addr::new(192, 0, 2, 1).into(), 20_001),
|
||||
State::Established,
|
||||
);
|
||||
assert!(new_socket.is_none());
|
||||
|
||||
drop(socket);
|
||||
}
|
||||
}
|
||||
@@ -1,373 +0,0 @@
|
||||
use std::{
|
||||
sync::Arc,
|
||||
task::{Context, Poll},
|
||||
};
|
||||
|
||||
use auto_impl::auto_impl;
|
||||
use futures::{Sink, SinkExt, Stream, StreamExt};
|
||||
|
||||
use crate::proto::common::TunnelInfo;
|
||||
|
||||
use self::stats::Throughput;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[auto_impl(Arc, Box)]
|
||||
pub trait TunnelFilter: Send + Sync {
|
||||
type FilterOutput;
|
||||
|
||||
fn before_send(&self, data: SinkItem) -> Option<SinkItem> {
|
||||
Some(data)
|
||||
}
|
||||
|
||||
fn after_received(&self, data: StreamItem) -> Option<StreamItem> {
|
||||
match data {
|
||||
Ok(v) => Some(Ok(v)),
|
||||
Err(e) => Some(Err(e)),
|
||||
}
|
||||
}
|
||||
|
||||
fn filter_output(&self) -> Self::FilterOutput;
|
||||
}
|
||||
|
||||
pub struct TunnelFilterChain<A, B> {
|
||||
a: A,
|
||||
b: B,
|
||||
}
|
||||
|
||||
impl<A, B, OA, OB> TunnelFilter for TunnelFilterChain<A, B>
|
||||
where
|
||||
A: TunnelFilter<FilterOutput = OA>,
|
||||
B: TunnelFilter<FilterOutput = OB>,
|
||||
{
|
||||
type FilterOutput = (OA, OB);
|
||||
fn before_send(&self, data: SinkItem) -> Option<SinkItem> {
|
||||
let data = self.a.before_send(data)?;
|
||||
self.b.before_send(data)
|
||||
}
|
||||
fn after_received(&self, data: StreamItem) -> Option<StreamItem> {
|
||||
let data = self.b.after_received(data)?;
|
||||
self.a.after_received(data)
|
||||
}
|
||||
fn filter_output(&self) -> Self::FilterOutput {
|
||||
(self.a.filter_output(), self.b.filter_output())
|
||||
}
|
||||
}
|
||||
|
||||
impl<A, B> TunnelFilterChain<A, B> {
|
||||
pub fn new(a: A, b: B) -> Self {
|
||||
Self { a, b }
|
||||
}
|
||||
|
||||
pub fn chain<T: TunnelFilter>(self, c: T) -> TunnelFilterChain<Self, T> {
|
||||
TunnelFilterChain::new(self, c)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct EmptyFilter;
|
||||
impl TunnelFilter for EmptyFilter {
|
||||
type FilterOutput = ();
|
||||
fn filter_output(&self) {}
|
||||
}
|
||||
|
||||
pub trait ToTunnelChain {
|
||||
fn to_chain(self) -> TunnelFilterChain<EmptyFilter, Self>
|
||||
where
|
||||
Self: Sized,
|
||||
{
|
||||
TunnelFilterChain::new(EmptyFilter, self)
|
||||
}
|
||||
}
|
||||
|
||||
impl<O, T: TunnelFilter<FilterOutput = O>> ToTunnelChain for T {}
|
||||
|
||||
pub struct TunnelWithFilter<T, F> {
|
||||
inner: T,
|
||||
filter: Arc<F>,
|
||||
}
|
||||
|
||||
impl<T, F> TunnelWithFilter<T, F>
|
||||
where
|
||||
T: Tunnel + Send + 'static,
|
||||
F: TunnelFilter + Send + 'static,
|
||||
{
|
||||
pub fn new(inner: T, filter: F) -> Self {
|
||||
Self {
|
||||
inner,
|
||||
filter: Arc::new(filter),
|
||||
}
|
||||
}
|
||||
|
||||
fn wrap_sink<S: ZCPacketSink + Unpin + 'static>(filter: Arc<F>, sink: S) -> impl ZCPacketSink {
|
||||
struct SinkWrapper<F, S> {
|
||||
sink: S,
|
||||
filter: Arc<F>,
|
||||
}
|
||||
|
||||
impl<F, S> Sink<ZCPacket> for SinkWrapper<F, S>
|
||||
where
|
||||
F: TunnelFilter + 'static,
|
||||
S: ZCPacketSink + 'static + Unpin,
|
||||
{
|
||||
type Error = SinkError;
|
||||
|
||||
fn poll_ready(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
self.get_mut().sink.poll_ready_unpin(cx)
|
||||
}
|
||||
|
||||
fn start_send(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
item: ZCPacket,
|
||||
) -> Result<(), Self::Error> {
|
||||
let Some(item) = self.filter.before_send(item) else {
|
||||
return Ok(());
|
||||
};
|
||||
self.get_mut().sink.start_send_unpin(item)
|
||||
}
|
||||
|
||||
fn poll_flush(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
self.get_mut().sink.poll_flush_unpin(cx)
|
||||
}
|
||||
|
||||
fn poll_close(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
self.get_mut().sink.poll_close_unpin(cx)
|
||||
}
|
||||
}
|
||||
|
||||
SinkWrapper { sink, filter }
|
||||
}
|
||||
|
||||
fn wrap_stream<S: ZCPacketStream + Unpin + 'static>(
|
||||
filter: Arc<F>,
|
||||
stream: S,
|
||||
) -> impl ZCPacketStream {
|
||||
struct StreamWrapper<F, S> {
|
||||
stream: S,
|
||||
filter: Arc<F>,
|
||||
}
|
||||
|
||||
impl<F, S> Stream for StreamWrapper<F, S>
|
||||
where
|
||||
F: TunnelFilter + 'static,
|
||||
S: ZCPacketStream + 'static + Unpin,
|
||||
{
|
||||
type Item = StreamItem;
|
||||
|
||||
fn poll_next(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Option<Self::Item>> {
|
||||
let self_mut = self.get_mut();
|
||||
loop {
|
||||
match self_mut.stream.poll_next_unpin(cx) {
|
||||
Poll::Ready(Some(ret)) => {
|
||||
let Some(ret) = self_mut.filter.after_received(ret) else {
|
||||
continue;
|
||||
};
|
||||
return Poll::Ready(Some(ret));
|
||||
}
|
||||
Poll::Ready(None) => {
|
||||
return Poll::Ready(None);
|
||||
}
|
||||
Poll::Pending => {
|
||||
return Poll::Pending;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
StreamWrapper { stream, filter }
|
||||
}
|
||||
}
|
||||
|
||||
impl<T, F> Tunnel for TunnelWithFilter<T, F>
|
||||
where
|
||||
T: Tunnel + Send + 'static,
|
||||
F: TunnelFilter + Send + 'static,
|
||||
{
|
||||
fn info(&self) -> Option<TunnelInfo> {
|
||||
self.inner.info()
|
||||
}
|
||||
|
||||
fn split(&self) -> (Pin<Box<dyn ZCPacketStream>>, Pin<Box<dyn ZCPacketSink>>) {
|
||||
let (stream, sink) = self.inner.split();
|
||||
let filter = self.filter.clone();
|
||||
(
|
||||
Box::pin(Self::wrap_stream(filter.clone(), stream)),
|
||||
Box::pin(Self::wrap_sink(filter, sink)),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PacketRecorderTunnelFilter {
|
||||
pub received: Arc<std::sync::Mutex<Vec<ZCPacket>>>,
|
||||
pub sent: Arc<std::sync::Mutex<Vec<ZCPacket>>>,
|
||||
}
|
||||
|
||||
impl TunnelFilter for PacketRecorderTunnelFilter {
|
||||
type FilterOutput = (Vec<ZCPacket>, Vec<ZCPacket>);
|
||||
|
||||
fn before_send(&self, data: SinkItem) -> Option<SinkItem> {
|
||||
self.sent.lock().unwrap().push(data.clone());
|
||||
Some(data)
|
||||
}
|
||||
|
||||
fn after_received(&self, data: StreamItem) -> Option<StreamItem> {
|
||||
match data {
|
||||
Ok(v) => {
|
||||
self.received.lock().unwrap().push(v.clone());
|
||||
Some(Ok(v))
|
||||
}
|
||||
Err(e) => Some(Err(e)),
|
||||
}
|
||||
}
|
||||
|
||||
fn filter_output(&self) -> Self::FilterOutput {
|
||||
(
|
||||
self.sent.lock().unwrap().clone(),
|
||||
self.received.lock().unwrap().clone(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for PacketRecorderTunnelFilter {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl PacketRecorderTunnelFilter {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
received: Arc::new(std::sync::Mutex::new(Vec::new())),
|
||||
sent: Arc::new(std::sync::Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct StatsRecorderTunnelFilter {
|
||||
throughput: Arc<Throughput>,
|
||||
}
|
||||
|
||||
impl TunnelFilter for StatsRecorderTunnelFilter {
|
||||
type FilterOutput = Arc<Throughput>;
|
||||
|
||||
fn before_send(&self, data: SinkItem) -> Option<SinkItem> {
|
||||
self.throughput.record_tx_bytes(data.buf_len() as u64);
|
||||
Some(data)
|
||||
}
|
||||
|
||||
fn after_received(&self, data: StreamItem) -> Option<StreamItem> {
|
||||
match data {
|
||||
Ok(v) => {
|
||||
self.throughput.record_rx_bytes(v.buf_len() as u64);
|
||||
Some(Ok(v))
|
||||
}
|
||||
Err(e) => Some(Err(e)),
|
||||
}
|
||||
}
|
||||
|
||||
fn filter_output(&self) -> Self::FilterOutput {
|
||||
self.throughput.clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for StatsRecorderTunnelFilter {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl StatsRecorderTunnelFilter {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
throughput: Arc::new(Throughput::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_throughput(&self) -> Arc<Throughput> {
|
||||
self.throughput.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub mod tests {
|
||||
use std::sync::atomic::{AtomicU32, Ordering};
|
||||
|
||||
use filter::ring::create_ring_tunnel_pair;
|
||||
|
||||
use super::*;
|
||||
|
||||
pub struct DropSendTunnelFilter {
|
||||
start: AtomicU32,
|
||||
end: AtomicU32,
|
||||
cur: AtomicU32,
|
||||
}
|
||||
|
||||
impl TunnelFilter for DropSendTunnelFilter {
|
||||
type FilterOutput = ();
|
||||
|
||||
fn before_send(&self, data: SinkItem) -> Option<SinkItem> {
|
||||
self.cur.fetch_add(1, Ordering::SeqCst);
|
||||
if self.cur.load(Ordering::SeqCst) >= self.start.load(Ordering::SeqCst)
|
||||
&& self.cur.load(std::sync::atomic::Ordering::SeqCst)
|
||||
< self.end.load(Ordering::SeqCst)
|
||||
{
|
||||
tracing::trace!("drop packet: {:?}", data);
|
||||
return None;
|
||||
}
|
||||
Some(data)
|
||||
}
|
||||
|
||||
fn filter_output(&self) {}
|
||||
}
|
||||
|
||||
impl DropSendTunnelFilter {
|
||||
pub fn new(start: u32, end: u32) -> Self {
|
||||
Self {
|
||||
start: AtomicU32::new(start),
|
||||
end: AtomicU32::new(end),
|
||||
cur: AtomicU32::new(0),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_nested_filter() {
|
||||
let filter = Arc::new(
|
||||
PacketRecorderTunnelFilter::new()
|
||||
.to_chain()
|
||||
.chain(PacketRecorderTunnelFilter::new())
|
||||
.chain(PacketRecorderTunnelFilter::new())
|
||||
.chain(PacketRecorderTunnelFilter::new()),
|
||||
);
|
||||
let (s, _b) = create_ring_tunnel_pair();
|
||||
let tunnel = TunnelWithFilter::new(s, filter.clone());
|
||||
|
||||
let (_r, mut s) = tunnel.split();
|
||||
s.send(ZCPacket::new_with_payload("ab".as_bytes()))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let out = filter.filter_output();
|
||||
|
||||
let a = out.0.0.0.1;
|
||||
let b = out.0.0.1;
|
||||
let c = out.0.1;
|
||||
let _d = out.1;
|
||||
|
||||
assert_eq!(1, a.0.len());
|
||||
assert_eq!(1, b.0.len());
|
||||
assert_eq!(1, c.0.len());
|
||||
}
|
||||
}
|
||||
@@ -1,86 +0,0 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime};
|
||||
|
||||
/// Dummy certificate verifier that treats any certificate as valid.
|
||||
/// NOTE, such verification is vulnerable to MITM attacks, but convenient for testing.
|
||||
#[derive(Debug)]
|
||||
struct SkipServerVerification(Arc<rustls::crypto::CryptoProvider>);
|
||||
|
||||
impl SkipServerVerification {
|
||||
fn new(provider: Arc<rustls::crypto::CryptoProvider>) -> Arc<Self> {
|
||||
Arc::new(Self(provider))
|
||||
}
|
||||
}
|
||||
|
||||
impl rustls::client::danger::ServerCertVerifier for SkipServerVerification {
|
||||
fn verify_server_cert(
|
||||
&self,
|
||||
_end_entity: &CertificateDer<'_>,
|
||||
_intermediates: &[CertificateDer<'_>],
|
||||
_server_name: &ServerName<'_>,
|
||||
_ocsp: &[u8],
|
||||
_now: UnixTime,
|
||||
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
|
||||
Ok(rustls::client::danger::ServerCertVerified::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls12_signature(
|
||||
&self,
|
||||
message: &[u8],
|
||||
cert: &CertificateDer<'_>,
|
||||
dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
|
||||
rustls::crypto::verify_tls12_signature(
|
||||
message,
|
||||
cert,
|
||||
dss,
|
||||
&self.0.signature_verification_algorithms,
|
||||
)
|
||||
}
|
||||
|
||||
fn verify_tls13_signature(
|
||||
&self,
|
||||
message: &[u8],
|
||||
cert: &CertificateDer<'_>,
|
||||
dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
|
||||
rustls::crypto::verify_tls13_signature(
|
||||
message,
|
||||
cert,
|
||||
dss,
|
||||
&self.0.signature_verification_algorithms,
|
||||
)
|
||||
}
|
||||
|
||||
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
|
||||
self.0.signature_verification_algorithms.supported_schemes()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn init_crypto_provider() {
|
||||
let _ =
|
||||
rustls::crypto::CryptoProvider::install_default(rustls::crypto::ring::default_provider());
|
||||
}
|
||||
|
||||
pub fn get_insecure_tls_client_config() -> rustls::ClientConfig {
|
||||
init_crypto_provider();
|
||||
let provider = rustls::crypto::CryptoProvider::get_default().unwrap();
|
||||
let mut config = rustls::ClientConfig::builder()
|
||||
.dangerous()
|
||||
.with_custom_certificate_verifier(SkipServerVerification::new(provider.clone()))
|
||||
.with_no_client_auth();
|
||||
config.enable_sni = true;
|
||||
config.enable_early_data = false;
|
||||
config
|
||||
}
|
||||
|
||||
pub fn get_insecure_tls_cert<'a>() -> (Vec<CertificateDer<'a>>, PrivateKeyDer<'a>) {
|
||||
let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
|
||||
let cert_der = cert.serialize_der().unwrap();
|
||||
let priv_key = cert.serialize_private_key_der();
|
||||
let priv_key = rustls::pki_types::PrivatePkcs8KeyDer::from(priv_key);
|
||||
let cert_chain = vec![cert_der.into()];
|
||||
|
||||
(cert_chain, priv_key.into())
|
||||
}
|
||||
+17
-255
@@ -1,34 +1,15 @@
|
||||
use std::{
|
||||
collections::hash_map::DefaultHasher, hash::Hasher, net::SocketAddr, pin::Pin, sync::Arc,
|
||||
};
|
||||
use std::{collections::hash_map::DefaultHasher, hash::Hasher, net::SocketAddr};
|
||||
|
||||
use crate::{
|
||||
common::{dns::socket_addrs, error::Error},
|
||||
proto::common::TunnelInfo,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))]
|
||||
use crate::common::dns::socket_addrs;
|
||||
use crate::common::error::Error;
|
||||
use derive_more::{From, TryInto};
|
||||
use futures::{Sink, Stream};
|
||||
use socket2::Protocol;
|
||||
use std::fmt::Debug;
|
||||
#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))]
|
||||
use easytier_core::tunnel::{IpVersion, TunnelError};
|
||||
use strum::{Display, EnumString, IntoStaticStr, VariantArray};
|
||||
use tokio::time::error::Elapsed;
|
||||
|
||||
use self::packet_def::ZCPacket;
|
||||
|
||||
pub mod buf;
|
||||
pub mod common;
|
||||
pub mod filter;
|
||||
pub mod mpsc;
|
||||
pub mod packet_def;
|
||||
pub mod ring;
|
||||
pub mod stats;
|
||||
pub mod tcp;
|
||||
pub mod udp;
|
||||
pub(crate) mod udp_src;
|
||||
|
||||
#[cfg(feature = "faketcp")]
|
||||
pub mod fake_tcp;
|
||||
pub(crate) mod protocol;
|
||||
|
||||
#[cfg(feature = "wireguard")]
|
||||
pub mod wireguard;
|
||||
@@ -39,116 +20,6 @@ pub mod quic;
|
||||
#[cfg(feature = "websocket")]
|
||||
pub mod websocket;
|
||||
|
||||
#[cfg(any(feature = "quic", feature = "websocket"))]
|
||||
pub mod insecure_tls;
|
||||
|
||||
#[cfg(unix)]
|
||||
pub mod unix;
|
||||
|
||||
#[derive(thiserror::Error, Debug)]
|
||||
pub enum TunnelError {
|
||||
#[error("io error: {0}")]
|
||||
IOError(#[from] std::io::Error),
|
||||
#[error("invalid packet. msg: {0}")]
|
||||
InvalidPacket(String),
|
||||
#[error("exceed max packet size. max: {0}, input: {1}")]
|
||||
ExceedMaxPacketSize(usize, usize),
|
||||
|
||||
#[error("invalid protocol: {0}")]
|
||||
InvalidProtocol(String),
|
||||
#[error("invalid addr: {0}")]
|
||||
InvalidAddr(String),
|
||||
|
||||
#[error("internal error {0}")]
|
||||
InternalError(String),
|
||||
|
||||
#[error("conn id not match, expect: {0}, actual: {1}")]
|
||||
ConnIdNotMatch(u32, u32),
|
||||
#[error("buffer full")]
|
||||
BufferFull,
|
||||
|
||||
#[error("timeout")]
|
||||
Timeout(#[from] Elapsed),
|
||||
|
||||
#[error("anyhow error: {0}")]
|
||||
Anyhow(#[from] anyhow::Error),
|
||||
|
||||
#[error("shutdown")]
|
||||
Shutdown,
|
||||
|
||||
#[error("no dns record found")]
|
||||
NoDnsRecordFound(IpVersion),
|
||||
|
||||
#[cfg(feature = "websocket")]
|
||||
#[error("websocket error: {0}")]
|
||||
WebSocketError(#[from] tokio_websockets::Error),
|
||||
|
||||
#[error("tunnel error: {0}")]
|
||||
TunError(String),
|
||||
}
|
||||
|
||||
pub type StreamT = packet_def::ZCPacket;
|
||||
pub type StreamItem = Result<StreamT, TunnelError>;
|
||||
pub type SinkItem = packet_def::ZCPacket;
|
||||
pub type SinkError = TunnelError;
|
||||
|
||||
pub trait ZCPacketStream: Stream<Item = StreamItem> + Send {}
|
||||
impl<T> ZCPacketStream for T where T: Stream<Item = StreamItem> + Send {}
|
||||
pub trait ZCPacketSink: Sink<SinkItem, Error = SinkError> + Send {}
|
||||
impl<T> ZCPacketSink for T where T: Sink<SinkItem, Error = SinkError> + Send {}
|
||||
|
||||
pub type SplitTunnel = (Pin<Box<dyn ZCPacketStream>>, Pin<Box<dyn ZCPacketSink>>);
|
||||
|
||||
#[auto_impl::auto_impl(Box, Arc)]
|
||||
pub trait Tunnel: Send {
|
||||
fn split(&self) -> SplitTunnel;
|
||||
fn info(&self) -> Option<TunnelInfo>;
|
||||
}
|
||||
|
||||
#[auto_impl::auto_impl(Arc)]
|
||||
pub trait TunnelConnCounter: 'static + Send + Sync + Debug {
|
||||
fn get(&self) -> Option<u32>;
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub enum IpVersion {
|
||||
V4,
|
||||
V6,
|
||||
Both,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
#[auto_impl::auto_impl(Box)]
|
||||
pub trait TunnelListener: Send {
|
||||
async fn listen(&mut self) -> Result<(), TunnelError>;
|
||||
async fn accept(&mut self) -> Result<Box<dyn Tunnel>, TunnelError>;
|
||||
fn local_url(&self) -> url::Url;
|
||||
fn get_conn_counter(&self) -> Arc<Box<dyn TunnelConnCounter>> {
|
||||
#[derive(Debug)]
|
||||
struct FakeTunnelConnCounter {}
|
||||
impl TunnelConnCounter for FakeTunnelConnCounter {
|
||||
fn get(&self) -> Option<u32> {
|
||||
None
|
||||
}
|
||||
}
|
||||
Arc::new(Box::new(FakeTunnelConnCounter {}))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
#[auto_impl::auto_impl(Box, &mut)]
|
||||
pub trait TunnelConnector: Send {
|
||||
async fn connect(&mut self) -> Result<Box<dyn Tunnel>, TunnelError>;
|
||||
fn remote_url(&self) -> url::Url;
|
||||
fn set_bind_addrs(&mut self, _addrs: Vec<SocketAddr>) {}
|
||||
fn set_ip_version(&mut self, _ip_version: IpVersion) {}
|
||||
fn set_resolved_addr(&mut self, _addr: SocketAddr) {}
|
||||
/// Linux SO_MARK to apply to outbound sockets. `None` leaves SO_MARK
|
||||
/// untouched; `Some(mark)` applies that exact value (including `Some(0)`).
|
||||
/// Default impl is a no-op; IP-based connectors override.
|
||||
fn set_socket_mark(&mut self, _socket_mark: Option<u32>) {}
|
||||
}
|
||||
|
||||
pub fn build_url_from_socket_addr(addr: &String, scheme: &str) -> url::Url {
|
||||
if let Ok(sock_addr) = addr.parse::<SocketAddr>() {
|
||||
let url_str = format!("{}://0.0.0.0", scheme);
|
||||
@@ -162,31 +33,8 @@ pub fn build_url_from_socket_addr(addr: &String, scheme: &str) -> url::Url {
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for dyn Tunnel {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("Tunnel")
|
||||
.field("info", &self.info())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for dyn TunnelConnector {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("TunnelConnector")
|
||||
.field("remote_url", &self.remote_url())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for dyn TunnelListener {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("TunnelListener")
|
||||
.field("local_url", &self.local_url())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))]
|
||||
pub(crate) trait FromUrl {
|
||||
async fn from_url(url: url::Url, ip_version: IpVersion) -> Result<Self, TunnelError>
|
||||
where
|
||||
@@ -194,6 +42,7 @@ pub(crate) trait FromUrl {
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[cfg(any(feature = "faketcp", feature = "websocket", feature = "wireguard"))]
|
||||
impl FromUrl for SocketAddr {
|
||||
async fn from_url(url: url::Url, ip_version: IpVersion) -> Result<Self, TunnelError> {
|
||||
let addrs = socket_addrs(&url, || {
|
||||
@@ -229,15 +78,6 @@ impl FromUrl for SocketAddr {
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl FromUrl for uuid::Uuid {
|
||||
async fn from_url(url: url::Url, _ip_version: IpVersion) -> Result<Self, TunnelError> {
|
||||
let o = url.host_str().unwrap();
|
||||
let o = uuid::Uuid::parse_str(o).map_err(|e| TunnelError::InvalidAddr(e.to_string()))?;
|
||||
Ok(o)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct TunnelUrl {
|
||||
inner: url::Url,
|
||||
}
|
||||
@@ -284,12 +124,6 @@ pub fn generate_digest_from_str(str1: &str, str2: &str, digest: &mut [u8]) {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct IpSchemeAttributes {
|
||||
protocol: Protocol,
|
||||
port_offset: u16,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Display, EnumString, IntoStaticStr, VariantArray)]
|
||||
#[strum(serialize_all = "lowercase")]
|
||||
pub enum IpScheme {
|
||||
@@ -308,42 +142,16 @@ pub enum IpScheme {
|
||||
}
|
||||
|
||||
impl IpScheme {
|
||||
const fn attributes(self) -> IpSchemeAttributes {
|
||||
let (protocol, port_offset) = match self {
|
||||
Self::Tcp => (Protocol::TCP, 0),
|
||||
Self::Udp => (Protocol::UDP, 0),
|
||||
#[cfg(feature = "wireguard")]
|
||||
Self::Wg => (Protocol::UDP, 1),
|
||||
#[cfg(feature = "quic")]
|
||||
Self::Quic => (Protocol::UDP, 2),
|
||||
#[cfg(feature = "websocket")]
|
||||
Self::Ws => (Protocol::TCP, 1),
|
||||
#[cfg(feature = "websocket")]
|
||||
Self::Wss => (Protocol::TCP, 2),
|
||||
#[cfg(feature = "faketcp")]
|
||||
Self::FakeTcp => (Protocol::TCP, 3),
|
||||
};
|
||||
IpSchemeAttributes {
|
||||
protocol,
|
||||
port_offset,
|
||||
}
|
||||
}
|
||||
pub const fn protocol(self) -> Protocol {
|
||||
self.attributes().protocol
|
||||
pub fn port_offset(self) -> u16 {
|
||||
let scheme: &'static str = self.into();
|
||||
easytier_core::connectivity::protocol::protocol_port_offset(scheme)
|
||||
.expect("IpScheme must have core protocol metadata")
|
||||
}
|
||||
|
||||
pub const fn port_offset(self) -> u16 {
|
||||
self.attributes().port_offset
|
||||
}
|
||||
|
||||
pub const fn default_port(self) -> u16 {
|
||||
match self {
|
||||
#[cfg(feature = "websocket")]
|
||||
Self::Ws => 80,
|
||||
#[cfg(feature = "websocket")]
|
||||
Self::Wss => 443,
|
||||
_ => 11010 + self.port_offset(),
|
||||
}
|
||||
pub fn default_port(self) -> u16 {
|
||||
let scheme: &'static str = self.into();
|
||||
easytier_core::connectivity::protocol::protocol_default_port(scheme)
|
||||
.expect("IpScheme must have core protocol metadata")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -376,49 +184,3 @@ impl TryFrom<&url::Url> for TunnelScheme {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn get_scheme_by_url(l: &url::Url) -> Result<TunnelScheme, Error> {
|
||||
l.try_into()
|
||||
}
|
||||
|
||||
macro_rules! __matches_scheme__ {
|
||||
($url:expr, $( $pattern:pat_param )|+ ) => {
|
||||
matches!($crate::tunnel::get_scheme_by_url(&$url), Ok($( $pattern )|+))
|
||||
};
|
||||
}
|
||||
|
||||
pub(crate) use __matches_scheme__ as matches_scheme;
|
||||
|
||||
pub fn get_protocol_by_url(l: &url::Url) -> Result<Protocol, Error> {
|
||||
let TunnelScheme::Ip(scheme) = l.try_into()? else {
|
||||
return Err(Error::InvalidUrl(l.to_string()));
|
||||
};
|
||||
Ok(scheme.protocol())
|
||||
}
|
||||
|
||||
macro_rules! __matches_protocol__ {
|
||||
($url:expr, $( $pattern:pat_param )|+ ) => {
|
||||
matches!($crate::tunnel::get_protocol_by_url($url), Ok($( $pattern )|+))
|
||||
};
|
||||
}
|
||||
|
||||
pub(crate) use __matches_protocol__ as matches_protocol;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{IpScheme, TunnelScheme, matches_scheme};
|
||||
|
||||
#[test]
|
||||
fn matches_scheme_accepts_owned_url() {
|
||||
let url: url::Url = "udp://[2001:db8::1]:11010".parse().unwrap();
|
||||
|
||||
assert!(matches_scheme!(url, TunnelScheme::Ip(IpScheme::Udp)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn matches_scheme_accepts_borrowed_url() {
|
||||
let url: url::Url = "udp://[2001:db8::1]:11010".parse().unwrap();
|
||||
|
||||
assert!(matches_scheme!(&url, TunnelScheme::Ip(IpScheme::Udp)));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,256 +0,0 @@
|
||||
// this mod wrap tunnel to a mpsc tunnel, based on crossbeam_channel
|
||||
|
||||
use std::{pin::Pin, time::Duration};
|
||||
|
||||
use anyhow::Context;
|
||||
use tokio::time::timeout;
|
||||
|
||||
use crate::proto::common::TunnelInfo;
|
||||
|
||||
use super::{Tunnel, TunnelError, ZCPacketSink, ZCPacketStream, packet_def::ZCPacket};
|
||||
|
||||
use tokio::sync::mpsc::{Receiver, Sender, channel, error::TrySendError};
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
// use tachyonix::{channel, Receiver, Sender, TrySendError};
|
||||
|
||||
use futures::SinkExt;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct MpscTunnelSender(Sender<ZCPacket>);
|
||||
|
||||
impl MpscTunnelSender {
|
||||
pub async fn send(&self, item: ZCPacket) -> Result<(), TunnelError> {
|
||||
self.0.send(item).await.with_context(|| "send error")?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn try_send(&self, item: ZCPacket) -> Result<(), TunnelError> {
|
||||
self.0.try_send(item).map_err(|e| match e {
|
||||
TrySendError::Full(_) => TunnelError::BufferFull,
|
||||
TrySendError::Closed(_) => TunnelError::Shutdown,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub struct MpscTunnel<T> {
|
||||
tx: Option<Sender<ZCPacket>>,
|
||||
|
||||
tunnel: T,
|
||||
stream: Option<Pin<Box<dyn ZCPacketStream>>>,
|
||||
|
||||
task: AbortOnDropHandle<()>,
|
||||
}
|
||||
|
||||
impl<T: Tunnel> MpscTunnel<T> {
|
||||
pub fn new(tunnel: T, send_timeout: Option<Duration>) -> Self {
|
||||
let (tx, mut rx) = channel(32);
|
||||
let (stream, mut sink) = tunnel.split();
|
||||
|
||||
let task = tokio::spawn(async move {
|
||||
loop {
|
||||
if let Err(e) = Self::forward_one_round(&mut rx, &mut sink, send_timeout).await {
|
||||
tracing::error!(?e, "forward error");
|
||||
break;
|
||||
}
|
||||
}
|
||||
rx.close();
|
||||
let close_ret = timeout(Duration::from_secs(5), sink.close()).await;
|
||||
tracing::warn!(?close_ret, "mpsc close sink");
|
||||
});
|
||||
|
||||
Self {
|
||||
tx: Some(tx),
|
||||
tunnel,
|
||||
stream: Some(stream),
|
||||
task: AbortOnDropHandle::new(task),
|
||||
}
|
||||
}
|
||||
|
||||
async fn forward_one_round(
|
||||
rx: &mut Receiver<ZCPacket>,
|
||||
sink: &mut Pin<Box<dyn ZCPacketSink>>,
|
||||
send_timeout_ms: Option<Duration>,
|
||||
) -> Result<(), TunnelError> {
|
||||
let item = rx.recv().await.with_context(|| "recv error")?;
|
||||
if let Some(timeout_ms) = send_timeout_ms {
|
||||
Self::forward_one_round_with_timeout(rx, sink, item, timeout_ms).await
|
||||
} else {
|
||||
Self::forward_one_round_no_timeout(rx, sink, item).await
|
||||
}
|
||||
}
|
||||
|
||||
async fn forward_one_round_no_timeout(
|
||||
rx: &mut Receiver<ZCPacket>,
|
||||
sink: &mut Pin<Box<dyn ZCPacketSink>>,
|
||||
initial_item: ZCPacket,
|
||||
) -> Result<(), TunnelError> {
|
||||
sink.feed(initial_item).await?;
|
||||
|
||||
while let Ok(item) = rx.try_recv() {
|
||||
if let Err(e) = sink.feed(item).await {
|
||||
tracing::error!(?e, "feed error");
|
||||
return Err(e);
|
||||
}
|
||||
}
|
||||
|
||||
sink.flush().await
|
||||
}
|
||||
|
||||
async fn forward_one_round_with_timeout(
|
||||
rx: &mut Receiver<ZCPacket>,
|
||||
sink: &mut Pin<Box<dyn ZCPacketSink>>,
|
||||
initial_item: ZCPacket,
|
||||
timeout_ms: Duration,
|
||||
) -> Result<(), TunnelError> {
|
||||
match timeout(timeout_ms, async move {
|
||||
Self::forward_one_round_no_timeout(rx, sink, initial_item).await
|
||||
})
|
||||
.await
|
||||
{
|
||||
Ok(Ok(_)) => Ok(()),
|
||||
Ok(Err(e)) => {
|
||||
tracing::error!(?e, "forward error");
|
||||
Err(e)
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(?e, "forward timeout");
|
||||
Err(e.into())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_stream(&mut self) -> Pin<Box<dyn ZCPacketStream>> {
|
||||
self.stream.take().unwrap()
|
||||
}
|
||||
|
||||
pub fn get_sink(&self) -> MpscTunnelSender {
|
||||
MpscTunnelSender(self.tx.as_ref().unwrap().clone())
|
||||
}
|
||||
|
||||
pub fn close(&mut self) {
|
||||
self.tx.take();
|
||||
self.task.abort();
|
||||
}
|
||||
|
||||
pub fn tunnel_info(&self) -> Option<TunnelInfo> {
|
||||
self.tunnel.info()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use futures::StreamExt;
|
||||
|
||||
use crate::tunnel::{
|
||||
TunnelConnector, TunnelListener,
|
||||
ring::{RING_TUNNEL_CAP, create_ring_tunnel_pair},
|
||||
tcp::{TcpTunnelConnector, TcpTunnelListener},
|
||||
};
|
||||
|
||||
use super::*;
|
||||
// test slow send lock in framed tunnel
|
||||
#[tokio::test]
|
||||
async fn mpsc_slow_receiver() {
|
||||
let mut listener = TcpTunnelListener::new("tcp://127.0.0.1:11014".parse().unwrap());
|
||||
let mut connector = TcpTunnelConnector::new("tcp://127.0.0.1:11014".parse().unwrap());
|
||||
|
||||
listener.listen().await.unwrap();
|
||||
let t1 = tokio::spawn(async move {
|
||||
let t = listener.accept().await.unwrap();
|
||||
let (mut stream, _sink) = t.split();
|
||||
let now = tokio::time::Instant::now();
|
||||
|
||||
let mut a_counter = 0;
|
||||
let mut b_counter = 0;
|
||||
|
||||
while let Some(Ok(msg)) = stream.next().await {
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
if now.elapsed().as_secs() > 5 {
|
||||
break;
|
||||
}
|
||||
|
||||
if msg.payload() == "hello".as_bytes() {
|
||||
a_counter += 1;
|
||||
} else if msg.payload() == "hello2".as_bytes() {
|
||||
b_counter += 1;
|
||||
}
|
||||
}
|
||||
|
||||
tracing::info!("t1 exit");
|
||||
assert_ne!(a_counter, 0);
|
||||
assert_ne!(b_counter, 0);
|
||||
});
|
||||
|
||||
let tunnel = connector.connect().await.unwrap();
|
||||
let mpsc_tunnel = MpscTunnel::new(tunnel, None);
|
||||
|
||||
let sink1 = mpsc_tunnel.get_sink();
|
||||
let t2 = tokio::spawn(async move {
|
||||
for i in 0..1000000 {
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
|
||||
let a = sink1
|
||||
.send(ZCPacket::new_with_payload("hello".as_bytes()))
|
||||
.await;
|
||||
if a.is_err() {
|
||||
tracing::info!(?a, "t2 exit with err");
|
||||
break;
|
||||
}
|
||||
|
||||
if i % 5000 == 0 {
|
||||
tracing::info!(i, "send2 1000");
|
||||
}
|
||||
}
|
||||
|
||||
tracing::info!("t2 exit");
|
||||
});
|
||||
|
||||
let sink2 = mpsc_tunnel.get_sink();
|
||||
let t3 = tokio::spawn(async move {
|
||||
for i in 0..1000000 {
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
let a = sink2
|
||||
.send(ZCPacket::new_with_payload("hello2".as_bytes()))
|
||||
.await;
|
||||
if a.is_err() {
|
||||
tracing::info!(?a, "t3 exit with err");
|
||||
break;
|
||||
}
|
||||
|
||||
if i % 5000 == 0 {
|
||||
tracing::info!(i, "send2 1000");
|
||||
}
|
||||
}
|
||||
|
||||
tracing::info!("t3 exit");
|
||||
});
|
||||
|
||||
let t4 = tokio::spawn(async move {
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(5)).await;
|
||||
tracing::info!("closing");
|
||||
drop(mpsc_tunnel);
|
||||
tracing::info!("closed");
|
||||
});
|
||||
|
||||
let _ = tokio::join!(t1, t2, t3, t4);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mpsc_slow_receiver_with_send_timeout() {
|
||||
let (a, _b) = create_ring_tunnel_pair();
|
||||
let mpsc_tunnel = MpscTunnel::new(a, Some(Duration::from_secs(1)));
|
||||
let s = mpsc_tunnel.get_sink();
|
||||
for _ in 0..RING_TUNNEL_CAP {
|
||||
s.send(ZCPacket::new_with_payload(&[0; 1024]))
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(1500)).await;
|
||||
let e = s.send(ZCPacket::new_with_payload(&[0; 1024])).await;
|
||||
assert!(e.is_ok());
|
||||
|
||||
tokio::time::sleep(Duration::from_millis(1500)).await;
|
||||
|
||||
let e = s.send(ZCPacket::new_with_payload(&[0; 1024])).await;
|
||||
assert!(e.is_err());
|
||||
}
|
||||
}
|
||||
@@ -1,820 +0,0 @@
|
||||
use bytes::Buf;
|
||||
use bytes::Bytes;
|
||||
use bytes::BytesMut;
|
||||
use zerocopy::AsBytes;
|
||||
use zerocopy::FromBytes;
|
||||
use zerocopy::FromZeroes;
|
||||
use zerocopy::byteorder::*;
|
||||
|
||||
type DefaultEndian = LittleEndian;
|
||||
|
||||
const fn max(a: usize, b: usize) -> usize {
|
||||
[a, b][(a < b) as usize]
|
||||
}
|
||||
|
||||
// TCP TunnelHeader
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct TCPTunnelHeader {
|
||||
pub len: U32<DefaultEndian>,
|
||||
}
|
||||
pub const TCP_TUNNEL_HEADER_SIZE: usize = std::mem::size_of::<TCPTunnelHeader>();
|
||||
|
||||
#[derive(AsBytes, FromZeroes, Clone, Debug)]
|
||||
#[repr(u8)]
|
||||
pub enum UdpPacketType {
|
||||
Invalid = 0,
|
||||
Syn = 1,
|
||||
Sack = 2,
|
||||
Data = 3,
|
||||
Fin = 4,
|
||||
HolePunch = 5,
|
||||
V4HolePunch = 6, // when receiving v4 hole punch packet, the packet contains a socket addr of other peer, we
|
||||
// will send a hole punch packet to that peer. we only accept this packet from loopback interface.
|
||||
V6HolePunch = 7, // when receiving v6 hole punch packet, the packet contains a socket addr of other peer, we
|
||||
// will send a hole punch packet to that peer. we only accept this packet from lookback interface.
|
||||
}
|
||||
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct V4HolePunchPacket {
|
||||
pub dst_ipv4: [u8; 4],
|
||||
pub dst_port: U16<DefaultEndian>,
|
||||
}
|
||||
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct V6HolePunchPacket {
|
||||
pub dst_ipv6: [u8; 16],
|
||||
pub dst_port: U16<DefaultEndian>,
|
||||
pub preferred_src_ipv6: [u8; 16],
|
||||
pub preferred_src_ifindex: U32<DefaultEndian>,
|
||||
}
|
||||
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct UDPTunnelHeader {
|
||||
pub conn_id: U32<DefaultEndian>,
|
||||
pub msg_type: u8,
|
||||
pub padding: u8,
|
||||
pub len: U16<DefaultEndian>,
|
||||
}
|
||||
pub const UDP_TUNNEL_HEADER_SIZE: usize = std::mem::size_of::<UDPTunnelHeader>();
|
||||
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct WGTunnelHeader {
|
||||
pub ipv4_header: [u8; 20],
|
||||
}
|
||||
pub const WG_TUNNEL_HEADER_SIZE: usize = std::mem::size_of::<WGTunnelHeader>();
|
||||
|
||||
#[derive(AsBytes, FromZeroes, Copy, Clone, Debug)]
|
||||
#[repr(u8)]
|
||||
pub enum PacketType {
|
||||
Invalid = 0,
|
||||
Data = 1,
|
||||
HandShake = 2,
|
||||
RoutePacket = 3, // deprecated
|
||||
Ping = 4,
|
||||
Pong = 5,
|
||||
TaRpc = 6, // deprecated
|
||||
Route = 7, // deprecated
|
||||
RpcReq = 8,
|
||||
RpcResp = 9,
|
||||
ForeignNetworkPacket = 10,
|
||||
KcpSrc = 11,
|
||||
KcpDst = 12,
|
||||
QuicSrc = 16,
|
||||
QuicDst = 17,
|
||||
NoiseHandshakeMsg1 = 13,
|
||||
NoiseHandshakeMsg2 = 14,
|
||||
NoiseHandshakeMsg3 = 15,
|
||||
RelayHandshake = 20,
|
||||
RelayHandshakeAck = 21,
|
||||
|
||||
// used internally,
|
||||
DataWithKcpSrcModified = 18,
|
||||
DataWithQuicSrcModified = 19,
|
||||
}
|
||||
|
||||
bitflags::bitflags! {
|
||||
struct PeerManagerHeaderFlags: u8 {
|
||||
const ENCRYPTED = 0b0000_0001;
|
||||
const LATENCY_FIRST = 0b0000_0010;
|
||||
const EXIT_NODE = 0b0000_0100;
|
||||
const NO_PROXY = 0b0000_1000;
|
||||
const COMPRESSED = 0b0001_0000;
|
||||
// deprecated flags, can be reused.
|
||||
// const KCP_SRC_MODIFIED = 0b0010_0000;
|
||||
// const QUIC_SRC_MODIFIED = 0b1000_0000;
|
||||
const NOT_SEND_TO_TUN = 0b0100_0000;
|
||||
|
||||
const _ = !0;
|
||||
}
|
||||
}
|
||||
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct PeerManagerHeader {
|
||||
pub from_peer_id: U32<DefaultEndian>,
|
||||
pub to_peer_id: U32<DefaultEndian>,
|
||||
pub packet_type: u8,
|
||||
pub flags: u8,
|
||||
pub forward_counter: u8,
|
||||
reserved: u8,
|
||||
pub len: U32<DefaultEndian>,
|
||||
}
|
||||
pub const PEER_MANAGER_HEADER_SIZE: usize = std::mem::size_of::<PeerManagerHeader>();
|
||||
|
||||
impl PeerManagerHeader {
|
||||
pub fn is_encrypted(&self) -> bool {
|
||||
PeerManagerHeaderFlags::from_bits(self.flags)
|
||||
.unwrap()
|
||||
.contains(PeerManagerHeaderFlags::ENCRYPTED)
|
||||
}
|
||||
|
||||
pub fn set_encrypted(&mut self, encrypted: bool) {
|
||||
let mut flags = PeerManagerHeaderFlags::from_bits(self.flags).unwrap();
|
||||
if encrypted {
|
||||
flags.insert(PeerManagerHeaderFlags::ENCRYPTED);
|
||||
} else {
|
||||
flags.remove(PeerManagerHeaderFlags::ENCRYPTED);
|
||||
}
|
||||
self.flags = flags.bits();
|
||||
}
|
||||
|
||||
pub fn is_latency_first(&self) -> bool {
|
||||
PeerManagerHeaderFlags::from_bits(self.flags)
|
||||
.unwrap()
|
||||
.contains(PeerManagerHeaderFlags::LATENCY_FIRST)
|
||||
}
|
||||
|
||||
pub fn is_exit_node(&self) -> bool {
|
||||
PeerManagerHeaderFlags::from_bits(self.flags)
|
||||
.unwrap()
|
||||
.contains(PeerManagerHeaderFlags::EXIT_NODE)
|
||||
}
|
||||
|
||||
pub fn is_no_proxy(&self) -> bool {
|
||||
PeerManagerHeaderFlags::from_bits(self.flags)
|
||||
.unwrap()
|
||||
.contains(PeerManagerHeaderFlags::NO_PROXY)
|
||||
}
|
||||
|
||||
pub fn is_compressed(&self) -> bool {
|
||||
PeerManagerHeaderFlags::from_bits(self.flags)
|
||||
.unwrap()
|
||||
.contains(PeerManagerHeaderFlags::COMPRESSED)
|
||||
}
|
||||
|
||||
pub fn set_latency_first(&mut self, latency_first: bool) -> &mut Self {
|
||||
let mut flags = PeerManagerHeaderFlags::from_bits(self.flags).unwrap();
|
||||
if latency_first {
|
||||
flags.insert(PeerManagerHeaderFlags::LATENCY_FIRST);
|
||||
} else {
|
||||
flags.remove(PeerManagerHeaderFlags::LATENCY_FIRST);
|
||||
}
|
||||
self.flags = flags.bits();
|
||||
self
|
||||
}
|
||||
|
||||
pub fn set_exit_node(&mut self, exit_node: bool) -> &mut Self {
|
||||
let mut flags = PeerManagerHeaderFlags::from_bits(self.flags).unwrap();
|
||||
if exit_node {
|
||||
flags.insert(PeerManagerHeaderFlags::EXIT_NODE);
|
||||
} else {
|
||||
flags.remove(PeerManagerHeaderFlags::EXIT_NODE);
|
||||
}
|
||||
self.flags = flags.bits();
|
||||
self
|
||||
}
|
||||
|
||||
pub fn set_no_proxy(&mut self, no_proxy: bool) -> &mut Self {
|
||||
let mut flags = PeerManagerHeaderFlags::from_bits(self.flags).unwrap();
|
||||
if no_proxy {
|
||||
flags.insert(PeerManagerHeaderFlags::NO_PROXY);
|
||||
} else {
|
||||
flags.remove(PeerManagerHeaderFlags::NO_PROXY);
|
||||
}
|
||||
self.flags = flags.bits();
|
||||
self
|
||||
}
|
||||
|
||||
pub fn set_compressed(&mut self, compressed: bool) -> &mut Self {
|
||||
let mut flags = PeerManagerHeaderFlags::from_bits(self.flags).unwrap();
|
||||
if compressed {
|
||||
flags.insert(PeerManagerHeaderFlags::COMPRESSED);
|
||||
} else {
|
||||
flags.remove(PeerManagerHeaderFlags::COMPRESSED);
|
||||
}
|
||||
self.flags = flags.bits();
|
||||
self
|
||||
}
|
||||
|
||||
pub fn mark_kcp_src_modified(&mut self) -> &mut Self {
|
||||
assert_eq!(self.packet_type, PacketType::Data as u8);
|
||||
self.packet_type = PacketType::DataWithKcpSrcModified as u8;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn is_kcp_src_modified(&self) -> bool {
|
||||
self.packet_type == PacketType::DataWithKcpSrcModified as u8
|
||||
}
|
||||
|
||||
pub fn mark_quic_src_modified(&mut self) -> &mut Self {
|
||||
assert_eq!(self.packet_type, PacketType::Data as u8);
|
||||
self.packet_type = PacketType::DataWithQuicSrcModified as u8;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn is_quic_src_modified(&self) -> bool {
|
||||
self.packet_type == PacketType::DataWithQuicSrcModified as u8
|
||||
}
|
||||
|
||||
pub fn set_not_send_to_tun(&mut self, not_send_to_tun: bool) -> &mut Self {
|
||||
let mut flags = PeerManagerHeaderFlags::from_bits(self.flags).unwrap();
|
||||
if not_send_to_tun {
|
||||
flags.insert(PeerManagerHeaderFlags::NOT_SEND_TO_TUN);
|
||||
} else {
|
||||
flags.remove(PeerManagerHeaderFlags::NOT_SEND_TO_TUN);
|
||||
}
|
||||
self.flags = flags.bits();
|
||||
self
|
||||
}
|
||||
|
||||
pub fn is_not_send_to_tun(&self) -> bool {
|
||||
PeerManagerHeaderFlags::from_bits(self.flags)
|
||||
.unwrap()
|
||||
.contains(PeerManagerHeaderFlags::NOT_SEND_TO_TUN)
|
||||
}
|
||||
}
|
||||
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct ForeignNetworkPacketHeader {
|
||||
pub header_len: U16<DefaultEndian>,
|
||||
pub dst_peer_id: U32<DefaultEndian>,
|
||||
pub network_name_offset: U16<DefaultEndian>,
|
||||
pub network_name_len: U16<DefaultEndian>,
|
||||
/* variable length network_name string */
|
||||
}
|
||||
|
||||
impl ForeignNetworkPacketHeader {
|
||||
pub fn new(dst_peer_id: u32, network_name: &str) -> Self {
|
||||
let network_name_offset = std::mem::size_of::<ForeignNetworkPacketHeader>() as u16;
|
||||
let network_name_len = network_name.len() as u16;
|
||||
let header_len = network_name_offset + network_name_len;
|
||||
Self {
|
||||
header_len: U16::new(header_len),
|
||||
dst_peer_id: U32::new(dst_peer_id),
|
||||
network_name_offset: U16::new(network_name_offset),
|
||||
network_name_len: U16::new(network_name_len),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_network_name(&self, zc_packet_payload: &[u8]) -> String {
|
||||
let offset = self.network_name_offset.get() as usize;
|
||||
let len = self.network_name_len.get() as usize;
|
||||
std::str::from_utf8(&zc_packet_payload[offset..offset + len])
|
||||
.unwrap()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
pub fn get_dst_peer_id(&self) -> u32 {
|
||||
self.dst_peer_id.get()
|
||||
}
|
||||
|
||||
pub fn get_header_len(&self) -> usize {
|
||||
self.header_len.get() as usize
|
||||
}
|
||||
}
|
||||
|
||||
// reserve space for AEAD authentication tag and nonce
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug)]
|
||||
pub struct AeadTail<const TAG_SIZE: usize, const NONCE_SIZE: usize> {
|
||||
pub tag: [u8; TAG_SIZE],
|
||||
pub nonce: [u8; NONCE_SIZE],
|
||||
}
|
||||
|
||||
impl<const TAG_SIZE: usize, const NONCE_SIZE: usize> AeadTail<TAG_SIZE, NONCE_SIZE> {
|
||||
pub const TAG_SIZE: usize = TAG_SIZE;
|
||||
pub const NONCE_SIZE: usize = NONCE_SIZE;
|
||||
|
||||
pub const SIZE: usize = std::mem::size_of::<Self>();
|
||||
}
|
||||
|
||||
pub type StandardAeadTail = AeadTail<16, 12>;
|
||||
|
||||
#[derive(AsBytes, FromZeroes, Clone, Debug, Copy, PartialEq, Hash, Eq)]
|
||||
#[repr(u8)]
|
||||
pub enum CompressorAlgo {
|
||||
None = 0,
|
||||
#[cfg(feature = "zstd")]
|
||||
ZstdDefault = 1,
|
||||
}
|
||||
|
||||
#[repr(C, packed)]
|
||||
#[derive(AsBytes, FromBytes, FromZeroes, Clone, Debug, Default)]
|
||||
pub struct CompressorTail {
|
||||
pub algo: u8,
|
||||
}
|
||||
pub const COMPRESSOR_TAIL_SIZE: usize = std::mem::size_of::<CompressorTail>();
|
||||
|
||||
impl CompressorTail {
|
||||
pub fn get_algo(&self) -> Option<CompressorAlgo> {
|
||||
match self.algo {
|
||||
#[cfg(feature = "zstd")]
|
||||
1 => Some(CompressorAlgo::ZstdDefault),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new(algo: CompressorAlgo) -> Self {
|
||||
Self { algo: algo as u8 }
|
||||
}
|
||||
}
|
||||
|
||||
pub const TAIL_RESERVED_SIZE: usize = max(StandardAeadTail::SIZE, COMPRESSOR_TAIL_SIZE);
|
||||
|
||||
#[derive(Default, Debug)]
|
||||
pub struct ZCPacketOffsets {
|
||||
pub payload_offset: usize,
|
||||
pub peer_manager_header_offset: usize,
|
||||
pub tcp_tunnel_header_offset: usize,
|
||||
pub udp_tunnel_header_offset: usize,
|
||||
pub wg_tunnel_header_offset: usize,
|
||||
pub dummy_tunnel_header_offset: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub enum ZCPacketType {
|
||||
// received from peer tcp connection
|
||||
TCP,
|
||||
// received from peer udp connection
|
||||
UDP,
|
||||
// received from peer wireguard connection
|
||||
WG,
|
||||
// received from local tun device, should reserve header space for tcp or udp tunnel
|
||||
NIC,
|
||||
// tunnel without header
|
||||
DummyTunnel,
|
||||
}
|
||||
|
||||
const PAYLOAD_OFFSET_FOR_NIC_PACKET: usize = max(
|
||||
max(TCP_TUNNEL_HEADER_SIZE, UDP_TUNNEL_HEADER_SIZE),
|
||||
WG_TUNNEL_HEADER_SIZE,
|
||||
) + PEER_MANAGER_HEADER_SIZE;
|
||||
|
||||
// UDP Tunnel: TUN MTU + 24 (Easy) + 20 (Encrypted) + 8(UDP) + 20(IP) = TUN MTU + 72
|
||||
// TCP Tunnel: TUN MTU + 20 (Easy) + 20 (Encrypted) + 20(TCP) + 20(IP) = TUN MTU + 80
|
||||
|
||||
const INVALID_OFFSET: usize = usize::MAX;
|
||||
|
||||
const fn get_converted_offset(old_hdr_size: usize, new_hdr_size: usize) -> usize {
|
||||
if old_hdr_size < new_hdr_size {
|
||||
INVALID_OFFSET
|
||||
} else {
|
||||
old_hdr_size - new_hdr_size
|
||||
}
|
||||
}
|
||||
|
||||
impl ZCPacketType {
|
||||
pub fn get_packet_offsets(&self) -> ZCPacketOffsets {
|
||||
match self {
|
||||
ZCPacketType::TCP => ZCPacketOffsets {
|
||||
payload_offset: TCP_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE,
|
||||
peer_manager_header_offset: TCP_TUNNEL_HEADER_SIZE,
|
||||
tcp_tunnel_header_offset: 0,
|
||||
udp_tunnel_header_offset: get_converted_offset(
|
||||
TCP_TUNNEL_HEADER_SIZE,
|
||||
UDP_TUNNEL_HEADER_SIZE,
|
||||
),
|
||||
wg_tunnel_header_offset: get_converted_offset(
|
||||
TCP_TUNNEL_HEADER_SIZE,
|
||||
WG_TUNNEL_HEADER_SIZE,
|
||||
),
|
||||
dummy_tunnel_header_offset: get_converted_offset(TCP_TUNNEL_HEADER_SIZE, 0),
|
||||
},
|
||||
ZCPacketType::UDP => ZCPacketOffsets {
|
||||
payload_offset: UDP_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE,
|
||||
peer_manager_header_offset: UDP_TUNNEL_HEADER_SIZE,
|
||||
tcp_tunnel_header_offset: get_converted_offset(
|
||||
UDP_TUNNEL_HEADER_SIZE,
|
||||
TCP_TUNNEL_HEADER_SIZE,
|
||||
),
|
||||
udp_tunnel_header_offset: 0,
|
||||
wg_tunnel_header_offset: get_converted_offset(
|
||||
UDP_TUNNEL_HEADER_SIZE,
|
||||
WG_TUNNEL_HEADER_SIZE,
|
||||
),
|
||||
dummy_tunnel_header_offset: get_converted_offset(UDP_TUNNEL_HEADER_SIZE, 0),
|
||||
},
|
||||
ZCPacketType::WG => ZCPacketOffsets {
|
||||
payload_offset: WG_TUNNEL_HEADER_SIZE + PEER_MANAGER_HEADER_SIZE,
|
||||
peer_manager_header_offset: WG_TUNNEL_HEADER_SIZE,
|
||||
tcp_tunnel_header_offset: get_converted_offset(
|
||||
WG_TUNNEL_HEADER_SIZE,
|
||||
TCP_TUNNEL_HEADER_SIZE,
|
||||
),
|
||||
udp_tunnel_header_offset: get_converted_offset(
|
||||
WG_TUNNEL_HEADER_SIZE,
|
||||
UDP_TUNNEL_HEADER_SIZE,
|
||||
),
|
||||
wg_tunnel_header_offset: 0,
|
||||
dummy_tunnel_header_offset: get_converted_offset(WG_TUNNEL_HEADER_SIZE, 0),
|
||||
},
|
||||
ZCPacketType::NIC => ZCPacketOffsets {
|
||||
payload_offset: PAYLOAD_OFFSET_FOR_NIC_PACKET,
|
||||
peer_manager_header_offset: PAYLOAD_OFFSET_FOR_NIC_PACKET
|
||||
- PEER_MANAGER_HEADER_SIZE,
|
||||
tcp_tunnel_header_offset: PAYLOAD_OFFSET_FOR_NIC_PACKET
|
||||
- PEER_MANAGER_HEADER_SIZE
|
||||
- TCP_TUNNEL_HEADER_SIZE,
|
||||
udp_tunnel_header_offset: PAYLOAD_OFFSET_FOR_NIC_PACKET
|
||||
- PEER_MANAGER_HEADER_SIZE
|
||||
- UDP_TUNNEL_HEADER_SIZE,
|
||||
wg_tunnel_header_offset: PAYLOAD_OFFSET_FOR_NIC_PACKET
|
||||
- PEER_MANAGER_HEADER_SIZE
|
||||
- WG_TUNNEL_HEADER_SIZE,
|
||||
dummy_tunnel_header_offset: PAYLOAD_OFFSET_FOR_NIC_PACKET
|
||||
- PEER_MANAGER_HEADER_SIZE,
|
||||
},
|
||||
ZCPacketType::DummyTunnel => ZCPacketOffsets {
|
||||
payload_offset: PEER_MANAGER_HEADER_SIZE,
|
||||
peer_manager_header_offset: 0,
|
||||
tcp_tunnel_header_offset: get_converted_offset(0, TCP_TUNNEL_HEADER_SIZE),
|
||||
udp_tunnel_header_offset: get_converted_offset(0, UDP_TUNNEL_HEADER_SIZE),
|
||||
wg_tunnel_header_offset: get_converted_offset(0, WG_TUNNEL_HEADER_SIZE),
|
||||
dummy_tunnel_header_offset: 0,
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ZCPacket {
|
||||
inner: BytesMut,
|
||||
packet_type: ZCPacketType,
|
||||
}
|
||||
|
||||
impl ZCPacket {
|
||||
fn bytes_from_offset(&self, offset: usize) -> Option<&[u8]> {
|
||||
self.inner.get(offset..)
|
||||
}
|
||||
|
||||
fn mut_bytes_from_offset(&mut self, offset: usize) -> Option<&mut [u8]> {
|
||||
self.inner.get_mut(offset..)
|
||||
}
|
||||
|
||||
pub fn new_nic_packet() -> Self {
|
||||
Self {
|
||||
inner: BytesMut::new(),
|
||||
packet_type: ZCPacketType::NIC,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_from_buf(buf: BytesMut, packet_type: ZCPacketType) -> Self {
|
||||
Self {
|
||||
inner: buf,
|
||||
packet_type,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_with_payload(payload: &[u8]) -> Self {
|
||||
let mut ret = Self::new_nic_packet();
|
||||
let payload_off = ret.packet_type.get_packet_offsets().payload_offset;
|
||||
let total_len = payload_off + payload.len();
|
||||
ret.inner.reserve(total_len);
|
||||
unsafe { ret.inner.set_len(total_len) };
|
||||
ret.mut_payload().copy_from_slice(payload);
|
||||
ret
|
||||
}
|
||||
|
||||
pub fn new_for_tun(cap: usize, packet_info_len: usize) -> Self {
|
||||
let mut ret = Self::new_nic_packet();
|
||||
ret.inner.reserve(cap);
|
||||
let total_len = ret.packet_type.get_packet_offsets().payload_offset - packet_info_len;
|
||||
unsafe { ret.inner.set_len(total_len) };
|
||||
ret
|
||||
}
|
||||
|
||||
pub fn new_for_foreign_network(
|
||||
network_name: &String,
|
||||
dst_peer_id: u32,
|
||||
foreign_zc_packet: &ZCPacket,
|
||||
) -> Self {
|
||||
let foreign_network_hdr = ForeignNetworkPacketHeader::new(dst_peer_id, network_name);
|
||||
let total_payload_len =
|
||||
foreign_network_hdr.get_header_len() + foreign_zc_packet.tunnel_payload().len();
|
||||
|
||||
let mut ret = Self::new_nic_packet();
|
||||
let payload_off = ret.packet_type.get_packet_offsets().payload_offset;
|
||||
ret.inner.reserve(payload_off + total_payload_len);
|
||||
unsafe { ret.inner.set_len(payload_off + total_payload_len) };
|
||||
|
||||
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_len = foreign_network_hdr.network_name_len.get() as usize;
|
||||
ret.mut_payload()[name_offset..name_offset + name_len]
|
||||
.copy_from_slice(network_name.as_bytes());
|
||||
|
||||
ret.mut_payload()[foreign_network_hdr.get_header_len()..]
|
||||
.copy_from_slice(foreign_zc_packet.tunnel_payload());
|
||||
|
||||
let hdr = ret.mut_peer_manager_header().unwrap();
|
||||
hdr.from_peer_id = 0.into();
|
||||
hdr.to_peer_id = 0.into();
|
||||
hdr.packet_type = PacketType::ForeignNetworkPacket as u8;
|
||||
hdr.len.set(total_payload_len as u32);
|
||||
|
||||
ret
|
||||
}
|
||||
|
||||
pub fn packet_type(&self) -> ZCPacketType {
|
||||
self.packet_type
|
||||
}
|
||||
|
||||
pub fn payload_offset(&self) -> usize {
|
||||
self.packet_type.get_packet_offsets().payload_offset
|
||||
}
|
||||
|
||||
pub fn mut_payload(&mut self) -> &mut [u8] {
|
||||
let offset = self.payload_offset();
|
||||
&mut self.inner[offset..]
|
||||
}
|
||||
|
||||
pub fn mut_peer_manager_header(&mut self) -> Option<&mut PeerManagerHeader> {
|
||||
let offset = self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.peer_manager_header_offset;
|
||||
let bytes = self.mut_bytes_from_offset(offset)?;
|
||||
PeerManagerHeader::mut_from_prefix(bytes)
|
||||
}
|
||||
|
||||
pub fn mut_tcp_tunnel_header(&mut self) -> Option<&mut TCPTunnelHeader> {
|
||||
let offset = self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.tcp_tunnel_header_offset;
|
||||
let bytes = self.mut_bytes_from_offset(offset)?;
|
||||
TCPTunnelHeader::mut_from_prefix(bytes)
|
||||
}
|
||||
|
||||
pub fn mut_udp_tunnel_header(&mut self) -> Option<&mut UDPTunnelHeader> {
|
||||
let offset = self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.udp_tunnel_header_offset;
|
||||
let bytes = self.mut_bytes_from_offset(offset)?;
|
||||
UDPTunnelHeader::mut_from_prefix(bytes)
|
||||
}
|
||||
|
||||
pub fn mut_wg_tunnel_header(&mut self) -> Option<&mut WGTunnelHeader> {
|
||||
let offset = self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.wg_tunnel_header_offset;
|
||||
let bytes = self.mut_bytes_from_offset(offset)?;
|
||||
WGTunnelHeader::mut_from_prefix(bytes)
|
||||
}
|
||||
|
||||
// ref versions
|
||||
pub fn payload(&self) -> &[u8] {
|
||||
&self.inner[self.payload_offset()..]
|
||||
}
|
||||
|
||||
pub fn payload_bytes(mut self) -> BytesMut {
|
||||
self.inner.advance(self.payload_offset());
|
||||
self.inner
|
||||
}
|
||||
|
||||
pub fn peer_manager_header(&self) -> Option<&PeerManagerHeader> {
|
||||
let offset = self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.peer_manager_header_offset;
|
||||
let bytes = self.bytes_from_offset(offset)?;
|
||||
PeerManagerHeader::ref_from_prefix(bytes)
|
||||
}
|
||||
|
||||
pub fn tcp_tunnel_header(&self) -> Option<&TCPTunnelHeader> {
|
||||
let offset = self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.tcp_tunnel_header_offset;
|
||||
let bytes = self.bytes_from_offset(offset)?;
|
||||
TCPTunnelHeader::ref_from_prefix(bytes)
|
||||
}
|
||||
|
||||
pub fn udp_tunnel_header(&self) -> Option<&UDPTunnelHeader> {
|
||||
let offset = self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.udp_tunnel_header_offset;
|
||||
let bytes = self.bytes_from_offset(offset)?;
|
||||
UDPTunnelHeader::ref_from_prefix(bytes)
|
||||
}
|
||||
|
||||
pub fn udp_payload(&self) -> &[u8] {
|
||||
&self.inner[self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.udp_tunnel_header_offset
|
||||
+ UDP_TUNNEL_HEADER_SIZE..]
|
||||
}
|
||||
|
||||
pub fn payload_len(&self) -> usize {
|
||||
self.inner.len() - self.payload_offset()
|
||||
}
|
||||
|
||||
pub fn buf_len(&self) -> usize {
|
||||
self.inner.len()
|
||||
}
|
||||
|
||||
pub fn fill_peer_manager_hdr(&mut self, from_peer_id: u32, to_peer_id: u32, packet_type: u8) {
|
||||
let payload_len = self.payload_len();
|
||||
let hdr = self.mut_peer_manager_header().unwrap();
|
||||
hdr.from_peer_id.set(from_peer_id);
|
||||
hdr.to_peer_id.set(to_peer_id);
|
||||
hdr.packet_type = packet_type;
|
||||
hdr.flags = 0;
|
||||
hdr.forward_counter = 1;
|
||||
hdr.len.set(payload_len as u32);
|
||||
}
|
||||
|
||||
pub fn tunnel_payload(&self) -> &[u8] {
|
||||
&self.inner[self
|
||||
.packet_type
|
||||
.get_packet_offsets()
|
||||
.peer_manager_header_offset..]
|
||||
}
|
||||
|
||||
pub fn tunnel_payload_bytes(mut self) -> BytesMut {
|
||||
self.inner.advance(
|
||||
self.packet_type
|
||||
.get_packet_offsets()
|
||||
.peer_manager_header_offset,
|
||||
);
|
||||
self.inner
|
||||
}
|
||||
|
||||
pub fn convert_type(mut self, target_packet_type: ZCPacketType) -> Self {
|
||||
if target_packet_type == self.packet_type {
|
||||
return self;
|
||||
}
|
||||
|
||||
let new_offset = match target_packet_type {
|
||||
ZCPacketType::TCP => {
|
||||
self.packet_type
|
||||
.get_packet_offsets()
|
||||
.tcp_tunnel_header_offset
|
||||
}
|
||||
ZCPacketType::UDP => {
|
||||
self.packet_type
|
||||
.get_packet_offsets()
|
||||
.udp_tunnel_header_offset
|
||||
}
|
||||
ZCPacketType::WG => {
|
||||
self.packet_type
|
||||
.get_packet_offsets()
|
||||
.wg_tunnel_header_offset
|
||||
}
|
||||
ZCPacketType::DummyTunnel => {
|
||||
self.packet_type
|
||||
.get_packet_offsets()
|
||||
.dummy_tunnel_header_offset
|
||||
}
|
||||
ZCPacketType::NIC => unreachable!(),
|
||||
};
|
||||
|
||||
tracing::trace!(?self.packet_type, ?target_packet_type, ?new_offset, "convert zc packet type");
|
||||
|
||||
if new_offset == INVALID_OFFSET {
|
||||
// copy peer manager header and payload to new buffer
|
||||
let tunnel_payload = self.tunnel_payload();
|
||||
let new_pm_offset = target_packet_type
|
||||
.get_packet_offsets()
|
||||
.peer_manager_header_offset;
|
||||
let mut buf = BytesMut::with_capacity(new_pm_offset + tunnel_payload.len());
|
||||
unsafe { buf.set_len(new_pm_offset) };
|
||||
buf.extend_from_slice(tunnel_payload);
|
||||
return Self::new_from_buf(buf, target_packet_type);
|
||||
}
|
||||
|
||||
self.inner.advance(new_offset);
|
||||
Self::new_from_buf(self.inner, target_packet_type)
|
||||
}
|
||||
|
||||
pub fn into_bytes(self) -> Bytes {
|
||||
self.inner.freeze()
|
||||
}
|
||||
|
||||
pub fn inner(self) -> BytesMut {
|
||||
self.inner
|
||||
}
|
||||
|
||||
pub fn mut_inner(&mut self) -> &mut BytesMut {
|
||||
&mut self.inner
|
||||
}
|
||||
|
||||
pub fn is_lossy(&self) -> bool {
|
||||
self.peer_manager_header()
|
||||
.map(|hdr| hdr.packet_type == PacketType::Data as u8)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
pub fn foreign_network_hdr(&self) -> Option<&ForeignNetworkPacketHeader> {
|
||||
if self.peer_manager_header().unwrap().packet_type == PacketType::ForeignNetworkPacket as u8
|
||||
{
|
||||
ForeignNetworkPacketHeader::ref_from_prefix(self.payload())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
pub fn foreign_network_inner_packet_type(&self) -> Option<u8> {
|
||||
if self.peer_manager_header()?.packet_type != PacketType::ForeignNetworkPacket as u8 {
|
||||
return None;
|
||||
}
|
||||
|
||||
let payload = self.payload();
|
||||
let hdr = ForeignNetworkPacketHeader::ref_from_prefix(payload)?;
|
||||
let inner_packet = payload.get(hdr.get_header_len()..)?;
|
||||
PeerManagerHeader::ref_from_prefix(inner_packet).map(|hdr| hdr.packet_type)
|
||||
}
|
||||
|
||||
pub fn foreign_network_packet(mut self) -> Self {
|
||||
let hdr = self.foreign_network_hdr().unwrap();
|
||||
let foreign_hdr_len = hdr.get_header_len();
|
||||
|
||||
Self::new_from_buf(
|
||||
{
|
||||
self.inner.advance(foreign_hdr_len + self.payload_offset());
|
||||
self.inner
|
||||
},
|
||||
ZCPacketType::DummyTunnel,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn get_src_peer_id(&self) -> Option<u32> {
|
||||
self.peer_manager_header().map(|hdr| hdr.from_peer_id.get())
|
||||
}
|
||||
|
||||
pub fn get_dst_peer_id(&self) -> Option<u32> {
|
||||
self.peer_manager_header().map(|hdr| hdr.to_peer_id.get())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_zc_packet() {
|
||||
let payload = b"hello world";
|
||||
let mut packet = ZCPacket::new_with_payload(payload);
|
||||
let peer_manager_header = packet.mut_peer_manager_header().unwrap();
|
||||
peer_manager_header.packet_type = PacketType::Data as u8;
|
||||
peer_manager_header.len.set(payload.len() as u32);
|
||||
|
||||
let tcp_tunnel_header = packet.mut_tcp_tunnel_header().unwrap();
|
||||
tcp_tunnel_header.len.set(payload.len() as u32);
|
||||
|
||||
// let udp_tunnel_header = packet.mut_udp_tunnel_header().unwrap();
|
||||
// udp_tunnel_header.conn_id = 1;
|
||||
// udp_tunnel_header.msg_type = 2;
|
||||
// udp_tunnel_header.len = payload.len() as u32;
|
||||
|
||||
assert_eq!(packet.payload(), b"hello world");
|
||||
assert_eq!(packet.payload_len(), 11);
|
||||
println!("{:?}", packet.inner);
|
||||
|
||||
let tcp_packet = packet.convert_type(ZCPacketType::TCP).into_bytes();
|
||||
assert_eq!(&tcp_packet[..1], b"\x0b");
|
||||
println!("{:?}", tcp_packet);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_short_tcp_packet_header_access_is_safe() {
|
||||
let mut packet = ZCPacket::new_from_buf(BytesMut::from(&b"\x01"[..]), ZCPacketType::TCP);
|
||||
|
||||
assert!(packet.peer_manager_header().is_none());
|
||||
assert!(packet.tcp_tunnel_header().is_none());
|
||||
assert!(packet.udp_tunnel_header().is_none());
|
||||
assert!(packet.mut_peer_manager_header().is_none());
|
||||
assert!(packet.mut_tcp_tunnel_header().is_none());
|
||||
assert!(packet.mut_udp_tunnel_header().is_none());
|
||||
assert!(packet.mut_wg_tunnel_header().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_invalid_converted_header_offset_is_safe() {
|
||||
let mut packet = ZCPacket::new_from_buf(BytesMut::from(&b"\x01"[..]), ZCPacketType::UDP);
|
||||
|
||||
assert!(packet.mut_wg_tunnel_header().is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,481 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use easytier_core::{
|
||||
connectivity::{
|
||||
protocol::{
|
||||
ClientProtocolUpgrader, CoreClientProtocolConfig, CoreClientProtocolUpgrader,
|
||||
CoreServerProtocolConfig, CoreServerProtocolUpgrader, ServerProtocolAdmission,
|
||||
ServerProtocolUpgrade, ServerProtocolUpgrader,
|
||||
},
|
||||
transport::ConnectedTransport,
|
||||
},
|
||||
socket::udp::UdpSession,
|
||||
tunnel::Tunnel,
|
||||
};
|
||||
|
||||
use crate::{common::global_ctx::ArcGlobalCtx, socket::tcp::RuntimeTcpSocket};
|
||||
|
||||
mod adapters;
|
||||
|
||||
pub(crate) struct RuntimeClientProtocolUpgrader {
|
||||
adapters: Vec<adapters::ClientAdapter>,
|
||||
}
|
||||
|
||||
pub(crate) struct RuntimeServerProtocolUpgrader {
|
||||
adapters: Vec<adapters::ServerAdapter>,
|
||||
}
|
||||
|
||||
fn runtime_client_protocol_adapter(global_ctx: &ArcGlobalCtx) -> RuntimeClientProtocolUpgrader {
|
||||
RuntimeClientProtocolUpgrader {
|
||||
adapters: adapters::client_adapters(global_ctx),
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_server_protocol_adapter(global_ctx: &ArcGlobalCtx) -> RuntimeServerProtocolUpgrader {
|
||||
RuntimeServerProtocolUpgrader {
|
||||
adapters: adapters::server_adapters(global_ctx),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn runtime_client_protocol_upgrader(
|
||||
global_ctx: ArcGlobalCtx,
|
||||
) -> Arc<dyn ClientProtocolUpgrader<RuntimeTcpSocket>> {
|
||||
Arc::new(CoreClientProtocolUpgrader::with_external(
|
||||
CoreClientProtocolConfig {
|
||||
unix: cfg!(unix),
|
||||
faketcp: cfg!(feature = "faketcp"),
|
||||
},
|
||||
Arc::new(runtime_client_protocol_adapter(&global_ctx)),
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) fn runtime_server_protocol_upgrader(
|
||||
global_ctx: ArcGlobalCtx,
|
||||
) -> Arc<dyn ServerProtocolUpgrader<RuntimeTcpSocket>> {
|
||||
Arc::new(CoreServerProtocolUpgrader::with_external(
|
||||
CoreServerProtocolConfig {
|
||||
unix: cfg!(unix),
|
||||
faketcp: cfg!(feature = "faketcp"),
|
||||
},
|
||||
Arc::new(runtime_server_protocol_adapter(&global_ctx)),
|
||||
))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ClientProtocolUpgrader<RuntimeTcpSocket> for RuntimeClientProtocolUpgrader {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
self.adapters
|
||||
.iter()
|
||||
.any(|adapter| adapter.supports_scheme(scheme))
|
||||
}
|
||||
|
||||
fn connect_timeout(&self, scheme: &str) -> Option<std::time::Duration> {
|
||||
self.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports_scheme(scheme))
|
||||
.and_then(|adapter| adapter.connect_timeout(scheme))
|
||||
}
|
||||
|
||||
async fn upgrade_client(
|
||||
&self,
|
||||
connected: ConnectedTransport<RuntimeTcpSocket>,
|
||||
requested_url: url::Url,
|
||||
) -> anyhow::Result<Box<dyn Tunnel>> {
|
||||
let scheme = requested_url.scheme().to_owned();
|
||||
let adapter = self
|
||||
.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports_scheme(&scheme))
|
||||
.ok_or_else(|| anyhow::anyhow!("unsupported client protocol upgrader: {scheme}"))?;
|
||||
adapter.upgrade_client(connected, requested_url).await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ServerProtocolUpgrader<RuntimeTcpSocket> for RuntimeServerProtocolUpgrader {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
self.adapters
|
||||
.iter()
|
||||
.any(|adapter| adapter.supports_scheme(scheme))
|
||||
}
|
||||
|
||||
fn max_pending_tcp_upgrades(&self, scheme: &str) -> Option<std::num::NonZeroUsize> {
|
||||
self.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports_scheme(scheme))
|
||||
.and_then(|adapter| adapter.max_pending_tcp_upgrades(scheme))
|
||||
}
|
||||
|
||||
async fn upgrade_tcp(
|
||||
&self,
|
||||
socket: RuntimeTcpSocket,
|
||||
local_url: url::Url,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
let scheme = local_url.scheme().to_owned();
|
||||
let adapter = self
|
||||
.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports_scheme(&scheme))
|
||||
.ok_or_else(|| {
|
||||
anyhow::anyhow!("unsupported native TCP server protocol upgrader: {scheme}")
|
||||
})?;
|
||||
adapter.upgrade_tcp(socket, local_url).await
|
||||
}
|
||||
|
||||
async fn upgrade_udp(
|
||||
&self,
|
||||
session: UdpSession,
|
||||
local_url: url::Url,
|
||||
admission: Option<ServerProtocolAdmission>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
let scheme = local_url.scheme().to_owned();
|
||||
let adapter = self
|
||||
.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports_scheme(&scheme))
|
||||
.ok_or_else(|| {
|
||||
anyhow::anyhow!("unsupported native UDP server protocol upgrader: {scheme}")
|
||||
})?;
|
||||
adapter.upgrade_udp(session, local_url, admission).await
|
||||
}
|
||||
|
||||
async fn upgrade_byte_stream(
|
||||
&self,
|
||||
socket: RuntimeTcpSocket,
|
||||
local_url: url::Url,
|
||||
remote_url: Option<url::Url>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
let scheme = local_url.scheme().to_owned();
|
||||
let adapter = self
|
||||
.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports_scheme(&scheme))
|
||||
.ok_or_else(|| {
|
||||
anyhow::anyhow!("unsupported native byte-stream server protocol upgrader: {scheme}")
|
||||
})?;
|
||||
adapter
|
||||
.upgrade_byte_stream(socket, local_url, remote_url)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::common::global_ctx::tests::get_mock_global_ctx;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn protocol_capabilities_follow_enabled_features() {
|
||||
let global_ctx = get_mock_global_ctx();
|
||||
let external = runtime_client_protocol_adapter(&global_ctx);
|
||||
|
||||
assert!(!external.supports_scheme("tcp"));
|
||||
assert!(!external.supports_scheme("faketcp"));
|
||||
assert_eq!(external.supports_scheme("ws"), cfg!(feature = "websocket"));
|
||||
assert_eq!(external.supports_scheme("wss"), cfg!(feature = "websocket"));
|
||||
assert_eq!(external.supports_scheme("wg"), cfg!(feature = "wireguard"));
|
||||
assert_eq!(external.supports_scheme("quic"), cfg!(feature = "quic"));
|
||||
|
||||
let upgrader = runtime_client_protocol_upgrader(global_ctx.clone());
|
||||
|
||||
assert!(upgrader.supports_scheme("tcp"));
|
||||
assert!(upgrader.supports_scheme("udp"));
|
||||
assert!(upgrader.supports_scheme("ring"));
|
||||
assert_eq!(upgrader.supports_scheme("unix"), cfg!(unix));
|
||||
assert_eq!(upgrader.supports_scheme("ws"), cfg!(feature = "websocket"));
|
||||
assert_eq!(upgrader.supports_scheme("wss"), cfg!(feature = "websocket"));
|
||||
assert_eq!(upgrader.supports_scheme("wg"), cfg!(feature = "wireguard"));
|
||||
assert_eq!(upgrader.supports_scheme("quic"), cfg!(feature = "quic"));
|
||||
assert_eq!(
|
||||
upgrader.supports_scheme("faketcp"),
|
||||
cfg!(feature = "faketcp")
|
||||
);
|
||||
|
||||
let server_external = runtime_server_protocol_adapter(&global_ctx);
|
||||
assert!(!server_external.supports_scheme("tcp"));
|
||||
assert!(!server_external.supports_scheme("udp"));
|
||||
assert!(!server_external.supports_scheme("ring"));
|
||||
assert_eq!(
|
||||
server_external.supports_scheme("ws"),
|
||||
cfg!(feature = "websocket")
|
||||
);
|
||||
assert_eq!(
|
||||
server_external.max_pending_tcp_upgrades("ws"),
|
||||
cfg!(feature = "websocket").then_some(std::num::NonZeroUsize::MIN)
|
||||
);
|
||||
assert_eq!(
|
||||
server_external.supports_scheme("wg"),
|
||||
cfg!(feature = "wireguard")
|
||||
);
|
||||
assert_eq!(
|
||||
server_external.supports_scheme("quic"),
|
||||
cfg!(feature = "quic")
|
||||
);
|
||||
|
||||
let server = runtime_server_protocol_upgrader(global_ctx);
|
||||
assert!(server.supports_scheme("tcp"));
|
||||
assert!(server.supports_scheme("udp"));
|
||||
assert!(server.supports_scheme("ring"));
|
||||
assert_eq!(server.supports_scheme("unix"), cfg!(unix));
|
||||
assert_eq!(server.supports_scheme("ws"), cfg!(feature = "websocket"));
|
||||
assert_eq!(server.supports_scheme("wg"), cfg!(feature = "wireguard"));
|
||||
assert_eq!(server.supports_scheme("quic"), cfg!(feature = "quic"));
|
||||
}
|
||||
|
||||
#[cfg(feature = "websocket")]
|
||||
#[rstest::rstest]
|
||||
#[case("ws")]
|
||||
#[case("wss")]
|
||||
#[tokio::test]
|
||||
async fn runtime_websocket_upgraders_share_one_native_engine(#[case] scheme: &str) {
|
||||
use easytier_core::packet::ZCPacket;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
|
||||
let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0))
|
||||
.await
|
||||
.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
let url: url::Url = format!("{scheme}://{addr}").parse().unwrap();
|
||||
let global_ctx = get_mock_global_ctx();
|
||||
let server = runtime_server_protocol_upgrader(global_ctx.clone());
|
||||
let client = runtime_client_protocol_upgrader(global_ctx);
|
||||
|
||||
assert_eq!(
|
||||
client.connect_timeout(scheme),
|
||||
Some(crate::tunnel::websocket::CONNECT_TIMEOUT)
|
||||
);
|
||||
|
||||
let server_url = url.clone();
|
||||
let server_task = tokio::spawn(async move {
|
||||
let (socket, _) = listener.accept().await.unwrap();
|
||||
let ServerProtocolUpgrade::Tunnel(tunnel) = server
|
||||
.upgrade_tcp(RuntimeTcpSocket::new(socket), server_url)
|
||||
.await
|
||||
.unwrap()
|
||||
else {
|
||||
panic!("WebSocket must upgrade directly to a tunnel");
|
||||
};
|
||||
crate::tunnel::common::tests::_tunnel_echo_server(tunnel, true).await;
|
||||
});
|
||||
|
||||
tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
||||
let socket = tokio::net::TcpStream::connect(addr).await.unwrap();
|
||||
let tunnel = client
|
||||
.upgrade_client(ConnectedTransport::Tcp(RuntimeTcpSocket::new(socket)), url)
|
||||
.await
|
||||
.unwrap();
|
||||
let (mut recv, mut send) = tunnel.split();
|
||||
send.send(ZCPacket::new_with_payload(b"runtime websocket seam"))
|
||||
.await
|
||||
.unwrap();
|
||||
let packet = recv.next().await.unwrap().unwrap();
|
||||
assert_eq!(packet.payload(), b"runtime websocket seam".as_slice());
|
||||
send.close().await.unwrap();
|
||||
server_task.await.unwrap();
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[cfg(feature = "websocket")]
|
||||
#[tokio::test]
|
||||
async fn runtime_websocket_upgraders_reject_ws_client_for_wss_server() {
|
||||
let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0))
|
||||
.await
|
||||
.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
let server_url: url::Url = format!("wss://{addr}").parse().unwrap();
|
||||
let client_url: url::Url = format!("ws://{addr}").parse().unwrap();
|
||||
let global_ctx = get_mock_global_ctx();
|
||||
let server = runtime_server_protocol_upgrader(global_ctx.clone());
|
||||
let client = runtime_client_protocol_upgrader(global_ctx);
|
||||
|
||||
let server_task = tokio::spawn(async move {
|
||||
let (socket, _) = listener.accept().await.unwrap();
|
||||
assert!(
|
||||
server
|
||||
.upgrade_tcp(RuntimeTcpSocket::new(socket), server_url)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
});
|
||||
|
||||
tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
||||
let socket = tokio::net::TcpStream::connect(addr).await.unwrap();
|
||||
assert!(
|
||||
client
|
||||
.upgrade_client(
|
||||
ConnectedTransport::Tcp(RuntimeTcpSocket::new(socket)),
|
||||
client_url,
|
||||
)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
server_task.await.unwrap();
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[cfg(feature = "wireguard")]
|
||||
#[rstest::rstest]
|
||||
#[case("127.0.0.1:0")]
|
||||
#[case("[::1]:0")]
|
||||
#[tokio::test]
|
||||
async fn runtime_wireguard_upgraders_consume_core_udp_sessions(#[case] bind_addr: &str) {
|
||||
use crate::{
|
||||
common::netns::NetNS, host_runtime::native_host_runtime,
|
||||
socket::udp::new_runtime_udp_session_listener,
|
||||
};
|
||||
use easytier_core::{
|
||||
connectivity::transport::{UdpSessionMode, connect_udp},
|
||||
packet::ZCPacket,
|
||||
socket::SocketListener,
|
||||
socket::udp::{
|
||||
UdpBindOptions, UdpSessionAcceptKind, UdpSessionListenRequest, UdpSessionProtocol,
|
||||
VirtualUdpSocket,
|
||||
},
|
||||
};
|
||||
use futures::{SinkExt, StreamExt};
|
||||
|
||||
let bind_addr = bind_addr.parse().unwrap();
|
||||
let mut listener = new_runtime_udp_session_listener(
|
||||
format!("wg://{bind_addr}").parse().unwrap(),
|
||||
UdpSessionListenRequest::new(
|
||||
UdpBindOptions::port_bound_listener(bind_addr).with_only_v6(bind_addr.is_ipv6()),
|
||||
),
|
||||
UdpSessionAcceptKind::Classified(UdpSessionProtocol::WireGuard),
|
||||
NetNS::new(None),
|
||||
);
|
||||
listener.listen().await.unwrap();
|
||||
let remote_addr = listener.bound_socket().unwrap().local_addr().unwrap();
|
||||
let remote_url: url::Url = format!("wg://{remote_addr}").parse().unwrap();
|
||||
|
||||
let global_ctx = get_mock_global_ctx();
|
||||
let server = runtime_server_protocol_upgrader(global_ctx.clone());
|
||||
let client = runtime_client_protocol_upgrader(global_ctx);
|
||||
let server_url = remote_url.clone();
|
||||
let server_task = tokio::spawn(async move {
|
||||
let session = listener.accept().await.unwrap();
|
||||
let ServerProtocolUpgrade::Tunnel(tunnel) =
|
||||
server.upgrade_udp(session, server_url, None).await.unwrap()
|
||||
else {
|
||||
panic!("WireGuard must upgrade directly to a tunnel");
|
||||
};
|
||||
crate::tunnel::common::tests::_tunnel_echo_server(tunnel, false).await;
|
||||
});
|
||||
|
||||
tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
||||
let connected = connect_udp(
|
||||
native_host_runtime(),
|
||||
remote_addr,
|
||||
Vec::new(),
|
||||
UdpBindOptions::direct_connect(),
|
||||
UdpSessionMode::Classified(UdpSessionProtocol::WireGuard),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let tunnel = client
|
||||
.upgrade_client(ConnectedTransport::Udp(connected), remote_url)
|
||||
.await
|
||||
.unwrap();
|
||||
let (mut recv, mut send) = tunnel.split();
|
||||
send.send(ZCPacket::new_with_payload(b"runtime WireGuard seam"))
|
||||
.await
|
||||
.unwrap();
|
||||
let packet = recv.next().await.unwrap().unwrap();
|
||||
assert_eq!(packet.payload(), b"runtime WireGuard seam".as_slice());
|
||||
let _ = send.close().await;
|
||||
server_task.abort();
|
||||
let _ = server_task.await;
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[cfg(feature = "quic")]
|
||||
#[rstest::rstest]
|
||||
#[case("127.0.0.1:0")]
|
||||
#[case("[::1]:0")]
|
||||
#[tokio::test]
|
||||
async fn runtime_quic_upgraders_consume_core_udp_sessions(#[case] bind_addr: &str) {
|
||||
use crate::{
|
||||
common::netns::NetNS, host_runtime::native_host_runtime,
|
||||
socket::udp::new_runtime_udp_session_listener,
|
||||
};
|
||||
use easytier_core::{
|
||||
connectivity::{
|
||||
protocol::ServerProtocolAdmissionController,
|
||||
transport::{UdpSessionMode, connect_udp},
|
||||
},
|
||||
packet::ZCPacket,
|
||||
socket::SocketListener,
|
||||
socket::udp::{
|
||||
UdpBindOptions, UdpSessionAcceptKind, UdpSessionListenRequest, UdpSessionProtocol,
|
||||
VirtualUdpSocket,
|
||||
},
|
||||
};
|
||||
use futures::{SinkExt, StreamExt};
|
||||
|
||||
let bind_addr = bind_addr.parse().unwrap();
|
||||
let mut listener = new_runtime_udp_session_listener(
|
||||
format!("quic://{bind_addr}").parse().unwrap(),
|
||||
UdpSessionListenRequest::new(
|
||||
UdpBindOptions::port_bound_listener(bind_addr).with_only_v6(bind_addr.is_ipv6()),
|
||||
),
|
||||
UdpSessionAcceptKind::Classified(UdpSessionProtocol::Quic),
|
||||
NetNS::new(None),
|
||||
);
|
||||
listener.listen().await.unwrap();
|
||||
let remote_addr = listener.bound_socket().unwrap().local_addr().unwrap();
|
||||
let remote_url: url::Url = format!("quic://{remote_addr}").parse().unwrap();
|
||||
|
||||
let global_ctx = get_mock_global_ctx();
|
||||
let server = runtime_server_protocol_upgrader(global_ctx.clone());
|
||||
let client = runtime_client_protocol_upgrader(global_ctx);
|
||||
let server_url = remote_url.clone();
|
||||
let server_task = tokio::spawn(async move {
|
||||
let session = listener.accept().await.unwrap();
|
||||
let admission = ServerProtocolAdmissionController::quic()
|
||||
.try_admit()
|
||||
.unwrap();
|
||||
let ServerProtocolUpgrade::Acceptor(mut accepted) = server
|
||||
.upgrade_udp(session, server_url, Some(admission))
|
||||
.await
|
||||
.unwrap()
|
||||
else {
|
||||
panic!("QUIC must keep accepting connections from its UDP session");
|
||||
};
|
||||
let tunnel = accepted.accept().await.unwrap();
|
||||
crate::tunnel::common::tests::_tunnel_echo_server(tunnel, false).await;
|
||||
});
|
||||
|
||||
tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
||||
let connected = connect_udp(
|
||||
native_host_runtime(),
|
||||
remote_addr,
|
||||
Vec::new(),
|
||||
UdpBindOptions::direct_connect(),
|
||||
UdpSessionMode::Classified(UdpSessionProtocol::Quic),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let tunnel = client
|
||||
.upgrade_client(ConnectedTransport::Udp(connected), remote_url)
|
||||
.await
|
||||
.unwrap();
|
||||
let (mut recv, mut send) = tunnel.split();
|
||||
send.send(ZCPacket::new_with_payload(b"runtime QUIC seam"))
|
||||
.await
|
||||
.unwrap();
|
||||
let packet = recv.next().await.unwrap().unwrap();
|
||||
assert_eq!(packet.payload(), b"runtime QUIC seam".as_slice());
|
||||
let _ = send.close().await;
|
||||
server_task.await.unwrap();
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use easytier_core::connectivity::protocol::{ClientProtocolUpgrader, ServerProtocolUpgrader};
|
||||
|
||||
use crate::{common::global_ctx::ArcGlobalCtx, socket::tcp::RuntimeTcpSocket};
|
||||
|
||||
#[cfg(feature = "quic")]
|
||||
mod quic;
|
||||
#[cfg(feature = "websocket")]
|
||||
mod websocket;
|
||||
#[cfg(feature = "wireguard")]
|
||||
mod wireguard;
|
||||
|
||||
pub(super) type ClientAdapter = Arc<dyn ClientProtocolUpgrader<RuntimeTcpSocket>>;
|
||||
pub(super) type ServerAdapter = Arc<dyn ServerProtocolUpgrader<RuntimeTcpSocket>>;
|
||||
|
||||
pub(super) fn client_adapters(global_ctx: &ArcGlobalCtx) -> Vec<ClientAdapter> {
|
||||
let _ = global_ctx;
|
||||
[
|
||||
#[cfg(feature = "websocket")]
|
||||
websocket::client_adapter(global_ctx),
|
||||
#[cfg(feature = "wireguard")]
|
||||
wireguard::client_adapter(global_ctx),
|
||||
#[cfg(feature = "quic")]
|
||||
quic::client_adapter(global_ctx),
|
||||
]
|
||||
.into_iter()
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn server_adapters(global_ctx: &ArcGlobalCtx) -> Vec<ServerAdapter> {
|
||||
let _ = global_ctx;
|
||||
[
|
||||
#[cfg(feature = "websocket")]
|
||||
websocket::server_adapter(global_ctx),
|
||||
#[cfg(feature = "wireguard")]
|
||||
wireguard::server_adapter(global_ctx),
|
||||
#[cfg(feature = "quic")]
|
||||
quic::server_adapter(global_ctx),
|
||||
]
|
||||
.into_iter()
|
||||
.collect()
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use easytier_core::{
|
||||
connectivity::{
|
||||
protocol::{
|
||||
ClientProtocolUpgrader, ServerProtocolAdmission, ServerProtocolUpgrade,
|
||||
ServerProtocolUpgrader,
|
||||
},
|
||||
transport::ConnectedTransport,
|
||||
},
|
||||
socket::udp::UdpSession,
|
||||
tunnel::Tunnel,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
common::global_ctx::ArcGlobalCtx,
|
||||
socket::tcp::RuntimeTcpSocket,
|
||||
tunnel::quic::{QuicAcceptedSession, upgrade_connected},
|
||||
};
|
||||
|
||||
use super::{ClientAdapter, ServerAdapter};
|
||||
|
||||
#[derive(Default)]
|
||||
struct QuicAdapter;
|
||||
|
||||
pub(super) fn client_adapter(_global_ctx: &ArcGlobalCtx) -> ClientAdapter {
|
||||
Arc::new(QuicAdapter)
|
||||
}
|
||||
|
||||
pub(super) fn server_adapter(_global_ctx: &ArcGlobalCtx) -> ServerAdapter {
|
||||
Arc::new(QuicAdapter)
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ClientProtocolUpgrader<RuntimeTcpSocket> for QuicAdapter {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
scheme == "quic"
|
||||
}
|
||||
|
||||
async fn upgrade_client(
|
||||
&self,
|
||||
connected: ConnectedTransport<RuntimeTcpSocket>,
|
||||
requested_url: url::Url,
|
||||
) -> anyhow::Result<Box<dyn Tunnel>> {
|
||||
let ConnectedTransport::Udp(session) = connected else {
|
||||
anyhow::bail!("QUIC protocol requires a UDP session");
|
||||
};
|
||||
Ok(upgrade_connected(session, requested_url).await?)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ServerProtocolUpgrader<RuntimeTcpSocket> for QuicAdapter {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
scheme == "quic"
|
||||
}
|
||||
|
||||
async fn upgrade_tcp(
|
||||
&self,
|
||||
_socket: RuntimeTcpSocket,
|
||||
_local_url: url::Url,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
anyhow::bail!("unsupported native TCP server protocol upgrader: quic")
|
||||
}
|
||||
|
||||
async fn upgrade_udp(
|
||||
&self,
|
||||
session: UdpSession,
|
||||
local_url: url::Url,
|
||||
admission: Option<ServerProtocolAdmission>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
let admission =
|
||||
admission.ok_or_else(|| anyhow::anyhow!("QUIC server admission permit is missing"))?;
|
||||
Ok(ServerProtocolUpgrade::Acceptor(Box::new(
|
||||
QuicAcceptedSession::new(session, local_url, admission)?,
|
||||
)))
|
||||
}
|
||||
|
||||
async fn upgrade_byte_stream(
|
||||
&self,
|
||||
_socket: RuntimeTcpSocket,
|
||||
_local_url: url::Url,
|
||||
_remote_url: Option<url::Url>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
anyhow::bail!("unsupported native byte-stream server protocol upgrader: quic")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
use std::{num::NonZeroUsize, sync::Arc, time::Duration};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use easytier_core::{
|
||||
connectivity::{
|
||||
protocol::{
|
||||
ClientProtocolUpgrader, ServerProtocolAdmission, ServerProtocolUpgrade,
|
||||
ServerProtocolUpgrader,
|
||||
},
|
||||
transport::ConnectedTransport,
|
||||
},
|
||||
socket::udp::UdpSession,
|
||||
tunnel::Tunnel,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
common::global_ctx::ArcGlobalCtx,
|
||||
socket::tcp::RuntimeTcpSocket,
|
||||
tunnel::websocket::{
|
||||
CONNECT_TIMEOUT, SERVER_HANDSHAKE_TIMEOUT, upgrade_accepted, upgrade_connected,
|
||||
},
|
||||
};
|
||||
|
||||
use super::{ClientAdapter, ServerAdapter};
|
||||
|
||||
#[derive(Default)]
|
||||
struct WebSocketAdapter;
|
||||
|
||||
fn supports_scheme(scheme: &str) -> bool {
|
||||
matches!(scheme, "ws" | "wss")
|
||||
}
|
||||
|
||||
pub(super) fn client_adapter(_global_ctx: &ArcGlobalCtx) -> ClientAdapter {
|
||||
Arc::new(WebSocketAdapter)
|
||||
}
|
||||
|
||||
pub(super) fn server_adapter(_global_ctx: &ArcGlobalCtx) -> ServerAdapter {
|
||||
Arc::new(WebSocketAdapter)
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ClientProtocolUpgrader<RuntimeTcpSocket> for WebSocketAdapter {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
supports_scheme(scheme)
|
||||
}
|
||||
|
||||
fn connect_timeout(&self, scheme: &str) -> Option<Duration> {
|
||||
supports_scheme(scheme).then_some(CONNECT_TIMEOUT)
|
||||
}
|
||||
|
||||
async fn upgrade_client(
|
||||
&self,
|
||||
connected: ConnectedTransport<RuntimeTcpSocket>,
|
||||
requested_url: url::Url,
|
||||
) -> anyhow::Result<Box<dyn Tunnel>> {
|
||||
let ConnectedTransport::Tcp(socket) = connected else {
|
||||
anyhow::bail!("WebSocket protocol requires a TCP transport");
|
||||
};
|
||||
Ok(upgrade_connected(socket, requested_url).await?)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ServerProtocolUpgrader<RuntimeTcpSocket> for WebSocketAdapter {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
supports_scheme(scheme)
|
||||
}
|
||||
|
||||
fn max_pending_tcp_upgrades(&self, scheme: &str) -> Option<NonZeroUsize> {
|
||||
supports_scheme(scheme).then_some(NonZeroUsize::MIN)
|
||||
}
|
||||
|
||||
async fn upgrade_tcp(
|
||||
&self,
|
||||
socket: RuntimeTcpSocket,
|
||||
local_url: url::Url,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
Ok(ServerProtocolUpgrade::Tunnel(
|
||||
tokio::time::timeout(
|
||||
SERVER_HANDSHAKE_TIMEOUT,
|
||||
upgrade_accepted(socket, local_url),
|
||||
)
|
||||
.await??,
|
||||
))
|
||||
}
|
||||
|
||||
async fn upgrade_udp(
|
||||
&self,
|
||||
_session: UdpSession,
|
||||
_local_url: url::Url,
|
||||
_admission: Option<ServerProtocolAdmission>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
anyhow::bail!("WebSocket protocol requires a TCP transport")
|
||||
}
|
||||
|
||||
async fn upgrade_byte_stream(
|
||||
&self,
|
||||
_socket: RuntimeTcpSocket,
|
||||
_local_url: url::Url,
|
||||
_remote_url: Option<url::Url>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
anyhow::bail!("WebSocket protocol requires a TCP transport")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use easytier_core::{
|
||||
connectivity::{
|
||||
protocol::{
|
||||
ClientProtocolUpgrader, ServerProtocolAdmission, ServerProtocolUpgrade,
|
||||
ServerProtocolUpgrader,
|
||||
},
|
||||
transport::ConnectedTransport,
|
||||
},
|
||||
socket::udp::UdpSession,
|
||||
tunnel::Tunnel,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
common::global_ctx::ArcGlobalCtx,
|
||||
socket::tcp::RuntimeTcpSocket,
|
||||
tunnel::wireguard::{WgConfig, upgrade_accepted, upgrade_connected},
|
||||
};
|
||||
|
||||
use super::{ClientAdapter, ServerAdapter};
|
||||
|
||||
struct WireGuardAdapter {
|
||||
config: WgConfig,
|
||||
}
|
||||
|
||||
impl WireGuardAdapter {
|
||||
fn new(global_ctx: &ArcGlobalCtx) -> Self {
|
||||
let identity = global_ctx.get_network_identity();
|
||||
Self {
|
||||
config: WgConfig::new_from_network_identity(
|
||||
&identity.network_name,
|
||||
&identity.network_secret.unwrap_or_default(),
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn client_adapter(global_ctx: &ArcGlobalCtx) -> ClientAdapter {
|
||||
Arc::new(WireGuardAdapter::new(global_ctx))
|
||||
}
|
||||
|
||||
pub(super) fn server_adapter(global_ctx: &ArcGlobalCtx) -> ServerAdapter {
|
||||
Arc::new(WireGuardAdapter::new(global_ctx))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ClientProtocolUpgrader<RuntimeTcpSocket> for WireGuardAdapter {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
scheme == "wg"
|
||||
}
|
||||
|
||||
async fn upgrade_client(
|
||||
&self,
|
||||
connected: ConnectedTransport<RuntimeTcpSocket>,
|
||||
requested_url: url::Url,
|
||||
) -> anyhow::Result<Box<dyn Tunnel>> {
|
||||
let ConnectedTransport::Udp(session) = connected else {
|
||||
anyhow::bail!("WireGuard protocol requires a UDP session");
|
||||
};
|
||||
Ok(upgrade_connected(session, requested_url, self.config.clone()).await?)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ServerProtocolUpgrader<RuntimeTcpSocket> for WireGuardAdapter {
|
||||
fn supports_scheme(&self, scheme: &str) -> bool {
|
||||
scheme == "wg"
|
||||
}
|
||||
|
||||
async fn upgrade_tcp(
|
||||
&self,
|
||||
_socket: RuntimeTcpSocket,
|
||||
_local_url: url::Url,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
anyhow::bail!("unsupported native TCP server protocol upgrader: wg")
|
||||
}
|
||||
|
||||
async fn upgrade_udp(
|
||||
&self,
|
||||
session: UdpSession,
|
||||
_local_url: url::Url,
|
||||
_admission: Option<ServerProtocolAdmission>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
Ok(ServerProtocolUpgrade::Tunnel(upgrade_accepted(
|
||||
session,
|
||||
self.config.clone(),
|
||||
)?))
|
||||
}
|
||||
|
||||
async fn upgrade_byte_stream(
|
||||
&self,
|
||||
_socket: RuntimeTcpSocket,
|
||||
_local_url: url::Url,
|
||||
_remote_url: Option<url::Url>,
|
||||
) -> anyhow::Result<ServerProtocolUpgrade> {
|
||||
anyhow::bail!("unsupported native byte-stream server protocol upgrader: wg")
|
||||
}
|
||||
}
|
||||
+342
-720
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,279 @@
|
||||
use std::{
|
||||
fmt,
|
||||
future::Future,
|
||||
io::{self, IoSliceMut},
|
||||
pin::Pin,
|
||||
sync::{Arc, Mutex},
|
||||
task::{Context, Poll},
|
||||
};
|
||||
|
||||
use easytier_core::{
|
||||
connectivity::transport::ConnectedUdpSession,
|
||||
socket::udp::{UdpSession, UdpSessionSocket},
|
||||
};
|
||||
use quinn::{
|
||||
AsyncUdpSocket, UdpPoller,
|
||||
udp::{RecvMeta, Transmit},
|
||||
};
|
||||
use tokio::sync::mpsc::{self, Receiver, Sender, error::TrySendError};
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
|
||||
const DATAGRAM_QUEUE_CAPACITY: usize = 1024;
|
||||
|
||||
type SendBatch = Vec<Vec<u8>>;
|
||||
type WritableFuture = Pin<Box<dyn Future<Output = io::Result<()>> + Send>>;
|
||||
|
||||
struct ReceivedDatagram {
|
||||
payload: Vec<u8>,
|
||||
dst_ip: Option<std::net::IpAddr>,
|
||||
}
|
||||
|
||||
pub(crate) struct QuicUdpSessionSocket {
|
||||
_session: Arc<dyn UdpSessionSocket>,
|
||||
local_addr: std::net::SocketAddr,
|
||||
peer_addr: std::net::SocketAddr,
|
||||
incoming: Mutex<Receiver<io::Result<ReceivedDatagram>>>,
|
||||
outgoing: Sender<SendBatch>,
|
||||
_recv_task: AbortOnDropHandle<()>,
|
||||
_send_task: AbortOnDropHandle<()>,
|
||||
_session_guard: Box<dyn Send + Sync>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for QuicUdpSessionSocket {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("QuicUdpSessionSocket")
|
||||
.field("local_addr", &self.local_addr)
|
||||
.field("peer_addr", &self.peer_addr)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl QuicUdpSessionSocket {
|
||||
pub(crate) fn new(connected: ConnectedUdpSession) -> io::Result<Self> {
|
||||
let (session, session_guard) = connected.into_parts();
|
||||
Self::from_session(Arc::new(session), session_guard)
|
||||
}
|
||||
|
||||
pub(crate) fn from_accepted<T>(session: UdpSession, session_guard: T) -> io::Result<Self>
|
||||
where
|
||||
T: Send + Sync + 'static,
|
||||
{
|
||||
Self::from_session(Arc::new(session), Box::new(session_guard))
|
||||
}
|
||||
|
||||
fn from_session(
|
||||
session: Arc<dyn UdpSessionSocket>,
|
||||
session_guard: Box<dyn Send + Sync>,
|
||||
) -> io::Result<Self> {
|
||||
let local_addr = session.local_addr()?;
|
||||
let peer_addr = session.peer_addr()?;
|
||||
let (incoming_tx, incoming) = mpsc::channel(DATAGRAM_QUEUE_CAPACITY);
|
||||
let (outgoing, mut outgoing_rx) = mpsc::channel::<SendBatch>(DATAGRAM_QUEUE_CAPACITY);
|
||||
|
||||
let recv_session = session.clone();
|
||||
let recv_errors = incoming_tx.clone();
|
||||
let recv_task = AbortOnDropHandle::new(tokio::spawn(async move {
|
||||
let mut buffer = vec![0; 64 * 1024];
|
||||
loop {
|
||||
match recv_session.recv_with_meta(&mut buffer).await {
|
||||
Ok((length, meta)) => {
|
||||
let datagram = ReceivedDatagram {
|
||||
payload: buffer[..length].to_vec(),
|
||||
dst_ip: meta.dst_ip,
|
||||
};
|
||||
if incoming_tx.send(Ok(datagram)).await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = incoming_tx.send(Err(error)).await;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}));
|
||||
|
||||
let send_session = session.clone();
|
||||
let send_task = AbortOnDropHandle::new(tokio::spawn(async move {
|
||||
while let Some(batch) = outgoing_rx.recv().await {
|
||||
for datagram in batch {
|
||||
match send_session.send(&datagram).await {
|
||||
Ok(length) if length == datagram.len() => {}
|
||||
Ok(_) => {
|
||||
let _ = recv_errors
|
||||
.send(Err(io::Error::new(
|
||||
io::ErrorKind::WriteZero,
|
||||
"QUIC UDP session partially sent a datagram",
|
||||
)))
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = recv_errors.send(Err(error)).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}));
|
||||
|
||||
Ok(Self {
|
||||
_session: session,
|
||||
local_addr,
|
||||
peer_addr,
|
||||
incoming: Mutex::new(incoming),
|
||||
outgoing,
|
||||
_recv_task: recv_task,
|
||||
_send_task: send_task,
|
||||
_session_guard: session_guard,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn peer_addr(&self) -> std::net::SocketAddr {
|
||||
self.peer_addr
|
||||
}
|
||||
|
||||
fn send_batch(&self, transmit: &Transmit<'_>) -> io::Result<()> {
|
||||
if transmit.destination != self.peer_addr {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::AddrNotAvailable,
|
||||
format!(
|
||||
"QUIC UDP session is connected to {}, not {}",
|
||||
self.peer_addr, transmit.destination
|
||||
),
|
||||
));
|
||||
}
|
||||
|
||||
let batch = match transmit.segment_size {
|
||||
Some(0) => {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"QUIC segment size cannot be zero",
|
||||
));
|
||||
}
|
||||
Some(segment_size) => transmit
|
||||
.contents
|
||||
.chunks(segment_size)
|
||||
.map(<[u8]>::to_vec)
|
||||
.collect(),
|
||||
None => vec![transmit.contents.to_vec()],
|
||||
};
|
||||
self.outgoing.try_send(batch).map_err(map_try_send_error)
|
||||
}
|
||||
}
|
||||
|
||||
fn map_try_send_error(error: TrySendError<SendBatch>) -> io::Error {
|
||||
match error {
|
||||
TrySendError::Full(_) => io::Error::new(io::ErrorKind::WouldBlock, "QUIC send queue full"),
|
||||
TrySendError::Closed(_) => {
|
||||
io::Error::new(io::ErrorKind::BrokenPipe, "QUIC UDP session closed")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct QuicUdpSessionPoller {
|
||||
outgoing: Sender<SendBatch>,
|
||||
writable: Mutex<Option<WritableFuture>>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for QuicUdpSessionPoller {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("QuicUdpSessionPoller")
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl UdpPoller for QuicUdpSessionPoller {
|
||||
fn poll_writable(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
|
||||
let mut writable = self.writable.lock().unwrap();
|
||||
if writable.is_none() {
|
||||
let outgoing = self.outgoing.clone();
|
||||
*writable = Some(Box::pin(async move {
|
||||
let permit = outgoing.reserve_owned().await.map_err(|_| {
|
||||
io::Error::new(io::ErrorKind::BrokenPipe, "QUIC UDP session closed")
|
||||
})?;
|
||||
drop(permit);
|
||||
Ok(())
|
||||
}));
|
||||
}
|
||||
|
||||
match writable.as_mut().unwrap().as_mut().poll(context) {
|
||||
Poll::Ready(result) => {
|
||||
*writable = None;
|
||||
Poll::Ready(result)
|
||||
}
|
||||
Poll::Pending => Poll::Pending,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncUdpSocket for QuicUdpSessionSocket {
|
||||
fn create_io_poller(self: Arc<Self>) -> Pin<Box<dyn UdpPoller>> {
|
||||
Box::pin(QuicUdpSessionPoller {
|
||||
outgoing: self.outgoing.clone(),
|
||||
writable: Mutex::new(None),
|
||||
})
|
||||
}
|
||||
|
||||
fn try_send(&self, transmit: &Transmit<'_>) -> io::Result<()> {
|
||||
self.send_batch(transmit)
|
||||
}
|
||||
|
||||
fn poll_recv(
|
||||
&self,
|
||||
context: &mut Context<'_>,
|
||||
buffers: &mut [IoSliceMut<'_>],
|
||||
meta: &mut [RecvMeta],
|
||||
) -> Poll<io::Result<usize>> {
|
||||
if buffers.is_empty() || meta.is_empty() {
|
||||
return Poll::Ready(Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"QUIC UDP recv buffers are empty",
|
||||
)));
|
||||
}
|
||||
|
||||
let mut incoming = self.incoming.lock().unwrap();
|
||||
loop {
|
||||
match Pin::new(&mut *incoming).poll_recv(context) {
|
||||
Poll::Ready(Some(Ok(datagram))) => {
|
||||
if buffers[0].len() < datagram.payload.len() {
|
||||
tracing::debug!(
|
||||
payload_len = datagram.payload.len(),
|
||||
recv_buf_len = buffers[0].len(),
|
||||
peer_addr = ?self.peer_addr,
|
||||
"drop oversized QUIC UDP session datagram"
|
||||
);
|
||||
continue;
|
||||
}
|
||||
buffers[0][..datagram.payload.len()].copy_from_slice(&datagram.payload);
|
||||
meta[0] = RecvMeta {
|
||||
addr: self.peer_addr,
|
||||
len: datagram.payload.len(),
|
||||
stride: datagram.payload.len(),
|
||||
ecn: None,
|
||||
dst_ip: datagram.dst_ip,
|
||||
};
|
||||
return Poll::Ready(Ok(1));
|
||||
}
|
||||
Poll::Ready(Some(Err(error))) => return Poll::Ready(Err(error)),
|
||||
Poll::Ready(None) => {
|
||||
return Poll::Ready(Err(io::Error::new(
|
||||
io::ErrorKind::UnexpectedEof,
|
||||
"QUIC UDP session closed",
|
||||
)));
|
||||
}
|
||||
Poll::Pending => return Poll::Pending,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn local_addr(&self) -> io::Result<std::net::SocketAddr> {
|
||||
Ok(self.local_addr)
|
||||
}
|
||||
|
||||
fn may_fragment(&self) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
@@ -1,391 +0,0 @@
|
||||
use async_ringbuf::{AsyncHeapCons, AsyncHeapProd, AsyncHeapRb, traits::*};
|
||||
use crossbeam::atomic::AtomicCell;
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
fmt::Debug,
|
||||
sync::Arc,
|
||||
task::{Poll, ready},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use futures::{Sink, SinkExt, Stream, StreamExt};
|
||||
use once_cell::sync::Lazy;
|
||||
|
||||
use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
|
||||
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::tunnel::{FromUrl, IpVersion, SinkError, SinkItem};
|
||||
|
||||
use super::{
|
||||
StreamItem, Tunnel, TunnelConnector, TunnelError, TunnelInfo, TunnelListener,
|
||||
build_url_from_socket_addr, common::TunnelWrapper,
|
||||
};
|
||||
|
||||
pub static RING_TUNNEL_CAP: usize = 128;
|
||||
static RING_TUNNEL_RESERVED_CAP: usize = 4;
|
||||
|
||||
type RingLock = parking_lot::Mutex<()>;
|
||||
|
||||
type RingItem = SinkItem;
|
||||
|
||||
pub struct RingTunnel {
|
||||
id: Uuid,
|
||||
|
||||
ring_cons_impl: AtomicCell<Option<AsyncHeapCons<RingItem>>>,
|
||||
ring_prod_impl: AtomicCell<Option<AsyncHeapProd<RingItem>>>,
|
||||
}
|
||||
|
||||
impl RingTunnel {
|
||||
fn id(&self) -> &Uuid {
|
||||
&self.id
|
||||
}
|
||||
|
||||
pub fn new(cap: usize) -> Self {
|
||||
let id = Uuid::new_v4();
|
||||
let ring_impl = AsyncHeapRb::new(std::cmp::max(RING_TUNNEL_RESERVED_CAP * 2, cap));
|
||||
let (ring_prod_impl, ring_cons_impl) = ring_impl.split();
|
||||
Self {
|
||||
id,
|
||||
ring_cons_impl: AtomicCell::new(Some(ring_cons_impl)),
|
||||
ring_prod_impl: AtomicCell::new(Some(ring_prod_impl)),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_with_id(id: Uuid, cap: usize) -> Self {
|
||||
let mut ret = Self::new(cap);
|
||||
ret.id = id;
|
||||
ret
|
||||
}
|
||||
}
|
||||
|
||||
impl Debug for RingTunnel {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("RingTunnel").field("id", &self.id).finish()
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RingStream {
|
||||
id: Uuid,
|
||||
ring_cons_impl: AsyncHeapCons<RingItem>,
|
||||
}
|
||||
|
||||
impl RingStream {
|
||||
pub fn new(tunnel: Arc<RingTunnel>) -> Self {
|
||||
Self {
|
||||
id: tunnel.id,
|
||||
ring_cons_impl: tunnel.ring_cons_impl.take().unwrap(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Stream for RingStream {
|
||||
type Item = StreamItem;
|
||||
|
||||
fn poll_next(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> Poll<Option<Self::Item>> {
|
||||
let ret = ready!(self.get_mut().ring_cons_impl.poll_next_unpin(cx));
|
||||
match ret {
|
||||
Some(item) => Poll::Ready(Some(Ok(item))),
|
||||
None => Poll::Ready(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Debug for RingStream {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("RingStream")
|
||||
.field("id", &self.id)
|
||||
.field("len", &self.ring_cons_impl.base().occupied_len())
|
||||
.field("cap", &self.ring_cons_impl.base().capacity())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RingSink {
|
||||
id: Uuid,
|
||||
ring_prod_impl: AsyncHeapProd<RingItem>,
|
||||
}
|
||||
|
||||
impl RingSink {
|
||||
pub fn new(tunnel: Arc<RingTunnel>) -> Self {
|
||||
Self {
|
||||
id: tunnel.id,
|
||||
ring_prod_impl: tunnel.ring_prod_impl.take().unwrap(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn try_send(&mut self, item: RingItem) -> Result<(), RingItem> {
|
||||
let base = self.ring_prod_impl.base();
|
||||
if base.occupied_len() >= base.capacity().get() - RING_TUNNEL_RESERVED_CAP {
|
||||
return Err(item);
|
||||
}
|
||||
self.ring_prod_impl.try_push(item)
|
||||
}
|
||||
|
||||
pub fn force_send(&mut self, item: RingItem) -> Result<(), RingItem> {
|
||||
self.ring_prod_impl.try_push(item)
|
||||
}
|
||||
}
|
||||
|
||||
impl Sink<SinkItem> for RingSink {
|
||||
type Error = SinkError;
|
||||
|
||||
fn poll_ready(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Result<(), Self::Error>> {
|
||||
let ret = ready!(self.get_mut().ring_prod_impl.poll_ready_unpin(cx));
|
||||
Poll::Ready(ret.map_err(|_| TunnelError::Shutdown))
|
||||
}
|
||||
|
||||
fn start_send(self: std::pin::Pin<&mut Self>, item: SinkItem) -> Result<(), Self::Error> {
|
||||
self.get_mut()
|
||||
.ring_prod_impl
|
||||
.start_send_unpin(item)
|
||||
.map_err(|_| TunnelError::Shutdown)
|
||||
}
|
||||
|
||||
fn poll_flush(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Result<(), Self::Error>> {
|
||||
let ret = ready!(self.get_mut().ring_prod_impl.poll_flush_unpin(cx));
|
||||
Poll::Ready(ret.map_err(|_| TunnelError::Shutdown))
|
||||
}
|
||||
|
||||
fn poll_close(
|
||||
self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Result<(), Self::Error>> {
|
||||
let ret = ready!(self.get_mut().ring_prod_impl.poll_close_unpin(cx));
|
||||
Poll::Ready(ret.map_err(|_| TunnelError::Shutdown))
|
||||
}
|
||||
}
|
||||
|
||||
impl Debug for RingSink {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("RingSink")
|
||||
.field("id", &self.id)
|
||||
.field("len", &self.ring_prod_impl.base().occupied_len())
|
||||
.field("cap", &self.ring_prod_impl.base().capacity())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
struct Connection {
|
||||
client: Arc<RingTunnel>,
|
||||
server: Arc<RingTunnel>,
|
||||
}
|
||||
|
||||
type ConnectionMap = HashMap<uuid::Uuid, UnboundedSender<Arc<Connection>>>;
|
||||
|
||||
static CONNECTION_MAP: Lazy<Arc<std::sync::Mutex<ConnectionMap>>> =
|
||||
Lazy::new(|| Arc::new(std::sync::Mutex::new(HashMap::new())));
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct RingTunnelListener {
|
||||
listener_addr: url::Url,
|
||||
conn_sender: UnboundedSender<Arc<Connection>>,
|
||||
conn_receiver: UnboundedReceiver<Arc<Connection>>,
|
||||
|
||||
key_in_conn_map: Option<uuid::Uuid>,
|
||||
}
|
||||
|
||||
impl RingTunnelListener {
|
||||
pub fn new(key: url::Url) -> Self {
|
||||
let (conn_sender, conn_receiver) = unbounded_channel();
|
||||
RingTunnelListener {
|
||||
listener_addr: key,
|
||||
conn_sender,
|
||||
conn_receiver,
|
||||
key_in_conn_map: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn get_tunnel_for_client(conn: Arc<Connection>) -> impl Tunnel {
|
||||
TunnelWrapper::new(
|
||||
RingStream::new(conn.client.clone()),
|
||||
RingSink::new(conn.server.clone()),
|
||||
Some(TunnelInfo {
|
||||
tunnel_type: "ring".to_owned(),
|
||||
local_addr: Some(build_url_from_socket_addr(&conn.client.id.into(), "ring").into()),
|
||||
remote_addr: Some(build_url_from_socket_addr(&conn.server.id.into(), "ring").into()),
|
||||
resolved_remote_addr: Some(
|
||||
build_url_from_socket_addr(&conn.server.id.into(), "ring").into(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn get_tunnel_for_server(conn: Arc<Connection>) -> impl Tunnel {
|
||||
TunnelWrapper::new(
|
||||
RingStream::new(conn.server.clone()),
|
||||
RingSink::new(conn.client.clone()),
|
||||
Some(TunnelInfo {
|
||||
tunnel_type: "ring".to_owned(),
|
||||
local_addr: Some(build_url_from_socket_addr(&conn.server.id.into(), "ring").into()),
|
||||
remote_addr: Some(build_url_from_socket_addr(&conn.client.id.into(), "ring").into()),
|
||||
resolved_remote_addr: Some(
|
||||
build_url_from_socket_addr(&conn.client.id.into(), "ring").into(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
impl RingTunnelListener {
|
||||
async fn get_addr(&self) -> Result<Uuid, TunnelError> {
|
||||
Uuid::from_url(self.listener_addr.clone(), IpVersion::Both).await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TunnelListener for RingTunnelListener {
|
||||
async fn listen(&mut self) -> Result<(), TunnelError> {
|
||||
tracing::info!("listen new conn of key: {}", self.listener_addr);
|
||||
let addr = self.get_addr().await?;
|
||||
CONNECTION_MAP
|
||||
.lock()
|
||||
.unwrap()
|
||||
.insert(addr, self.conn_sender.clone());
|
||||
self.key_in_conn_map = Some(addr);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn accept(&mut self) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
tracing::info!("waiting accept new conn of key: {}", self.listener_addr);
|
||||
let my_addr = self.get_addr().await?;
|
||||
if let Some(conn) = self.conn_receiver.recv().await {
|
||||
if conn.server.id == my_addr {
|
||||
tracing::info!("accept new conn of key: {}", self.listener_addr);
|
||||
return Ok(Box::new(get_tunnel_for_server(conn)));
|
||||
} else {
|
||||
tracing::error!(?conn.server.id, ?my_addr, "got new conn with wrong id");
|
||||
return Err(TunnelError::InternalError(
|
||||
"accept got wrong ring server id".to_owned(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
return Err(TunnelError::InternalError(
|
||||
"conn receiver stopped".to_owned(),
|
||||
));
|
||||
}
|
||||
|
||||
fn local_url(&self) -> url::Url {
|
||||
self.listener_addr.clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for RingTunnelListener {
|
||||
fn drop(&mut self) {
|
||||
if let Some(addr) = self.key_in_conn_map {
|
||||
CONNECTION_MAP.lock().unwrap().remove(&addr);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RingTunnelConnector {
|
||||
remote_addr: url::Url,
|
||||
}
|
||||
|
||||
impl RingTunnelConnector {
|
||||
pub fn new(remote_addr: url::Url) -> Self {
|
||||
RingTunnelConnector { remote_addr }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TunnelConnector for RingTunnelConnector {
|
||||
async fn connect(&mut self) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
let remote_addr = Uuid::from_url(self.remote_addr.clone(), IpVersion::Both).await?;
|
||||
let entry = CONNECTION_MAP
|
||||
.lock()
|
||||
.unwrap()
|
||||
.get(&remote_addr)
|
||||
.unwrap()
|
||||
.clone();
|
||||
tracing::info!("connecting");
|
||||
let conn = Arc::new(Connection {
|
||||
client: Arc::new(RingTunnel::new(RING_TUNNEL_CAP)),
|
||||
server: Arc::new(RingTunnel::new_with_id(remote_addr, RING_TUNNEL_CAP)),
|
||||
});
|
||||
entry
|
||||
.send(conn.clone())
|
||||
.map_err(|_| TunnelError::InternalError("send conn to listner failed".to_owned()))?;
|
||||
Ok(Box::new(get_tunnel_for_client(conn)))
|
||||
}
|
||||
|
||||
fn remote_url(&self) -> url::Url {
|
||||
self.remote_addr.clone()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn create_ring_tunnel_pair() -> (Box<dyn Tunnel>, Box<dyn Tunnel>) {
|
||||
let conn = Arc::new(Connection {
|
||||
client: Arc::new(RingTunnel::new(RING_TUNNEL_CAP)),
|
||||
server: Arc::new(RingTunnel::new(RING_TUNNEL_CAP)),
|
||||
});
|
||||
(
|
||||
Box::new(get_tunnel_for_server(conn.clone())),
|
||||
Box::new(get_tunnel_for_client(conn)),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use futures::StreamExt;
|
||||
use tokio::time::timeout;
|
||||
|
||||
use crate::tunnel::common::tests::{_tunnel_bench, _tunnel_pingpong};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn ring_pingpong() {
|
||||
let id: url::Url = format!("ring://{}", Uuid::new_v4()).parse().unwrap();
|
||||
let listener = RingTunnelListener::new(id.clone());
|
||||
let connector = RingTunnelConnector::new(id.clone());
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ring_bench() {
|
||||
let id: url::Url = format!("ring://{}", Uuid::new_v4()).parse().unwrap();
|
||||
let listener = RingTunnelListener::new(id.clone());
|
||||
let connector = RingTunnelConnector::new(id);
|
||||
_tunnel_bench(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ring_close() {
|
||||
let (stunnel, ctunnel) = create_ring_tunnel_pair();
|
||||
drop(stunnel);
|
||||
|
||||
let mut stream = ctunnel.split().0;
|
||||
let ret = stream.next().await;
|
||||
assert!(ret.as_ref().is_none(), "expect none, got {:?}", ret);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn abort_ring_stream() {
|
||||
let (_stunnel, ctunnel) = create_ring_tunnel_pair();
|
||||
let mut stream = ctunnel.split().0;
|
||||
let task = tokio::spawn(async move {
|
||||
let _ = stream.next().await;
|
||||
});
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||
task.abort();
|
||||
let _ = tokio::join!(task);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ring_stream_recv_timeout() {
|
||||
let (_stunnel, ctunnel) = create_ring_tunnel_pair();
|
||||
let mut stream = ctunnel.split().0;
|
||||
let _ = timeout(tokio::time::Duration::from_millis(10), stream.next()).await;
|
||||
}
|
||||
}
|
||||
@@ -1,132 +0,0 @@
|
||||
use std::{
|
||||
cell::UnsafeCell,
|
||||
sync::atomic::{AtomicU32, Ordering::Relaxed},
|
||||
};
|
||||
|
||||
pub struct WindowLatency {
|
||||
latency_us_window: Vec<AtomicU32>,
|
||||
latency_us_window_index: AtomicU32,
|
||||
latency_us_window_size: u32,
|
||||
|
||||
sum: AtomicU32,
|
||||
count: AtomicU32,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for WindowLatency {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("WindowLatency")
|
||||
.field("count", &self.count)
|
||||
.field("window_size", &self.latency_us_window_size)
|
||||
.field("window_latency", &self.get_latency_us::<u32>())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl WindowLatency {
|
||||
pub fn new(window_size: u32) -> Self {
|
||||
Self {
|
||||
latency_us_window: (0..window_size).map(|_| AtomicU32::new(0)).collect(),
|
||||
latency_us_window_index: AtomicU32::new(0),
|
||||
latency_us_window_size: window_size,
|
||||
|
||||
sum: AtomicU32::new(0),
|
||||
count: AtomicU32::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn record_latency(&self, latency_us: u32) {
|
||||
let index = self.latency_us_window_index.fetch_add(1, Relaxed);
|
||||
if self.count.load(Relaxed) < self.latency_us_window_size {
|
||||
self.count.fetch_add(1, Relaxed);
|
||||
}
|
||||
|
||||
let index = index % self.latency_us_window_size;
|
||||
let old_lat = self.latency_us_window[index as usize].swap(latency_us, Relaxed);
|
||||
|
||||
if old_lat < latency_us {
|
||||
self.sum.fetch_add(latency_us - old_lat, Relaxed);
|
||||
} else {
|
||||
self.sum.fetch_sub(old_lat - latency_us, Relaxed);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_latency_us<T: From<u32> + std::ops::Div<Output = T>>(&self) -> T {
|
||||
let count = self.count.load(Relaxed);
|
||||
let sum = self.sum.load(Relaxed);
|
||||
if count == 0 {
|
||||
0.into()
|
||||
} else {
|
||||
(T::from(sum)) / T::from(count)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct Throughput {
|
||||
tx_bytes: UnsafeCell<u64>,
|
||||
rx_bytes: UnsafeCell<u64>,
|
||||
tx_packets: UnsafeCell<u64>,
|
||||
rx_packets: UnsafeCell<u64>,
|
||||
}
|
||||
|
||||
impl Clone for Throughput {
|
||||
fn clone(&self) -> Self {
|
||||
Self {
|
||||
tx_bytes: UnsafeCell::new(unsafe { *self.tx_bytes.get() }),
|
||||
rx_bytes: UnsafeCell::new(unsafe { *self.rx_bytes.get() }),
|
||||
tx_packets: UnsafeCell::new(unsafe { *self.tx_packets.get() }),
|
||||
rx_packets: UnsafeCell::new(unsafe { *self.rx_packets.get() }),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// add sync::Send and sync::Sync traits to Throughput
|
||||
unsafe impl Send for Throughput {}
|
||||
unsafe impl Sync for Throughput {}
|
||||
|
||||
impl Default for Throughput {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
tx_bytes: UnsafeCell::new(0),
|
||||
rx_bytes: UnsafeCell::new(0),
|
||||
tx_packets: UnsafeCell::new(0),
|
||||
rx_packets: UnsafeCell::new(0),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Throughput {
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
pub fn tx_bytes(&self) -> u64 {
|
||||
unsafe { *self.tx_bytes.get() }
|
||||
}
|
||||
|
||||
pub fn rx_bytes(&self) -> u64 {
|
||||
unsafe { *self.rx_bytes.get() }
|
||||
}
|
||||
|
||||
pub fn tx_packets(&self) -> u64 {
|
||||
unsafe { *self.tx_packets.get() }
|
||||
}
|
||||
|
||||
pub fn rx_packets(&self) -> u64 {
|
||||
unsafe { *self.rx_packets.get() }
|
||||
}
|
||||
|
||||
pub fn record_tx_bytes(&self, bytes: u64) {
|
||||
unsafe {
|
||||
*self.tx_bytes.get() += bytes;
|
||||
*self.tx_packets.get() += 1;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn record_rx_bytes(&self, bytes: u64) {
|
||||
unsafe {
|
||||
*self.rx_bytes.get() += bytes;
|
||||
*self.rx_packets.get() += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,378 +0,0 @@
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use super::{FromUrl, TunnelInfo};
|
||||
use crate::tunnel::common::{apply_socket_mark, bind};
|
||||
use async_trait::async_trait;
|
||||
use futures::stream::FuturesUnordered;
|
||||
use tokio::net::{TcpListener, TcpSocket, TcpStream};
|
||||
|
||||
use super::{
|
||||
IpVersion, Tunnel, TunnelError, TunnelListener,
|
||||
common::{FramedReader, FramedWriter, TunnelWrapper, wait_for_connect_futures},
|
||||
};
|
||||
|
||||
const TCP_MTU_BYTES: usize = 2000;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct TcpTunnelListener {
|
||||
addr: url::Url,
|
||||
listener: Option<TcpListener>,
|
||||
socket_mark: Option<u32>,
|
||||
}
|
||||
|
||||
impl TcpTunnelListener {
|
||||
pub fn new(addr: url::Url) -> Self {
|
||||
TcpTunnelListener {
|
||||
addr,
|
||||
listener: None,
|
||||
socket_mark: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn set_socket_mark(&mut self, socket_mark: Option<u32>) {
|
||||
self.socket_mark = socket_mark;
|
||||
}
|
||||
|
||||
async fn do_accept(&self) -> Result<Box<dyn Tunnel>, std::io::Error> {
|
||||
let listener = self.listener.as_ref().unwrap();
|
||||
let (stream, _) = listener.accept().await?;
|
||||
|
||||
if let Err(e) = stream.set_nodelay(true) {
|
||||
tracing::warn!(?e, "set_nodelay fail in accept");
|
||||
}
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: "tcp".to_owned(),
|
||||
local_addr: Some(self.local_url().into()),
|
||||
remote_addr: Some(
|
||||
super::build_url_from_socket_addr(&stream.peer_addr()?.to_string(), "tcp").into(),
|
||||
),
|
||||
resolved_remote_addr: Some(
|
||||
super::build_url_from_socket_addr(&stream.peer_addr()?.to_string(), "tcp").into(),
|
||||
),
|
||||
};
|
||||
|
||||
let (r, w) = stream.into_split();
|
||||
Ok(Box::new(TunnelWrapper::new(
|
||||
FramedReader::new(r, TCP_MTU_BYTES),
|
||||
FramedWriter::new(w),
|
||||
Some(info),
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TunnelListener for TcpTunnelListener {
|
||||
async fn listen(&mut self) -> Result<(), TunnelError> {
|
||||
self.listener = None;
|
||||
|
||||
let addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?;
|
||||
let listener = bind::<TcpListener>()
|
||||
.addr(addr)
|
||||
.only_v6(true)
|
||||
.maybe_socket_mark(self.socket_mark)
|
||||
.call()?;
|
||||
|
||||
self.addr
|
||||
.set_port(Some(listener.local_addr()?.port()))
|
||||
.unwrap();
|
||||
self.listener = Some(listener);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn accept(&mut self) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
loop {
|
||||
match self.do_accept().await {
|
||||
Ok(ret) => return Ok(ret),
|
||||
Err(e) => {
|
||||
use std::io::ErrorKind::*;
|
||||
if matches!(
|
||||
e.kind(),
|
||||
NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset
|
||||
) {
|
||||
tracing::warn!(?e, "accept fail with retryable error: {:?}", e);
|
||||
continue;
|
||||
}
|
||||
tracing::warn!(?e, "accept fail");
|
||||
return Err(e.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn local_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
}
|
||||
|
||||
fn get_tunnel_with_tcp_stream(
|
||||
stream: TcpStream,
|
||||
remote_url: url::Url,
|
||||
) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
if let Err(e) = stream.set_nodelay(true) {
|
||||
tracing::warn!(?e, "set_nodelay fail in get_tunnel_with_tcp_stream");
|
||||
}
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: "tcp".to_owned(),
|
||||
local_addr: Some(
|
||||
super::build_url_from_socket_addr(&stream.local_addr()?.to_string(), "tcp").into(),
|
||||
),
|
||||
remote_addr: Some(remote_url.into()),
|
||||
resolved_remote_addr: Some(
|
||||
super::build_url_from_socket_addr(&stream.peer_addr()?.to_string(), "tcp").into(),
|
||||
),
|
||||
};
|
||||
|
||||
let (r, w) = stream.into_split();
|
||||
Ok(Box::new(TunnelWrapper::new(
|
||||
FramedReader::new(r, TCP_MTU_BYTES),
|
||||
FramedWriter::new(w),
|
||||
Some(info),
|
||||
)))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct TcpTunnelConnector {
|
||||
addr: url::Url,
|
||||
|
||||
bind_addrs: Vec<SocketAddr>,
|
||||
ip_version: IpVersion,
|
||||
resolved_addr: Option<SocketAddr>,
|
||||
socket_mark: Option<u32>,
|
||||
}
|
||||
|
||||
impl TcpTunnelConnector {
|
||||
pub fn new(addr: url::Url) -> Self {
|
||||
TcpTunnelConnector {
|
||||
addr,
|
||||
bind_addrs: vec![],
|
||||
ip_version: IpVersion::Both,
|
||||
resolved_addr: None,
|
||||
socket_mark: None,
|
||||
}
|
||||
}
|
||||
|
||||
async fn connect_with_default_bind(
|
||||
&self,
|
||||
addr: SocketAddr,
|
||||
) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
tracing::info!(url = ?self.addr, ?addr, "connect tcp start, bind addrs: {:?}", self.bind_addrs);
|
||||
let stream = if self.socket_mark.is_some() {
|
||||
// SO_MARK requires applying the option on the socket before
|
||||
// connect, so go through TcpSocket rather than TcpStream::connect.
|
||||
let socket = if addr.is_ipv4() {
|
||||
TcpSocket::new_v4()?
|
||||
} else {
|
||||
TcpSocket::new_v6()?
|
||||
};
|
||||
apply_socket_mark(&socket2::SockRef::from(&socket), self.socket_mark)?;
|
||||
socket.connect(addr).await?
|
||||
} else {
|
||||
TcpStream::connect(addr).await?
|
||||
};
|
||||
tracing::info!(url = ?self.addr, ?addr, "connect tcp succ");
|
||||
get_tunnel_with_tcp_stream(stream, self.addr.clone())
|
||||
}
|
||||
|
||||
async fn connect_with_custom_bind(
|
||||
&self,
|
||||
addr: SocketAddr,
|
||||
) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
let futures = FuturesUnordered::new();
|
||||
|
||||
for bind_addr in self.bind_addrs.iter() {
|
||||
tracing::info!(?bind_addr, ?addr, "bind addr");
|
||||
match bind::<TcpSocket>()
|
||||
.addr(*bind_addr)
|
||||
.only_v6(true)
|
||||
.maybe_socket_mark(self.socket_mark)
|
||||
.call()
|
||||
{
|
||||
Ok(socket) => futures.push(socket.connect(addr)),
|
||||
Err(error) => {
|
||||
tracing::error!(?bind_addr, ?addr, ?error, "bind addr fail");
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let ret = wait_for_connect_futures(futures).await;
|
||||
get_tunnel_with_tcp_stream(ret?, self.addr.clone())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl super::TunnelConnector for TcpTunnelConnector {
|
||||
async fn connect(&mut self) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
let addr = match self.resolved_addr {
|
||||
Some(addr) => addr,
|
||||
None => SocketAddr::from_url(self.addr.clone(), self.ip_version).await?,
|
||||
};
|
||||
if self.bind_addrs.is_empty() {
|
||||
self.connect_with_default_bind(addr).await
|
||||
} else {
|
||||
self.connect_with_custom_bind(addr).await
|
||||
}
|
||||
}
|
||||
|
||||
fn remote_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
|
||||
fn set_bind_addrs(&mut self, addrs: Vec<SocketAddr>) {
|
||||
self.bind_addrs = addrs;
|
||||
}
|
||||
|
||||
fn set_ip_version(&mut self, ip_version: IpVersion) {
|
||||
self.ip_version = ip_version;
|
||||
}
|
||||
|
||||
fn set_resolved_addr(&mut self, addr: SocketAddr) {
|
||||
self.resolved_addr = Some(addr);
|
||||
}
|
||||
|
||||
fn set_socket_mark(&mut self, socket_mark: Option<u32>) {
|
||||
self.socket_mark = socket_mark;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tunnel::{
|
||||
TunnelConnector,
|
||||
common::tests::{_tunnel_bench, _tunnel_pingpong},
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn tcp_pingpong() {
|
||||
let listener = TcpTunnelListener::new("tcp://0.0.0.0:31011".parse().unwrap());
|
||||
let connector = TcpTunnelConnector::new("tcp://127.0.0.1:31011".parse().unwrap());
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tcp_bench() {
|
||||
let listener = TcpTunnelListener::new("tcp://0.0.0.0:31012".parse().unwrap());
|
||||
let connector = TcpTunnelConnector::new("tcp://127.0.0.1:31012".parse().unwrap());
|
||||
_tunnel_bench(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tcp_bench_with_bind() {
|
||||
let listener = TcpTunnelListener::new("tcp://127.0.0.1:11013".parse().unwrap());
|
||||
let mut connector = TcpTunnelConnector::new("tcp://127.0.0.1:11013".parse().unwrap());
|
||||
connector.set_bind_addrs(vec!["127.0.0.1:0".parse().unwrap()]);
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[should_panic]
|
||||
async fn tcp_bench_with_bind_fail() {
|
||||
let listener = TcpTunnelListener::new("tcp://127.0.0.1:11014".parse().unwrap());
|
||||
let mut connector = TcpTunnelConnector::new("tcp://127.0.0.1:11014".parse().unwrap());
|
||||
connector.set_bind_addrs(vec!["10.0.0.1:0".parse().unwrap()]);
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bind_same_port() {
|
||||
let mut listener = TcpTunnelListener::new("tcp://[::]:31014".parse().unwrap());
|
||||
let mut listener2 = TcpTunnelListener::new("tcp://0.0.0.0:31014".parse().unwrap());
|
||||
listener.listen().await.unwrap();
|
||||
listener2.listen().await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ipv6_pingpong() {
|
||||
let listener = TcpTunnelListener::new("tcp://[::1]:31015".parse().unwrap());
|
||||
let connector = TcpTunnelConnector::new("tcp://[::1]:31015".parse().unwrap());
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ipv6_domain_pingpong() {
|
||||
let listener = TcpTunnelListener::new("tcp://[::1]:31015".parse().unwrap());
|
||||
let mut connector =
|
||||
TcpTunnelConnector::new("tcp://test.easytier.top:31015".parse().unwrap());
|
||||
connector.set_ip_version(IpVersion::V6);
|
||||
_tunnel_pingpong(listener, connector).await;
|
||||
|
||||
let listener = TcpTunnelListener::new("tcp://127.0.0.1:31015".parse().unwrap());
|
||||
let mut connector =
|
||||
TcpTunnelConnector::new("tcp://test.easytier.top:31015".parse().unwrap());
|
||||
connector.set_ip_version(IpVersion::V4);
|
||||
_tunnel_pingpong(listener, connector).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn connector_keeps_source_addr_and_reports_resolved_addr() {
|
||||
let mut listener = TcpTunnelListener::new("tcp://127.0.0.1:0".parse().unwrap());
|
||||
listener.listen().await.unwrap();
|
||||
|
||||
let port = listener.local_url().port().unwrap();
|
||||
let source_url: url::Url = format!("tcp://localhost:{port}").parse().unwrap();
|
||||
let mut connector = TcpTunnelConnector::new(source_url.clone());
|
||||
connector.set_ip_version(IpVersion::V4);
|
||||
|
||||
let accept_task = tokio::spawn(async move { listener.accept().await.unwrap() });
|
||||
let tunnel = connector.connect().await.unwrap();
|
||||
let accepted_tunnel = accept_task.await.unwrap();
|
||||
|
||||
let info = tunnel.info().unwrap();
|
||||
assert_eq!(info.remote_addr.unwrap().url, source_url.to_string());
|
||||
|
||||
let resolved_remote_addr: url::Url = info.resolved_remote_addr.unwrap().into();
|
||||
assert_eq!(resolved_remote_addr.host_str(), Some("127.0.0.1"));
|
||||
assert_eq!(resolved_remote_addr.port(), Some(port));
|
||||
|
||||
let accepted_info = accepted_tunnel.info().unwrap();
|
||||
assert_eq!(
|
||||
accepted_info.remote_addr,
|
||||
accepted_info.resolved_remote_addr,
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn connector_uses_pre_resolved_addr_without_resolving_url() {
|
||||
let mut listener = TcpTunnelListener::new("tcp://127.0.0.1:0".parse().unwrap());
|
||||
listener.listen().await.unwrap();
|
||||
|
||||
let port = listener.local_url().port().unwrap();
|
||||
let source_url: url::Url = format!("tcp://unresolvable.invalid:{port}")
|
||||
.parse()
|
||||
.unwrap();
|
||||
let resolved_addr: SocketAddr = format!("127.0.0.1:{port}").parse().unwrap();
|
||||
let mut connector = TcpTunnelConnector::new(source_url.clone());
|
||||
connector.set_resolved_addr(resolved_addr);
|
||||
|
||||
let accept_task = tokio::spawn(async move { listener.accept().await.unwrap() });
|
||||
let tunnel = connector.connect().await.unwrap();
|
||||
let _accepted_tunnel = accept_task.await.unwrap();
|
||||
|
||||
let info = tunnel.info().unwrap();
|
||||
assert_eq!(info.remote_addr.unwrap().url, source_url.to_string());
|
||||
|
||||
let resolved_remote_addr: url::Url = info.resolved_remote_addr.unwrap().into();
|
||||
assert_eq!(resolved_remote_addr.host_str(), Some("127.0.0.1"));
|
||||
assert_eq!(resolved_remote_addr.port(), Some(port));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_alloc_port() {
|
||||
// v4
|
||||
let mut listener = TcpTunnelListener::new("tcp://0.0.0.0:0".parse().unwrap());
|
||||
listener.listen().await.unwrap();
|
||||
let port = listener.local_url().port().unwrap();
|
||||
assert!(port > 0);
|
||||
|
||||
// v6
|
||||
let mut listener = TcpTunnelListener::new("tcp://[::]:0".parse().unwrap());
|
||||
listener.listen().await.unwrap();
|
||||
let port = listener.local_url().port().unwrap();
|
||||
assert!(port > 0);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,210 +0,0 @@
|
||||
use std::{
|
||||
io,
|
||||
net::{Ipv6Addr, SocketAddrV6},
|
||||
};
|
||||
|
||||
use tokio::net::UdpSocket;
|
||||
|
||||
#[cfg(unix)]
|
||||
pub(crate) fn send_to_with_src_ipv6(
|
||||
socket: &UdpSocket,
|
||||
src_ip: Ipv6Addr,
|
||||
src_ifindex: u32,
|
||||
dst_addr: SocketAddrV6,
|
||||
buf: &[u8],
|
||||
) -> io::Result<usize> {
|
||||
#[cfg(target_env = "ohos")]
|
||||
{
|
||||
let _ = (socket, src_ip, src_ifindex, dst_addr, buf);
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::Unsupported,
|
||||
"sending UDP with a selected IPv6 source is not supported on OHOS",
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(not(target_env = "ohos"))]
|
||||
{
|
||||
use std::{mem, os::fd::AsRawFd, ptr};
|
||||
|
||||
use nix::libc;
|
||||
|
||||
#[repr(align(8))]
|
||||
struct ControlBuffer([u8; 128]);
|
||||
|
||||
#[cfg(target_os = "android")]
|
||||
let ipi6_ifindex: libc::c_int = i32::try_from(src_ifindex).map_err(|_| {
|
||||
io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"IPv6 source interface index is out of range",
|
||||
)
|
||||
})?;
|
||||
#[cfg(not(target_os = "android"))]
|
||||
let ipi6_ifindex: libc::c_uint = src_ifindex;
|
||||
|
||||
let pktinfo = libc::in6_pktinfo {
|
||||
ipi6_addr: libc::in6_addr {
|
||||
s6_addr: src_ip.octets(),
|
||||
},
|
||||
ipi6_ifindex,
|
||||
};
|
||||
let mut iov = libc::iovec {
|
||||
iov_base: buf.as_ptr() as *mut libc::c_void,
|
||||
iov_len: buf.len(),
|
||||
};
|
||||
let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr));
|
||||
let control_len = unsafe {
|
||||
libc::CMSG_SPACE(mem::size_of::<libc::in6_pktinfo>() as libc::c_uint) as usize
|
||||
};
|
||||
let mut control = ControlBuffer([0u8; 128]);
|
||||
if control_len > control.0.len() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"IPv6 packet info control buffer is too small",
|
||||
));
|
||||
}
|
||||
|
||||
let mut msg = unsafe { mem::zeroed::<libc::msghdr>() };
|
||||
msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void;
|
||||
msg.msg_namelen = dst_addr.len() as _;
|
||||
msg.msg_iov = &mut iov;
|
||||
msg.msg_iovlen = 1;
|
||||
msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void;
|
||||
msg.msg_controllen = control_len as _;
|
||||
msg.msg_flags = 0;
|
||||
|
||||
unsafe {
|
||||
let cmsg = libc::CMSG_FIRSTHDR(&msg);
|
||||
if cmsg.is_null() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"IPv6 packet info control buffer is invalid",
|
||||
));
|
||||
}
|
||||
(*cmsg).cmsg_level = libc::IPPROTO_IPV6;
|
||||
(*cmsg).cmsg_type = libc::IPV6_PKTINFO;
|
||||
(*cmsg).cmsg_len =
|
||||
libc::CMSG_LEN(mem::size_of::<libc::in6_pktinfo>() as libc::c_uint) as _;
|
||||
ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in6_pktinfo, pktinfo);
|
||||
|
||||
let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0);
|
||||
if ret < 0 {
|
||||
Err(io::Error::last_os_error())
|
||||
} else {
|
||||
Ok(ret as usize)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(windows)]
|
||||
pub(crate) fn send_to_with_src_ipv6(
|
||||
socket: &UdpSocket,
|
||||
src_ip: Ipv6Addr,
|
||||
src_ifindex: u32,
|
||||
dst_addr: SocketAddrV6,
|
||||
buf: &[u8],
|
||||
) -> io::Result<usize> {
|
||||
use std::{mem, os::windows::io::AsRawSocket, ptr};
|
||||
|
||||
use windows::{
|
||||
Win32::Networking::WinSock::{
|
||||
CMSGHDR, IN6_ADDR, IN6_ADDR_0, IN6_PKTINFO, IPPROTO_IPV6, IPV6_PKTINFO, SOCKET,
|
||||
SOCKET_ERROR, WSABUF, WSAGetLastError, WSAMSG, WSASendMsg,
|
||||
},
|
||||
core::PSTR,
|
||||
};
|
||||
|
||||
fn cmsghdr_align(length: usize) -> usize {
|
||||
(length + mem::align_of::<CMSGHDR>() - 1) & !(mem::align_of::<CMSGHDR>() - 1)
|
||||
}
|
||||
|
||||
fn cmsgdata_align(length: usize) -> usize {
|
||||
(length + mem::align_of::<usize>() - 1) & !(mem::align_of::<usize>() - 1)
|
||||
}
|
||||
|
||||
fn cmsg_len(length: usize) -> usize {
|
||||
cmsgdata_align(mem::size_of::<CMSGHDR>()) + length
|
||||
}
|
||||
|
||||
fn cmsg_space(length: usize) -> usize {
|
||||
cmsgdata_align(mem::size_of::<CMSGHDR>() + cmsghdr_align(length))
|
||||
}
|
||||
|
||||
fn cmsg_data(cmsg: *mut CMSGHDR) -> *mut u8 {
|
||||
(cmsg as usize + cmsgdata_align(mem::size_of::<CMSGHDR>())) as *mut u8
|
||||
}
|
||||
|
||||
#[repr(align(8))]
|
||||
struct ControlBuffer([u8; 128]);
|
||||
|
||||
let dst = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr));
|
||||
let mut data = WSABUF {
|
||||
len: buf.len() as u32,
|
||||
buf: PSTR(buf.as_ptr() as *mut u8),
|
||||
};
|
||||
let control_len = cmsg_space(mem::size_of::<IN6_PKTINFO>());
|
||||
let mut control = ControlBuffer([0u8; 128]);
|
||||
if control_len > control.0.len() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"IPv6 packet info control buffer is too small",
|
||||
));
|
||||
}
|
||||
let mut msg = WSAMSG {
|
||||
name: dst.as_ptr() as *mut _,
|
||||
namelen: dst.len(),
|
||||
lpBuffers: &mut data,
|
||||
dwBufferCount: 1,
|
||||
Control: WSABUF {
|
||||
len: control_len as u32,
|
||||
buf: PSTR(control.0.as_mut_ptr()),
|
||||
},
|
||||
dwFlags: 0,
|
||||
};
|
||||
|
||||
let pktinfo = IN6_PKTINFO {
|
||||
ipi6_addr: IN6_ADDR {
|
||||
u: IN6_ADDR_0 {
|
||||
Byte: src_ip.octets(),
|
||||
},
|
||||
},
|
||||
ipi6_ifindex: src_ifindex,
|
||||
};
|
||||
|
||||
unsafe {
|
||||
let cmsg = control.0.as_mut_ptr() as *mut CMSGHDR;
|
||||
(*cmsg).cmsg_level = IPPROTO_IPV6.0;
|
||||
(*cmsg).cmsg_type = IPV6_PKTINFO;
|
||||
(*cmsg).cmsg_len = cmsg_len(mem::size_of::<IN6_PKTINFO>());
|
||||
ptr::write(cmsg_data(cmsg) as *mut IN6_PKTINFO, pktinfo);
|
||||
msg.Control.len = control_len as u32;
|
||||
|
||||
let mut sent = 0;
|
||||
let ret = WSASendMsg(
|
||||
SOCKET(socket.as_raw_socket() as usize),
|
||||
&msg,
|
||||
0,
|
||||
Some(&mut sent),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
if ret == SOCKET_ERROR {
|
||||
return Err(io::Error::from_raw_os_error(WSAGetLastError().0));
|
||||
}
|
||||
Ok(sent as usize)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(any(unix, windows)))]
|
||||
pub(crate) fn send_to_with_src_ipv6(
|
||||
_socket: &UdpSocket,
|
||||
_src_ip: Ipv6Addr,
|
||||
_src_ifindex: u32,
|
||||
_dst_addr: SocketAddrV6,
|
||||
_buf: &[u8],
|
||||
) -> io::Result<usize> {
|
||||
Err(io::Error::new(
|
||||
io::ErrorKind::Unsupported,
|
||||
"sending UDP with a selected IPv6 source is not supported on this platform",
|
||||
))
|
||||
}
|
||||
@@ -1,218 +0,0 @@
|
||||
use std::path::Path;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tokio::net::{UnixListener, UnixStream, unix::SocketAddr};
|
||||
|
||||
use super::TunnelInfo;
|
||||
|
||||
use super::{
|
||||
IpVersion, Tunnel, TunnelError, TunnelListener,
|
||||
common::{FramedReader, FramedWriter, TunnelWrapper},
|
||||
};
|
||||
|
||||
const MAX_PACKET_SIZE: usize = 4096;
|
||||
|
||||
fn url_from_unix_socket_addr(addr: SocketAddr) -> Option<url::Url> {
|
||||
addr.as_pathname()
|
||||
.and_then(|p| p.to_str())
|
||||
.and_then(|s| format!("unix://{}", s).parse().ok())
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct UnixSocketTunnelListener {
|
||||
addr: url::Url,
|
||||
listener: Option<UnixListener>,
|
||||
unlink_on_drop: bool,
|
||||
}
|
||||
|
||||
impl UnixSocketTunnelListener {
|
||||
pub fn new(addr: url::Url) -> Self {
|
||||
UnixSocketTunnelListener {
|
||||
addr,
|
||||
listener: None,
|
||||
unlink_on_drop: true,
|
||||
}
|
||||
}
|
||||
|
||||
async fn do_accept(&self) -> Result<Box<dyn Tunnel>, std::io::Error> {
|
||||
let listener = self.listener.as_ref().unwrap();
|
||||
let (stream, _) = listener.accept().await?;
|
||||
|
||||
let remote_addr = stream.peer_addr().ok().and_then(url_from_unix_socket_addr);
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: "unix".to_owned(),
|
||||
local_addr: Some(self.local_url().into()),
|
||||
remote_addr: remote_addr.clone().map(Into::into),
|
||||
resolved_remote_addr: remote_addr.map(Into::into),
|
||||
};
|
||||
|
||||
let (r, w) = stream.into_split();
|
||||
Ok(Box::new(TunnelWrapper::new(
|
||||
FramedReader::new(r, MAX_PACKET_SIZE),
|
||||
FramedWriter::new(w),
|
||||
Some(info),
|
||||
)))
|
||||
}
|
||||
|
||||
fn set_unlink_on_drop(&mut self, unlink: bool) {
|
||||
self.unlink_on_drop = unlink;
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TunnelListener for UnixSocketTunnelListener {
|
||||
async fn listen(&mut self) -> Result<(), TunnelError> {
|
||||
self.listener = None;
|
||||
let path_str = self.addr.path();
|
||||
let path = Path::new(path_str);
|
||||
|
||||
let listener = UnixListener::bind(path)?;
|
||||
self.listener = Some(listener);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn accept(&mut self) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
loop {
|
||||
match self.do_accept().await {
|
||||
Ok(ret) => return Ok(ret),
|
||||
Err(e) => {
|
||||
use std::io::ErrorKind::*;
|
||||
if matches!(
|
||||
e.kind(),
|
||||
NotConnected | ConnectionAborted | ConnectionRefused | ConnectionReset
|
||||
) {
|
||||
tracing::warn!(?e, "accept fail with retryable error: {:?}", e);
|
||||
continue;
|
||||
}
|
||||
tracing::warn!(?e, "accept fail");
|
||||
return Err(e.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn local_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct UnixSocketTunnelConnector {
|
||||
addr: url::Url,
|
||||
}
|
||||
|
||||
impl UnixSocketTunnelConnector {
|
||||
pub fn new(addr: url::Url) -> Self {
|
||||
UnixSocketTunnelConnector { addr }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl super::TunnelConnector for UnixSocketTunnelConnector {
|
||||
async fn connect(&mut self) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
let path_str = self.addr.path();
|
||||
let path = Path::new(path_str);
|
||||
tracing::info!(url = ?self.addr, "connect unix socket start");
|
||||
let stream = UnixStream::connect(path).await?;
|
||||
tracing::info!(url = ?self.addr, "connect unix socket succ");
|
||||
|
||||
let local_addr = stream.local_addr().ok().and_then(url_from_unix_socket_addr);
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: "unix".to_owned(),
|
||||
local_addr: local_addr.map(Into::into),
|
||||
remote_addr: Some(self.addr.clone().into()),
|
||||
resolved_remote_addr: Some(self.addr.clone().into()),
|
||||
};
|
||||
|
||||
let (r, w) = stream.into_split();
|
||||
Ok(Box::new(TunnelWrapper::new(
|
||||
FramedReader::new(r, MAX_PACKET_SIZE),
|
||||
FramedWriter::new(w),
|
||||
Some(info),
|
||||
)))
|
||||
}
|
||||
|
||||
fn remote_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
|
||||
fn set_ip_version(&mut self, _ip_version: IpVersion) {
|
||||
// IP version is not applicable to UNIX sockets
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for UnixSocketTunnelListener {
|
||||
fn drop(&mut self) {
|
||||
if self.unlink_on_drop {
|
||||
let _ = std::fs::remove_file(self.addr.path());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tunnel::common::tests::{_tunnel_bench, _tunnel_pingpong};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn unix_socket_pingpong() {
|
||||
let listener =
|
||||
UnixSocketTunnelListener::new("unix:///tmp/easytier-test.sock".parse().unwrap());
|
||||
let connector =
|
||||
UnixSocketTunnelConnector::new("unix:///tmp/easytier-test.sock".parse().unwrap());
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unix_socket_bench() {
|
||||
let listener =
|
||||
UnixSocketTunnelListener::new("unix:///tmp/easytier-test-bench.sock".parse().unwrap());
|
||||
let connector =
|
||||
UnixSocketTunnelConnector::new("unix:///tmp/easytier-test-bench.sock".parse().unwrap());
|
||||
_tunnel_bench(listener, connector).await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unlink_on_drop() {
|
||||
let listener =
|
||||
UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap());
|
||||
let connector = UnixSocketTunnelConnector::new(
|
||||
"unix:///tmp/easytier-test-exists.sock".parse().unwrap(),
|
||||
);
|
||||
_tunnel_pingpong(listener, connector).await;
|
||||
|
||||
let mut listener =
|
||||
UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap());
|
||||
listener.set_unlink_on_drop(false);
|
||||
let connector = UnixSocketTunnelConnector::new(
|
||||
"unix:///tmp/easytier-test-exists.sock".parse().unwrap(),
|
||||
);
|
||||
_tunnel_pingpong(listener, connector).await;
|
||||
|
||||
let mut listener =
|
||||
UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap());
|
||||
let result = listener.listen().await;
|
||||
assert!(
|
||||
matches!(result, Err(TunnelError::IOError(err)) if err.kind() == std::io::ErrorKind::AddrInUse)
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bind_file_exists() {
|
||||
use std::fs;
|
||||
|
||||
let path = "/tmp/easytier-test-exists.sock";
|
||||
fs::File::create(path).unwrap();
|
||||
let mut listener =
|
||||
UnixSocketTunnelListener::new("unix:///tmp/easytier-test-exists.sock".parse().unwrap());
|
||||
let result = listener.listen().await;
|
||||
|
||||
fs::remove_file(path).unwrap();
|
||||
assert!(
|
||||
matches!(result, Err(TunnelError::IOError(err)) if err.kind() == std::io::ErrorKind::AddrInUse)
|
||||
)
|
||||
}
|
||||
}
|
||||
+289
-320
@@ -1,81 +1,250 @@
|
||||
use super::{
|
||||
FromUrl, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelListener,
|
||||
common::{TunnelWrapper, wait_for_connect_futures},
|
||||
insecure_tls::{get_insecure_tls_cert, init_crypto_provider},
|
||||
packet_def::{ZCPacket, ZCPacketType},
|
||||
};
|
||||
use super::FromUrl;
|
||||
use crate::tunnel::common::bind;
|
||||
use crate::{proto::common::TunnelInfo, tunnel::insecure_tls::get_insecure_tls_client_config};
|
||||
use anyhow::Context;
|
||||
use crate::{proto::common::TunnelInfo, socket::tcp::RuntimeTcpSocket};
|
||||
use anyhow::Context as _;
|
||||
use bytes::BytesMut;
|
||||
use cidr::IpCidr;
|
||||
use easytier_core::{
|
||||
packet::{ZCPacket, ZCPacketType},
|
||||
socket::tcp::VirtualTcpSocket,
|
||||
tunnel::{IpVersion, Tunnel, TunnelError, wrapper::TunnelWrapper},
|
||||
};
|
||||
use forwarded_header_value::ForwardedHeaderValue;
|
||||
use futures::{SinkExt, StreamExt, stream::FuturesUnordered};
|
||||
use pnet::ipnetwork::IpNetwork;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use std::{
|
||||
net::SocketAddr,
|
||||
net::{IpAddr, SocketAddr},
|
||||
sync::{Arc, LazyLock},
|
||||
time::Duration,
|
||||
};
|
||||
use tokio::{
|
||||
net::{TcpListener, TcpSocket, TcpStream},
|
||||
time::timeout,
|
||||
};
|
||||
use tokio::{net::TcpListener, time::timeout};
|
||||
use tokio_rustls::TlsAcceptor;
|
||||
use tokio_util::either::Either;
|
||||
use tokio_websockets::{ClientBuilder, Limits, MaybeTlsStream, Message, ServerBuilder};
|
||||
use zerocopy::AsBytes;
|
||||
use zerocopy::AsBytes as _;
|
||||
|
||||
fn is_wss(addr: &url::Url) -> Result<bool, TunnelError> {
|
||||
match addr.scheme() {
|
||||
pub(crate) const CONNECT_TIMEOUT: Duration = Duration::from_secs(20);
|
||||
pub(crate) const SERVER_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(3);
|
||||
|
||||
static TRUSTED_PROXIES: LazyLock<Vec<IpCidr>> = LazyLock::new(|| {
|
||||
[
|
||||
"127.0.0.0/8",
|
||||
"10.0.0.0/8",
|
||||
"172.16.0.0/12",
|
||||
"192.168.0.0/16",
|
||||
"::1/128",
|
||||
"fc00::/7",
|
||||
]
|
||||
.into_iter()
|
||||
.map(|cidr| cidr.parse().unwrap())
|
||||
.collect()
|
||||
});
|
||||
|
||||
fn trusted_proxy_contains(ip: IpAddr) -> bool {
|
||||
TRUSTED_PROXIES.iter().any(|cidr| match (cidr, ip) {
|
||||
(IpCidr::V4(cidr), IpAddr::V4(ip)) => cidr.contains(&ip),
|
||||
(IpCidr::V6(cidr), IpAddr::V6(ip)) => cidr.contains(&ip),
|
||||
_ => false,
|
||||
})
|
||||
}
|
||||
|
||||
fn websocket_error(error: impl std::fmt::Display) -> TunnelError {
|
||||
TunnelError::ProtocolError(format!("websocket error: {error}"))
|
||||
}
|
||||
|
||||
fn is_wss(url: &url::Url) -> Result<bool, TunnelError> {
|
||||
match url.scheme() {
|
||||
"ws" => Ok(false),
|
||||
"wss" => Ok(true),
|
||||
_ => Err(TunnelError::InvalidProtocol(addr.scheme().to_string())),
|
||||
scheme => Err(TunnelError::InvalidProtocol(scheme.to_owned())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn sink_from_zc_packet<E>(msg: ZCPacket) -> Result<Message, E> {
|
||||
Ok(Message::binary(msg.tunnel_payload_bytes().freeze()))
|
||||
async fn sink_from_zc_packet<E>(packet: ZCPacket) -> Result<Message, E> {
|
||||
Ok(Message::binary(packet.tunnel_payload_bytes().freeze()))
|
||||
}
|
||||
|
||||
async fn map_from_ws_message(
|
||||
msg: Result<Message, tokio_websockets::Error>,
|
||||
message: Result<Message, tokio_websockets::Error>,
|
||||
) -> Option<Result<ZCPacket, TunnelError>> {
|
||||
if let Err(e) = msg {
|
||||
tracing::error!(?e, "recv from websocket error");
|
||||
return Some(Err(TunnelError::WebSocketError(e)));
|
||||
}
|
||||
|
||||
let msg = msg.unwrap();
|
||||
if msg.is_close() {
|
||||
let message = match message {
|
||||
Ok(message) => message,
|
||||
Err(error) => {
|
||||
tracing::error!(?error, "recv from websocket error");
|
||||
return Some(Err(websocket_error(error)));
|
||||
}
|
||||
};
|
||||
if message.is_close() {
|
||||
tracing::warn!("recv close message from websocket");
|
||||
return None;
|
||||
}
|
||||
|
||||
if !msg.is_binary() {
|
||||
let msg = format!("{:?}", msg);
|
||||
tracing::error!(?msg, "Invalid packet");
|
||||
return Some(Err(TunnelError::InvalidPacket(msg)));
|
||||
if !message.is_binary() {
|
||||
let message = format!("{message:?}");
|
||||
tracing::error!(?message, "Invalid packet");
|
||||
return Some(Err(TunnelError::InvalidPacket(message)));
|
||||
}
|
||||
|
||||
Some(Ok(ZCPacket::new_from_buf(
|
||||
BytesMut::from(msg.into_payload().as_bytes()),
|
||||
BytesMut::from(message.into_payload().as_bytes()),
|
||||
ZCPacketType::DummyTunnel,
|
||||
)))
|
||||
}
|
||||
|
||||
static TRUSTED_PROXIES: LazyLock<Vec<IpNetwork>> = LazyLock::new(|| {
|
||||
[
|
||||
"127.0.0.0/8", // IPV4 Loopback
|
||||
"10.0.0.0/8", // IPV4 Private Networks
|
||||
"172.16.0.0/12",
|
||||
"192.168.0.0/16",
|
||||
"::1/128", // IPV6 Loopback
|
||||
"fc00::/7", // IPV6 Private network
|
||||
]
|
||||
.into_iter()
|
||||
.map(|s| s.parse().unwrap())
|
||||
.collect()
|
||||
});
|
||||
#[derive(Debug)]
|
||||
struct SkipServerVerification(Arc<rustls::crypto::CryptoProvider>);
|
||||
|
||||
impl SkipServerVerification {
|
||||
fn new(provider: Arc<rustls::crypto::CryptoProvider>) -> Arc<Self> {
|
||||
Arc::new(Self(provider))
|
||||
}
|
||||
}
|
||||
|
||||
impl rustls::client::danger::ServerCertVerifier for SkipServerVerification {
|
||||
fn verify_server_cert(
|
||||
&self,
|
||||
_end_entity: &rustls::pki_types::CertificateDer<'_>,
|
||||
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
|
||||
_server_name: &rustls::pki_types::ServerName<'_>,
|
||||
_ocsp: &[u8],
|
||||
_now: rustls::pki_types::UnixTime,
|
||||
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
|
||||
Ok(rustls::client::danger::ServerCertVerified::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls12_signature(
|
||||
&self,
|
||||
message: &[u8],
|
||||
cert: &rustls::pki_types::CertificateDer<'_>,
|
||||
dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
|
||||
rustls::crypto::verify_tls12_signature(
|
||||
message,
|
||||
cert,
|
||||
dss,
|
||||
&self.0.signature_verification_algorithms,
|
||||
)
|
||||
}
|
||||
|
||||
fn verify_tls13_signature(
|
||||
&self,
|
||||
message: &[u8],
|
||||
cert: &rustls::pki_types::CertificateDer<'_>,
|
||||
dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
|
||||
rustls::crypto::verify_tls13_signature(
|
||||
message,
|
||||
cert,
|
||||
dss,
|
||||
&self.0.signature_verification_algorithms,
|
||||
)
|
||||
}
|
||||
|
||||
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
|
||||
self.0.signature_verification_algorithms.supported_schemes()
|
||||
}
|
||||
}
|
||||
|
||||
fn init_crypto_provider() {
|
||||
let _ =
|
||||
rustls::crypto::CryptoProvider::install_default(rustls::crypto::ring::default_provider());
|
||||
}
|
||||
|
||||
fn get_insecure_tls_client_config() -> rustls::ClientConfig {
|
||||
init_crypto_provider();
|
||||
let provider = rustls::crypto::CryptoProvider::get_default().unwrap();
|
||||
let mut config = rustls::ClientConfig::builder()
|
||||
.dangerous()
|
||||
.with_custom_certificate_verifier(SkipServerVerification::new(provider.clone()))
|
||||
.with_no_client_auth();
|
||||
config.enable_sni = true;
|
||||
config.enable_early_data = false;
|
||||
config
|
||||
}
|
||||
|
||||
fn get_insecure_tls_cert<'a>() -> (
|
||||
Vec<rustls::pki_types::CertificateDer<'a>>,
|
||||
rustls::pki_types::PrivateKeyDer<'a>,
|
||||
) {
|
||||
let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
|
||||
let cert_der = cert.serialize_der().unwrap();
|
||||
let private_key = cert.serialize_private_key_der();
|
||||
let private_key = rustls::pki_types::PrivatePkcs8KeyDer::from(private_key);
|
||||
(vec![cert_der.into()], private_key.into())
|
||||
}
|
||||
|
||||
pub(crate) async fn upgrade_accepted<S>(
|
||||
stream: S,
|
||||
local_url: url::Url,
|
||||
) -> Result<Box<dyn Tunnel>, TunnelError>
|
||||
where
|
||||
S: VirtualTcpSocket,
|
||||
{
|
||||
let peer_addr = stream.peer_addr()?;
|
||||
let mut remote_url = socket_url(local_url.scheme(), peer_addr);
|
||||
let stream = if is_wss(&local_url)? {
|
||||
init_crypto_provider();
|
||||
let (certificates, private_key) = get_insecure_tls_cert();
|
||||
let config = rustls::ServerConfig::builder()
|
||||
.with_no_client_auth()
|
||||
.with_single_cert(certificates, private_key)
|
||||
.with_context(|| "Failed to create server config")?;
|
||||
Either::Left(TlsAcceptor::from(Arc::new(config)).accept(stream).await?)
|
||||
} else {
|
||||
Either::Right(stream)
|
||||
};
|
||||
|
||||
let (request, stream) = ServerBuilder::new()
|
||||
.limits(Limits::unlimited())
|
||||
.max_headers(128)
|
||||
.accept(stream)
|
||||
.await
|
||||
.map_err(websocket_error)?;
|
||||
|
||||
if trusted_proxy_contains(peer_addr.ip())
|
||||
&& let Some(forwarded) = request
|
||||
.headers()
|
||||
.get("Forwarded")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| ForwardedHeaderValue::from_forwarded(value).ok())
|
||||
.or_else(|| {
|
||||
request
|
||||
.headers()
|
||||
.get("X-Forwarded-For")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| ForwardedHeaderValue::from_x_forwarded_for(value).ok())
|
||||
})
|
||||
&& let Some(ip) = forwarded.remotest_forwarded_for_ip()
|
||||
{
|
||||
remote_url
|
||||
.set_host(Some(&ip.to_string()))
|
||||
.map_err(|_| TunnelError::InvalidAddr(format!("invalid forwarded ip {ip}")))?;
|
||||
remote_url
|
||||
.query_pairs_mut()
|
||||
.append_pair("proxy", &peer_addr.to_string());
|
||||
}
|
||||
|
||||
let (write, read) = stream.split();
|
||||
let remote_url: crate::proto::common::Url = remote_url.into();
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: local_url.scheme().to_owned(),
|
||||
local_addr: Some(local_url.into()),
|
||||
remote_addr: Some(remote_url.clone()),
|
||||
resolved_remote_addr: Some(remote_url),
|
||||
};
|
||||
Ok(Box::new(TunnelWrapper::new(
|
||||
read.filter_map(map_from_ws_message),
|
||||
write
|
||||
.sink_map_err(websocket_error)
|
||||
.with(sink_from_zc_packet::<TunnelError>),
|
||||
Some(info),
|
||||
)))
|
||||
}
|
||||
|
||||
fn socket_url(scheme: &str, addr: SocketAddr) -> url::Url {
|
||||
let mut url = url::Url::parse(&format!("{scheme}://0.0.0.0"))
|
||||
.expect("WebSocket transport scheme should be a valid URL scheme");
|
||||
url.set_ip_host(addr.ip()).unwrap();
|
||||
url.set_port(Some(addr.port())).unwrap();
|
||||
url
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct WsTunnelListener {
|
||||
@@ -97,77 +266,7 @@ impl WsTunnelListener {
|
||||
self.socket_mark = socket_mark;
|
||||
}
|
||||
|
||||
async fn try_accept(&self, stream: TcpStream) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
let peer_addr = stream.peer_addr()?;
|
||||
let mut remote_addr =
|
||||
super::build_url_from_socket_addr(&peer_addr.to_string(), self.addr.scheme());
|
||||
|
||||
let stream = if is_wss(&self.addr)? {
|
||||
init_crypto_provider();
|
||||
let (certs, key) = get_insecure_tls_cert();
|
||||
let config = rustls::ServerConfig::builder()
|
||||
.with_no_client_auth()
|
||||
.with_single_cert(certs, key)
|
||||
.with_context(|| "Failed to create server config")?;
|
||||
|
||||
let stream = TlsAcceptor::from(Arc::new(config)).accept(stream).await?;
|
||||
Either::Left(stream)
|
||||
} else {
|
||||
Either::Right(stream)
|
||||
};
|
||||
|
||||
let (request, stream) = ServerBuilder::new()
|
||||
.limits(Limits::unlimited())
|
||||
.max_headers(128)
|
||||
.accept(stream)
|
||||
.await?;
|
||||
|
||||
if TRUSTED_PROXIES
|
||||
.iter()
|
||||
.any(|net| net.contains(peer_addr.ip()))
|
||||
&& let Some(forwarded) = request
|
||||
.headers()
|
||||
.get("Forwarded")
|
||||
.and_then(|f| f.to_str().ok())
|
||||
.and_then(|f| ForwardedHeaderValue::from_forwarded(f).ok())
|
||||
.or_else(|| {
|
||||
request
|
||||
.headers()
|
||||
.get("X-Forwarded-For")
|
||||
.and_then(|f| f.to_str().ok())
|
||||
.and_then(|f| ForwardedHeaderValue::from_x_forwarded_for(f).ok())
|
||||
})
|
||||
&& let Some(ip) = forwarded.remotest_forwarded_for_ip()
|
||||
{
|
||||
remote_addr
|
||||
.set_host(Some(&ip.to_string()))
|
||||
.map_err(|_| TunnelError::InvalidAddr(format!("invalid forwarded ip {}", ip)))?;
|
||||
remote_addr
|
||||
.query_pairs_mut()
|
||||
.append_pair("proxy", &peer_addr.to_string());
|
||||
}
|
||||
|
||||
let (write, read) = stream.split();
|
||||
let remote_addr: crate::proto::common::Url = remote_addr.into();
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: self.addr.scheme().to_owned(),
|
||||
local_addr: Some(self.local_url().into()),
|
||||
remote_addr: Some(remote_addr.clone()),
|
||||
resolved_remote_addr: Some(remote_addr),
|
||||
};
|
||||
|
||||
Ok(Box::new(TunnelWrapper::new(
|
||||
read.filter_map(map_from_ws_message),
|
||||
write.with(sink_from_zc_packet),
|
||||
Some(info),
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl TunnelListener for WsTunnelListener {
|
||||
async fn listen(&mut self) -> Result<(), TunnelError> {
|
||||
async fn listen_tunnel(&mut self) -> Result<(), TunnelError> {
|
||||
self.listener = None;
|
||||
|
||||
let addr = SocketAddr::from_url(self.addr.clone(), IpVersion::Both).await?;
|
||||
@@ -185,13 +284,18 @@ impl TunnelListener for WsTunnelListener {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn accept(&mut self) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
async fn accept_tunnel(&mut self) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
loop {
|
||||
let listener = self.listener.as_ref().unwrap();
|
||||
// only fail on tcp accept error
|
||||
let (stream, _) = listener.accept().await?;
|
||||
stream.set_nodelay(true).unwrap();
|
||||
match timeout(Duration::from_secs(3), self.try_accept(stream)).await {
|
||||
match timeout(
|
||||
SERVER_HANDSHAKE_TIMEOUT,
|
||||
upgrade_accepted(RuntimeTcpSocket::new(stream), self.addr.clone()),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(tunnel)) => return Ok(tunnel),
|
||||
e => {
|
||||
tracing::error!(?e, ?self, "Failed to accept ws/wss tunnel");
|
||||
@@ -200,217 +304,82 @@ impl TunnelListener for WsTunnelListener {
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl easytier_core::socket::SocketListener for WsTunnelListener {
|
||||
type Accepted = Box<dyn Tunnel>;
|
||||
|
||||
async fn listen(&mut self) -> anyhow::Result<()> {
|
||||
Ok(self.listen_tunnel().await?)
|
||||
}
|
||||
|
||||
async fn accept(&mut self) -> anyhow::Result<Self::Accepted> {
|
||||
Ok(self.accept_tunnel().await?)
|
||||
}
|
||||
|
||||
fn local_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
}
|
||||
|
||||
pub struct WsTunnelConnector {
|
||||
addr: url::Url,
|
||||
ip_version: IpVersion,
|
||||
resolved_addr: Option<SocketAddr>,
|
||||
pub(crate) async fn upgrade_connected<S>(
|
||||
stream: S,
|
||||
remote_url: url::Url,
|
||||
) -> Result<Box<dyn Tunnel>, TunnelError>
|
||||
where
|
||||
S: VirtualTcpSocket,
|
||||
{
|
||||
let is_wss = is_wss(&remote_url)?;
|
||||
let local_addr = stream.local_addr()?;
|
||||
let resolved_remote_addr = stream.peer_addr()?;
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: remote_url.scheme().to_owned(),
|
||||
local_addr: Some(
|
||||
super::build_url_from_socket_addr(&local_addr.to_string(), remote_url.scheme()).into(),
|
||||
),
|
||||
remote_addr: Some(remote_url.clone().into()),
|
||||
resolved_remote_addr: Some(
|
||||
super::build_url_from_socket_addr(
|
||||
&resolved_remote_addr.to_string(),
|
||||
remote_url.scheme(),
|
||||
)
|
||||
.into(),
|
||||
),
|
||||
};
|
||||
|
||||
bind_addrs: Vec<SocketAddr>,
|
||||
socket_mark: Option<u32>,
|
||||
}
|
||||
let client = ClientBuilder::from_uri(http::Uri::try_from(remote_url.to_string()).unwrap())
|
||||
.max_headers(128);
|
||||
let stream: MaybeTlsStream<S> = if is_wss {
|
||||
init_crypto_provider();
|
||||
let tls = tokio_rustls::TlsConnector::from(Arc::new(get_insecure_tls_client_config()));
|
||||
let sni = remote_url.domain().unwrap_or("localhost").to_owned();
|
||||
let server_name = rustls::pki_types::ServerName::try_from(sni)
|
||||
.map_err(|_| TunnelError::InvalidProtocol("Invalid SNI".to_owned()))?;
|
||||
MaybeTlsStream::Rustls(tls.connect(server_name, stream).await?)
|
||||
} else {
|
||||
MaybeTlsStream::Plain(stream)
|
||||
};
|
||||
|
||||
impl WsTunnelConnector {
|
||||
pub fn new(addr: url::Url) -> Self {
|
||||
WsTunnelConnector {
|
||||
addr,
|
||||
ip_version: IpVersion::Both,
|
||||
resolved_addr: None,
|
||||
|
||||
bind_addrs: vec![],
|
||||
socket_mark: None,
|
||||
}
|
||||
}
|
||||
|
||||
async fn connect_with(
|
||||
addr: url::Url,
|
||||
socket_addr: SocketAddr,
|
||||
tcp_socket: TcpSocket,
|
||||
) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
let is_wss = is_wss(&addr)?;
|
||||
let stream = tcp_socket.connect(socket_addr).await?;
|
||||
if let Err(error) = stream.set_nodelay(true) {
|
||||
tracing::warn!(?error, "set_nodelay fail in ws connect");
|
||||
}
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: addr.scheme().to_owned(),
|
||||
local_addr: Some(
|
||||
super::build_url_from_socket_addr(
|
||||
&stream.local_addr()?.to_string(),
|
||||
addr.scheme().to_string().as_str(),
|
||||
)
|
||||
.into(),
|
||||
),
|
||||
remote_addr: Some(addr.clone().into()),
|
||||
resolved_remote_addr: Some(
|
||||
super::build_url_from_socket_addr(&socket_addr.to_string(), addr.scheme()).into(),
|
||||
),
|
||||
};
|
||||
|
||||
let c = ClientBuilder::from_uri(http::Uri::try_from(addr.to_string()).unwrap())
|
||||
.max_headers(128);
|
||||
let stream: MaybeTlsStream<TcpStream> = if is_wss {
|
||||
init_crypto_provider();
|
||||
let tls_conn =
|
||||
tokio_rustls::TlsConnector::from(Arc::new(get_insecure_tls_client_config()));
|
||||
// Modify SNI logic: use "localhost" as SNI for url without domain to avoid IP blocking.
|
||||
let sni = match addr.domain() {
|
||||
None => "localhost".to_string(),
|
||||
Some(domain) => domain.to_string(),
|
||||
};
|
||||
let server_name = rustls::pki_types::ServerName::try_from(sni)
|
||||
.map_err(|_| TunnelError::InvalidProtocol("Invalid SNI".to_string()))?;
|
||||
let stream = tls_conn.connect(server_name, stream).await?;
|
||||
MaybeTlsStream::Rustls(stream)
|
||||
} else {
|
||||
MaybeTlsStream::Plain(stream)
|
||||
};
|
||||
|
||||
let (client, _) = c.connect_on(stream).await?;
|
||||
let (write, read) = client.split();
|
||||
let read = read.filter_map(map_from_ws_message);
|
||||
let write = write.with(sink_from_zc_packet);
|
||||
Ok(Box::new(TunnelWrapper::new(read, write, Some(info))))
|
||||
}
|
||||
|
||||
async fn connect_with_default_bind(
|
||||
&self,
|
||||
addr: SocketAddr,
|
||||
) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
let socket = if addr.is_ipv4() {
|
||||
TcpSocket::new_v4()?
|
||||
} else {
|
||||
TcpSocket::new_v6()?
|
||||
};
|
||||
crate::tunnel::common::apply_socket_mark(
|
||||
&socket2::SockRef::from(&socket),
|
||||
self.socket_mark,
|
||||
)?;
|
||||
Self::connect_with(self.addr.clone(), addr, socket).await
|
||||
}
|
||||
|
||||
async fn connect_with_custom_bind(
|
||||
&self,
|
||||
addr: SocketAddr,
|
||||
) -> Result<Box<dyn Tunnel>, super::TunnelError> {
|
||||
let futures = FuturesUnordered::new();
|
||||
|
||||
for bind_addr in self.bind_addrs.iter() {
|
||||
tracing::info!(?bind_addr, ?addr, "bind addr");
|
||||
match bind()
|
||||
.addr(*bind_addr)
|
||||
.only_v6(true)
|
||||
.maybe_socket_mark(self.socket_mark)
|
||||
.call()
|
||||
{
|
||||
Ok(socket) => futures.push(Self::connect_with(self.addr.clone(), addr, socket)),
|
||||
Err(error) => {
|
||||
tracing::error!(?bind_addr, ?addr, ?error, "bind addr fail");
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
wait_for_connect_futures(futures).await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl TunnelConnector for WsTunnelConnector {
|
||||
async fn connect(&mut self) -> Result<Box<dyn Tunnel>, TunnelError> {
|
||||
let addr = match self.resolved_addr {
|
||||
Some(addr) => addr,
|
||||
None => SocketAddr::from_url(self.addr.clone(), self.ip_version).await?,
|
||||
};
|
||||
if self.bind_addrs.is_empty() || addr.is_ipv6() {
|
||||
self.connect_with_default_bind(addr).await
|
||||
} else {
|
||||
self.connect_with_custom_bind(addr).await
|
||||
}
|
||||
}
|
||||
|
||||
fn remote_url(&self) -> url::Url {
|
||||
self.addr.clone()
|
||||
}
|
||||
|
||||
fn set_ip_version(&mut self, ip_version: IpVersion) {
|
||||
self.ip_version = ip_version;
|
||||
}
|
||||
|
||||
fn set_bind_addrs(&mut self, addrs: Vec<SocketAddr>) {
|
||||
self.bind_addrs = addrs;
|
||||
}
|
||||
|
||||
fn set_resolved_addr(&mut self, addr: SocketAddr) {
|
||||
self.resolved_addr = Some(addr);
|
||||
}
|
||||
|
||||
fn set_socket_mark(&mut self, socket_mark: Option<u32>) {
|
||||
self.socket_mark = socket_mark;
|
||||
}
|
||||
let (client, _) = client.connect_on(stream).await.map_err(websocket_error)?;
|
||||
let (write, read) = client.split();
|
||||
Ok(Box::new(TunnelWrapper::new(
|
||||
read.filter_map(map_from_ws_message),
|
||||
write
|
||||
.sink_map_err(websocket_error)
|
||||
.with(sink_from_zc_packet::<TunnelError>),
|
||||
Some(info),
|
||||
)))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub mod tests {
|
||||
use super::*;
|
||||
use crate::tunnel::common::tests::_tunnel_pingpong;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
#[serial_test::serial]
|
||||
async fn ws_pingpong(#[values("ws", "wss")] proto: &str) {
|
||||
let listener = WsTunnelListener::new(format!("{}://0.0.0.0:25556", proto).parse().unwrap());
|
||||
let connector =
|
||||
WsTunnelConnector::new(format!("{}://127.0.0.1:25556", proto).parse().unwrap());
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
#[serial_test::serial]
|
||||
async fn ws_pingpong_bind(#[values("ws", "wss")] proto: &str) {
|
||||
let listener = WsTunnelListener::new(format!("{}://0.0.0.0:25557", proto).parse().unwrap());
|
||||
let mut connector =
|
||||
WsTunnelConnector::new(format!("{}://127.0.0.1:25557", proto).parse().unwrap());
|
||||
connector.set_bind_addrs(vec!["127.0.0.1:0".parse().unwrap()]);
|
||||
_tunnel_pingpong(listener, connector).await
|
||||
}
|
||||
|
||||
// TODO: tokio-websockets cannot correctly handle close, benchmark case is disabled
|
||||
// #[rstest::rstest]
|
||||
// #[tokio::test]
|
||||
// #[serial_test::serial]
|
||||
// async fn ws_bench(#[values("ws", "wss")] proto: &str) {
|
||||
// enable_log();
|
||||
// let listener = WSTunnelListener::new(format!("{}://0.0.0.0:25557", proto).parse().unwrap());
|
||||
// let connector =
|
||||
// WSTunnelConnector::new(format!("{}://127.0.0.1:25557", proto).parse().unwrap());
|
||||
// _tunnel_bench(listener, connector).await
|
||||
// }
|
||||
|
||||
#[tokio::test]
|
||||
async fn ws_accept_wss() {
|
||||
let mut listener = WsTunnelListener::new("wss://0.0.0.0:25558".parse().unwrap());
|
||||
listener.listen().await.unwrap();
|
||||
let j = tokio::spawn(async move {
|
||||
let _ = listener.accept().await;
|
||||
});
|
||||
|
||||
let mut connector = WsTunnelConnector::new("ws://127.0.0.1:25558".parse().unwrap());
|
||||
connector.connect().await.unwrap_err();
|
||||
|
||||
let mut connector = WsTunnelConnector::new("wss://127.0.0.1:25558".parse().unwrap());
|
||||
connector.connect().await.unwrap();
|
||||
|
||||
j.abort();
|
||||
}
|
||||
use easytier_core::socket::SocketListener;
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
net::TcpSocket,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn ws_forwarded() {
|
||||
|
||||
+342
-427
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user