[mio] 使用mio改写ip代理

This commit is contained in:
lubeilin
2024-02-29 22:26:53 +08:00
parent ff665d7dcf
commit fe499d0476
4 changed files with 789 additions and 508 deletions
+159 -90
View File
@@ -1,137 +1,206 @@
use std::io;
use std::mem::MaybeUninit;
use std::net::{IpAddr, Ipv4Addr, SocketAddrV4};
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use socket2::{Domain, SockAddr, Socket, Type};
use mio::net::UdpSocket;
use mio::{Events, Interest, Poll, Token, Waker};
use parking_lot::Mutex;
use packet::icmp::icmp;
use packet::icmp::icmp::HeaderOther;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::sender::ChannelSender;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::CurrentDeviceInfo;
use crate::ip_proxy::{send, ProxyHandler};
use crate::ip_proxy::ProxyHandler;
use crate::protocol;
use crate::protocol::{NetPacket, Version, MAX_TTL};
use crate::util::StopManager;
#[derive(Clone)]
pub struct IcmpProxy {
icmp_socket: Arc<Socket>,
icmp_socket: Arc<std::net::UdpSocket>,
// 对端-> 真实来源
icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
sender: ChannelSender,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
nat_map: Arc<Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>>,
}
impl IcmpProxy {
pub fn new(
addr: SocketAddrV4,
icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
sender: ChannelSender,
context: Context,
stop_manager: StopManager,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<IcmpProxy> {
let icmp_socket = Arc::new(Socket::new(
Domain::IPV4,
Type::RAW,
) -> io::Result<Self> {
let icmp_socket = socket2::Socket::new(
socket2::Domain::IPV4,
socket2::Type::RAW,
Some(socket2::Protocol::ICMPV4),
)?);
icmp_socket.bind(&SockAddr::from(addr))?;
Ok(IcmpProxy {
icmp_socket,
icmp_proxy_map,
sender,
current_device,
client_cipher,
)?;
let addr: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
icmp_socket.bind(&socket2::SockAddr::from(addr))?;
icmp_socket.set_nonblocking(true)?;
let std_socket: std::net::UdpSocket = icmp_socket.into();
let mio_icmp_socket = UdpSocket::from_std(std_socket.try_clone()?);
let nat_map: Arc<Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>> =
Arc::new(Mutex::new(HashMap::with_capacity(16)));
{
let nat_map = nat_map.clone();
thread::spawn(move || {
if let Err(e) = icmp_proxy(
mio_icmp_socket,
nat_map,
context,
stop_manager,
current_device,
client_cipher,
) {
log::warn!("icmp_proxy:{:?}", e);
}
});
}
Ok(Self {
icmp_socket: Arc::new(std_socket),
nat_map,
})
}
pub fn icmp_handler(&self) -> IcmpHandler {
IcmpHandler(self.icmp_socket.clone(), self.icmp_proxy_map.clone())
}
pub fn start(self) {
let mut buf = [0u8; 4096];
let data: &mut [MaybeUninit<u8>] = unsafe { std::mem::transmute(&mut buf[12..]) };
}
loop {
match self.recv(data) {
Ok((len, peer_ip)) => {
match peer_ip {
IpAddr::V4(peer_ip) => {
match IpV4Packet::new(&mut buf[12..12 + len]) {
Ok(mut ipv4_packet) => {
match icmp::IcmpPacket::new(ipv4_packet.payload()) {
Ok(icmp_packet) => {
match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => {
if let Some(entry) =
self.icmp_proxy_map.get(&(peer_ip, id, seq))
{
//将数据发送到真实的来源
let dest_ip = *entry.value();
drop(entry);
ipv4_packet.set_destination_ip(dest_ip);
ipv4_packet.update_checksum();
send(
&mut buf,
len,
dest_ip,
&self.sender,
&self.current_device,
&self.client_cipher,
);
}
}
_ => {
continue;
}
}
}
Err(_) => {}
};
}
Err(_) => {}
const SERVER_VAL: usize = 0;
const SERVER: Token = Token(SERVER_VAL);
const NOTIFY_VAL: usize = 1;
const NOTIFY: Token = Token(NOTIFY_VAL);
fn icmp_proxy(
mut icmp_socket: UdpSocket,
// 对端-> 真实来源
nat_map: Arc<Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>>,
context: Context,
stop_manager: StopManager,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<()> {
let mut poll = Poll::new()?;
poll.registry()
.register(&mut icmp_socket, SERVER, Interest::READABLE)?;
let mut events = Events::with_capacity(32);
let stop = Waker::new(poll.registry(), NOTIFY)?;
let _worker = stop_manager.add_listener("icmp_proxy".into(), move || {
if let Err(e) = stop.wake() {
log::warn!("stop icmp_proxy:{:?}", e);
}
})?;
let mut buf = [0u8; 65535 - 20 - 8];
loop {
poll.poll(&mut events, None)?;
for event in events.iter() {
match event.token() {
SERVER => loop {
let (len, addr) = match icmp_socket.recv_from(&mut buf[12..]) {
Ok(rs) => rs,
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::warn!("icmp_socket {:?}", e);
break;
}
IpAddr::V6(_) => {}
};
if let IpAddr::V4(peer_ip) = addr.ip() {
recv_handle(
&mut buf,
12 + len,
peer_ip,
&nat_map,
&context,
&current_device,
&client_cipher,
);
}
},
NOTIFY => {
return Ok(());
}
Err(e) => {
log::warn!("icmp代理异常:{:?}", e);
}
_ => {}
}
}
}
fn recv(&self, buf: &mut [MaybeUninit<u8>]) -> io::Result<(usize, IpAddr)> {
let (size, addr) = self.icmp_socket.recv_from(buf)?;
let addr = match addr.as_socket() {
None => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
Some(add) => add.ip(),
};
Ok((size, addr))
}
fn recv_handle(
buf: &mut [u8],
data_len: usize,
peer_ip: Ipv4Addr,
nat_map: &Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
client_cipher: &Cipher,
) {
match IpV4Packet::new(&mut buf[12..data_len]) {
Ok(mut ipv4_packet) => match icmp::IcmpPacket::new(ipv4_packet.payload()) {
Ok(icmp_packet) => match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => {
if let Some(dest_ip) = nat_map.lock().remove(&(peer_ip, id, seq)) {
ipv4_packet.set_destination_ip(dest_ip);
ipv4_packet.update_checksum();
let current_device = current_device.load();
let virtual_ip = current_device.virtual_ip();
let mut net_packet = NetPacket::new0(data_len, buf).unwrap();
net_packet.set_version(Version::V1);
net_packet.set_protocol(protocol::Protocol::IpTurn);
net_packet.set_transport_protocol(
protocol::ip_turn_packet::Protocol::Ipv4.into(),
);
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip);
net_packet.set_destination(dest_ip);
if let Err(e) = client_cipher.encrypt_ipv4(&mut net_packet) {
log::warn!("加密失败:{}", e);
return;
}
if context.send_by_id(net_packet.buffer(), &dest_ip).is_err() {
let connect_server = current_device.connect_server;
if let Err(e) =
context.send_default(net_packet.buffer(), connect_server)
{
log::warn!("发送到目标失败:{},{}", e, connect_server);
}
}
}
}
_ => {}
},
Err(_) => {}
},
Err(_) => {}
}
}
/// icmp用Identifier来区分,没有Identifier的一律不转发
#[derive(Clone)]
pub struct IcmpHandler(Arc<Socket>, Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>);
impl ProxyHandler for IcmpHandler {
/// icmp用Identifier来区分,没有Identifier的一律不转发
impl ProxyHandler for IcmpProxy {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
source: Ipv4Addr,
destination: Ipv4Addr,
) -> io::Result<bool> {
if ipv4.offset() != 0 || ipv4.flags() & 1 == 1 {
// ip分片的直接丢弃
return Ok(true);
}
let dest_ip = ipv4.destination_ip();
//转发到代理目标地址
let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?;
match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => {
self.1.insert((dest_ip, id, seq), source);
self.0.send_to(
self.nat_map.lock().insert((dest_ip, id, seq), source);
self.icmp_socket.send_to(
ipv4.payload(),
&SockAddr::from(SocketAddrV4::new(dest_ip, 0)),
SocketAddr::from(SocketAddrV4::new(dest_ip, 0)),
)?;
}
_ => {
+27 -152
View File
@@ -1,35 +1,24 @@
use std::io;
use std::net::Ipv4Addr;
use std::net::SocketAddrV4;
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use packet::ip::ipv4;
use tokio::net::UdpSocket;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::sender::ChannelSender;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::CurrentDeviceInfo;
#[cfg(not(target_os = "android"))]
use crate::ip_proxy::icmp_proxy::IcmpHandler;
use crate::ip_proxy::tcp_proxy::{TcpHandler, TcpProxy};
use crate::ip_proxy::udp_proxy::{UdpHandler, UdpProxy};
use crate::protocol;
use crate::protocol::{NetPacket, Version, MAX_TTL};
use crate::ip_proxy::icmp_proxy::IcmpProxy;
use crate::ip_proxy::tcp_proxy::TcpProxy;
use crate::ip_proxy::udp_proxy::UdpProxy;
use crate::util::{Scheduler, StopManager};
#[cfg(not(target_os = "android"))]
pub mod icmp_proxy;
pub mod tcp_proxy;
pub mod udp_proxy;
pub trait DashMapNew {
fn new0() -> Self;
fn new_cap(capacity: usize) -> Self;
}
pub trait ProxyHandler {
fn recv_handle(
&self,
@@ -40,139 +29,29 @@ pub trait ProxyHandler {
fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()>;
}
impl<'a, K: 'a + Eq + std::hash::Hash, V: 'a> DashMapNew for DashMap<K, V> {
fn new0() -> Self {
Self::new_cap(0)
}
fn new_cap(capacity: usize) -> Self {
let shard_amount = (thread::available_parallelism().map_or(4, |v| {
// https://github.com/rust-lang/rust/issues/115868
let n: usize = v.get() * 4;
if n == 0 {
log::warn!("available_parallelism=0");
println!("warn available_parallelism=0");
}
if n < 4 {
return 4;
}
n
}))
.next_power_of_two();
DashMap::with_capacity_and_shard_amount(capacity, shard_amount)
}
}
#[derive(Eq, PartialEq, Ord, PartialOrd, Copy, Clone, Debug)]
pub enum Protocol {
Icmp,
Tcp,
Udp,
}
#[derive(Clone)]
pub struct IpProxyMap {
#[cfg(not(target_os = "android"))]
pub(crate) icmp_handler: IcmpHandler,
pub(crate) tcp_handler: TcpHandler,
pub(crate) udp_handler: UdpHandler,
icmp_proxy: IcmpProxy,
tcp_proxy: TcpProxy,
udp_proxy: UdpProxy,
}
#[cfg(not(target_os = "android"))]
pub async fn init_proxy(
sender: ChannelSender,
pub fn init_proxy(
context: Context,
scheduler: Scheduler,
stop_manager: StopManager,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> {
let tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new0());
let udp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new0());
) -> io::Result<IpProxyMap> {
let icmp_proxy = IcmpProxy::new(context, stop_manager.clone(), current_device, client_cipher)?;
let tcp_proxy = TcpProxy::new(stop_manager.clone())?;
let udp_proxy = UdpProxy::new(scheduler, stop_manager)?;
let icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>> = Arc::new(DashMap::new0());
let udp_socket = UdpSocket::bind("0.0.0.0:0").await?;
let addr = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
let tcp_proxy = TcpProxy::new(addr, tcp_proxy_map.clone()).await?;
let tcp_handler = tcp_proxy.tcp_handler();
let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone())?;
let udp_handler = udp_proxy.udp_handler();
let icmp_handler = {
let icmp_proxy = icmp_proxy::IcmpProxy::new(
addr,
icmp_proxy_map.clone(),
sender.clone(),
current_device.clone(),
client_cipher.clone(),
)?;
let icmp_handler = icmp_proxy.icmp_handler();
thread::spawn(move || {
icmp_proxy.start();
});
icmp_handler
};
Ok((
Ok(IpProxyMap {
icmp_proxy,
tcp_proxy,
udp_proxy,
IpProxyMap {
tcp_handler,
udp_handler,
icmp_handler,
},
))
}
#[cfg(target_os = "android")]
pub async fn init_proxy() -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> {
let tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new0());
let udp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new0());
let addr = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
let udp_socket = UdpSocket::bind("0.0.0.0:0").await?;
let tcp_proxy = TcpProxy::new(addr, tcp_proxy_map.clone()).await?;
let tcp_handler = tcp_proxy.tcp_handler();
let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone())?;
let udp_handler = udp_proxy.udp_handler();
Ok((
tcp_proxy,
udp_proxy,
IpProxyMap {
tcp_handler,
udp_handler,
},
))
}
pub fn send(
buf: &mut [u8],
data_len: usize,
dest_ip: Ipv4Addr,
sender: &ChannelSender,
current_device: &AtomicCell<CurrentDeviceInfo>,
client_cipher: &Cipher,
) {
let current_device = current_device.load();
let virtual_ip = current_device.virtual_ip();
let mut net_packet = NetPacket::new0(12 + data_len, buf).unwrap();
net_packet.set_version(Version::V1);
net_packet.set_protocol(protocol::Protocol::IpTurn);
net_packet.set_transport_protocol(protocol::ip_turn_packet::Protocol::Ipv4.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip);
net_packet.set_destination(dest_ip);
if let Err(e) = client_cipher.encrypt_ipv4(&mut net_packet) {
log::warn!("加密失败:{}", e);
return;
}
if sender
.try_send_by_id(net_packet.buffer(), &dest_ip)
.is_err()
{
let connect_server = current_device.connect_server;
if let Err(e) = sender.send_main(net_packet.buffer(), connect_server) {
log::warn!("发送到目标失败:{},{}", e, connect_server);
}
}
})
}
impl ProxyHandler for IpProxyMap {
@@ -183,15 +62,10 @@ impl ProxyHandler for IpProxyMap {
destination: Ipv4Addr,
) -> io::Result<bool> {
match ipv4.protocol() {
ipv4::protocol::Protocol::Tcp => {
self.tcp_handler.recv_handle(ipv4, source, destination)
}
ipv4::protocol::Protocol::Udp => {
self.udp_handler.recv_handle(ipv4, source, destination)
}
#[cfg(not(target_os = "android"))]
ipv4::protocol::Protocol::Tcp => self.tcp_proxy.recv_handle(ipv4, source, destination),
ipv4::protocol::Protocol::Udp => self.udp_proxy.recv_handle(ipv4, source, destination),
ipv4::protocol::Protocol::Icmp => {
self.icmp_handler.recv_handle(ipv4, source, destination)
self.icmp_proxy.recv_handle(ipv4, source, destination)
}
_ => {
log::warn!(
@@ -208,8 +82,9 @@ impl ProxyHandler for IpProxyMap {
fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()> {
match ipv4.protocol() {
ipv4::protocol::Protocol::Tcp => self.tcp_handler.send_handle(ipv4),
ipv4::protocol::Protocol::Udp => self.udp_handler.send_handle(ipv4),
ipv4::protocol::Protocol::Tcp => self.tcp_proxy.send_handle(ipv4),
ipv4::protocol::Protocol::Udp => self.udp_proxy.send_handle(ipv4),
ipv4::protocol::Protocol::Icmp => self.icmp_proxy.send_handle(ipv4),
_ => Ok(()),
}
}
+341 -131
View File
@@ -1,141 +1,53 @@
use crate::ip_proxy::ProxyHandler;
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use std::io::{Read, Write};
use std::net::{Ipv4Addr, Shutdown, SocketAddrV4};
#[cfg(unix)]
use std::os::fd::AsRawFd;
#[cfg(windows)]
use std::os::windows::io::AsRawSocket;
use std::sync::Arc;
use std::{collections::HashMap, io, net::SocketAddr, thread};
use bytes::{BufMut, BytesMut};
use mio::net::TcpStream;
use mio::{net::TcpListener, Events, Interest, Poll, Registry, Token, Waker};
use parking_lot::Mutex;
use packet::ip::ipv4::packet::IpV4Packet;
use packet::tcp::tcp::TcpPacket;
use std::io;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::io::AsyncReadExt;
use tokio::io::AsyncWriteExt;
use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
use tokio::net::{TcpListener, TcpStream};
use crate::ip_proxy::ProxyHandler;
use crate::util::StopManager;
const SERVER_VAL: usize = 0;
const SERVER: Token = Token(SERVER_VAL);
const NOTIFY_VAL: usize = 1;
const NOTIFY: Token = Token(NOTIFY_VAL);
#[derive(Clone)]
pub struct TcpProxy {
tcp_proxy_port: u16,
tcp_listener: TcpListener,
//真实源地址 -> 目的地址
tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
port: u16,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
}
impl TcpProxy {
pub async fn new(
addr: SocketAddrV4,
tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
) -> io::Result<Self> {
let tcp_listener = TcpListener::bind(addr).await?;
Ok(Self {
tcp_proxy_port: tcp_listener.local_addr()?.port(),
tcp_listener,
tcp_proxy_map,
})
}
pub fn tcp_handler(&self) -> TcpHandler {
TcpHandler(self.tcp_proxy_port, self.tcp_proxy_map.clone())
}
pub async fn start(self) {
let tcp_listener = self.tcp_listener;
let tcp_proxy_map = self.tcp_proxy_map;
loop {
match tcp_listener.accept().await {
Ok((tcp_stream, sender_addr)) => match sender_addr {
SocketAddr::V4(sender_addr) => {
if let Some(entry) = tcp_proxy_map.get(&sender_addr) {
let dest_addr = *entry.value();
drop(entry);
tokio::spawn(async move {
let peer_tcp_stream = match tokio::time::timeout(
Duration::from_secs(5),
TcpStream::connect(dest_addr),
)
.await
{
Ok(peer_tcp_stream) => match peer_tcp_stream {
Ok(peer_tcp_stream) => peer_tcp_stream,
Err(e) => {
log::warn!(
"tcp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
return;
}
},
Err(e) => {
log::warn!(
"tcp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
return;
}
};
if let Err(e) = proxy(tcp_stream, peer_tcp_stream).await {
log::warn!("{}->{},{}", sender_addr, dest_addr, e);
}
});
} else {
log::warn!("tcp代理异常: 来源:{},未找到目标", sender_addr);
}
}
SocketAddr::V6(_) => {}
},
Err(e) => {
log::warn!("tcp代理监听:{:?}", e);
pub fn new(stop_manager: StopManager) -> io::Result<Self> {
let nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>> =
Arc::new(Mutex::new(HashMap::with_capacity(16)));
let tcp_listener = TcpListener::bind(format!("0.0.0.0:{}", 0).parse().unwrap())?;
let port = tcp_listener.local_addr()?.port();
{
let nat_map = nat_map.clone();
thread::spawn(move || {
if let Err(e) = tcp_proxy(tcp_listener, nat_map, stop_manager) {
log::warn!("tcp_proxy:{:?}", e);
}
}
});
}
Ok(Self { port, nat_map })
}
}
async fn proxy(client: TcpStream, server: TcpStream) -> io::Result<()> {
let (client_read, client_write) = client.into_split();
let (server_read, server_write) = server.into_split();
let time = Arc::new(AtomicCell::new(Instant::now()));
let time1 = time.clone();
tokio::spawn(async move {
if let Err(e) = copy(client_read, server_write, &time1).await {
log::warn!("{:?}", e);
}
});
copy(server_read, client_write, &time).await
}
async fn copy(
mut read: OwnedReadHalf,
mut write: OwnedWriteHalf,
time: &AtomicCell<Instant>,
) -> io::Result<()> {
let mut buf = [0; 10240];
loop {
tokio::select! {
result = read.read(&mut buf) =>{
let len = result?;
if len==0{
break;
}
write.write_all(&buf[..len]).await?;
time.store(Instant::now());
}
_ = tokio::time::sleep(Duration::from_secs(600)) =>{
if time.load().elapsed()>=Duration::from_secs(580){
//读写均超时再退出
break;
}
}
}
}
Ok(())
}
#[derive(Clone)]
pub struct TcpHandler(u16, Arc<DashMap<SocketAddrV4, SocketAddrV4>>);
impl ProxyHandler for TcpHandler {
impl ProxyHandler for TcpProxy {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
@@ -147,13 +59,14 @@ impl ProxyHandler for TcpHandler {
let mut tcp_packet = TcpPacket::new(source, destination, ipv4.payload_mut())?;
let source_port = tcp_packet.source_port();
let dest_port = tcp_packet.destination_port();
tcp_packet.set_destination_port(self.0);
tcp_packet.set_destination_port(self.port);
tcp_packet.update_checksum();
ipv4.set_destination_ip(destination);
ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port);
//https://github.com/crossbeam-rs/crossbeam/issues/1023
self.1.insert(key, SocketAddrV4::new(dest_ip, dest_port));
self.nat_map
.lock()
.insert(key, SocketAddrV4::new(dest_ip, dest_port));
Ok(false)
}
@@ -164,8 +77,7 @@ impl ProxyHandler for TcpHandler {
let tcp_packet = TcpPacket::new(src_ip, dest_ip, ipv4.payload_mut())?;
SocketAddrV4::new(dest_ip, tcp_packet.destination_port())
};
if let Some(entry) = self.1.get(&dest_addr) {
let source_addr = entry.value();
if let Some(source_addr) = self.nat_map.lock().get(&dest_addr) {
let source_ip = *source_addr.ip();
let mut tcp_packet = TcpPacket::new(source_ip, dest_ip, ipv4.payload_mut())?;
tcp_packet.set_source_port(source_addr.port());
@@ -176,3 +88,301 @@ impl ProxyHandler for TcpHandler {
Ok(())
}
}
fn tcp_proxy(
mut tcp_listener: TcpListener,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
stop_manager: StopManager,
) -> io::Result<()> {
let mut poll = Poll::new()?;
poll.registry()
.register(&mut tcp_listener, SERVER, Interest::READABLE)?;
let mut events = Events::with_capacity(32);
let mut tcp_map: HashMap<usize, ProxyValue> = HashMap::with_capacity(16);
let mut mapping: HashMap<usize, usize> = HashMap::with_capacity(16);
let stop = Waker::new(poll.registry(), NOTIFY)?;
let _worker = stop_manager.add_listener("tcp_proxy".into(), move || {
if let Err(e) = stop.wake() {
log::warn!("stop tcp_proxy:{:?}", e);
}
})?;
loop {
poll.poll(&mut events, None)?;
for event in events.iter() {
match event.token() {
SERVER => {
accept_handle(
poll.registry(),
&tcp_listener,
&nat_map,
&mut tcp_map,
&mut mapping,
);
}
NOTIFY => {
return Ok(());
}
Token(index) => {
let (val, src_index) = if let Some(v) = tcp_map.get_mut(&index) {
(v, index)
} else {
if let Some(dest_index) = mapping.get(&index) {
if let Some(v) = tcp_map.get_mut(dest_index) {
(v, *dest_index)
} else {
continue;
}
} else {
continue;
}
};
let (stream1, stream2, buf1, buf2) = val.as_mut(index);
if event.is_readable() {
if let Err(e) = readable_handle(stream1, stream2, buf1) {
log::warn!("tcp proxy {:?}", e);
close(src_index, &mut tcp_map, &mut mapping);
}
} else if event.is_writable() {
let read = buf2.len() >= BUF_LEN;
if let Err(e) = writable_handle(stream1, buf2) {
log::warn!("tcp proxy {:?}", e);
close(src_index, &mut tcp_map, &mut mapping);
} else if read {
if let Err(e) = readable_handle(stream2, stream1, buf2) {
log::warn!("tcp proxy {:?}", e);
close(src_index, &mut tcp_map, &mut mapping);
}
}
} else {
close(src_index, &mut tcp_map, &mut mapping);
}
}
}
}
}
}
fn accept_handle(
registry: &Registry,
tcp_listener: &TcpListener,
nat_map: &Mutex<HashMap<SocketAddrV4, SocketAddrV4>>,
tcp_map: &mut HashMap<usize, ProxyValue>,
mapping: &mut HashMap<usize, usize>,
) {
loop {
match tcp_listener.accept() {
Ok((mut src_stream, addr)) => {
#[cfg(windows)]
let src_fd = src_stream.as_raw_socket() as usize;
#[cfg(unix)]
let src_fd = src_stream.as_raw_fd() as usize;
if src_fd == SERVER_VAL || src_fd == NOTIFY_VAL {
log::error!("fd错误:{:?}", src_fd);
continue;
}
let addr = match addr {
SocketAddr::V4(addr) => addr,
SocketAddr::V6(_) => {
// 忽略ipv6
continue;
}
};
let _ = src_stream.set_nodelay(false);
if let Some(dest_addr) = nat_map.lock().get(&addr).cloned() {
match tcp_connect(addr.port(), dest_addr.into()) {
Ok(mut dest_stream) => {
#[cfg(windows)]
let dest_fd = dest_stream.as_raw_socket() as usize;
#[cfg(unix)]
let dest_fd = dest_stream.as_raw_fd() as usize;
if dest_fd == SERVER_VAL || dest_fd == NOTIFY_VAL {
log::error!("fd错误:{:?}", dest_fd);
continue;
}
if let Err(e) = registry.register(
&mut src_stream,
Token(src_fd),
Interest::READABLE,
) {
log::error!("register src_stream:{:?}", e);
continue;
}
if let Err(e) = registry.register(
&mut dest_stream,
Token(dest_fd),
Interest::READABLE,
) {
log::error!("register dest_stream:{:?}", e);
continue;
}
tcp_map.insert(
src_fd,
ProxyValue::new(src_stream, dest_stream, src_fd, dest_fd),
);
mapping.insert(dest_fd, src_fd);
}
Err(e) => {
log::error!("connect:{:?}", e);
}
}
}
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::error!("accept:{:?}", e);
}
}
}
}
fn tcp_connect(src_port: u16, addr: SocketAddr) -> io::Result<TcpStream> {
let socket = socket2::Socket::new(
socket2::Domain::IPV4,
socket2::Type::STREAM,
Some(socket2::Protocol::TCP),
)?;
if socket
.bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into())
.is_err()
{
socket.bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?;
}
socket.set_nonblocking(true)?;
let _ = socket.set_nodelay(false);
if let Err(e) = socket.connect(&addr.into()) {
if e.kind() != io::ErrorKind::WouldBlock {
return Err(e);
}
}
Ok(TcpStream::from_std(socket.into()))
}
struct ProxyValue {
src_stream: TcpStream,
dest_stream: TcpStream,
src_fd: usize,
dest_fd: usize,
src_buf: BytesMut,
dest_buf: BytesMut,
}
const BUF_LEN: usize = 10 * 4096;
impl ProxyValue {
fn new(src_stream: TcpStream, dest_stream: TcpStream, src_fd: usize, dest_fd: usize) -> Self {
Self {
src_stream,
dest_stream,
src_fd,
dest_fd,
src_buf: BytesMut::with_capacity(BUF_LEN),
dest_buf: BytesMut::with_capacity(BUF_LEN),
}
}
fn as_mut(
&mut self,
index: usize,
) -> (&mut TcpStream, &mut TcpStream, &mut BytesMut, &mut BytesMut) {
if index == self.src_fd {
(
&mut self.src_stream,
&mut self.dest_stream,
&mut self.src_buf,
&mut self.dest_buf,
)
} else {
(
&mut self.dest_stream,
&mut self.src_stream,
&mut self.dest_buf,
&mut self.src_buf,
)
}
}
}
fn readable_handle(
stream1: &mut TcpStream,
stream2: &mut TcpStream,
mid_buf: &mut BytesMut,
) -> io::Result<()> {
let mut buf = [0; BUF_LEN];
loop {
match stream1.read(&mut buf) {
Ok(len) => {
if len == 0 {
return Err(io::Error::from(io::ErrorKind::UnexpectedEof));
}
let mut buf = &buf[..len];
if mid_buf.is_empty() {
// 直接写入,避免在buf中过渡
while !buf.is_empty() {
match stream2.write(buf) {
Ok(end) => {
if end == 0 {
return Err(io::Error::from(io::ErrorKind::WriteZero));
}
buf = &buf[end..];
}
Err(e) => {
if e.kind() != io::ErrorKind::WouldBlock {
return Err(e);
}
break;
}
}
}
if buf.is_empty() {
continue;
}
}
mid_buf.reserve(buf.len());
mid_buf.put_slice(buf);
if mid_buf.len() >= BUF_LEN {
// 达到上限不再继续读取
return Ok(());
}
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
}
}
Ok(())
}
fn writable_handle(stream: &mut TcpStream, mid_buf: &mut BytesMut) -> io::Result<()> {
while !mid_buf.is_empty() {
match stream.write(&mid_buf) {
Ok(len) => {
let _ = mid_buf.split_to(len);
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
}
}
Ok(())
}
fn close(
index: usize,
tcp_map: &mut HashMap<usize, ProxyValue>,
mapping: &mut HashMap<usize, usize>,
) {
if let Some(val) = tcp_map.remove(&index) {
let _ = val.src_stream.shutdown(Shutdown::Both);
let _ = val.dest_stream.shutdown(Shutdown::Both);
mapping.remove(&val.src_fd);
mapping.remove(&val.dest_fd);
}
}
+262 -135
View File
@@ -1,146 +1,56 @@
use crate::ip_proxy::{DashMapNew, ProxyHandler};
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use std::net::{Ipv4Addr, SocketAddrV4};
#[cfg(unix)]
use std::os::fd::AsRawFd;
#[cfg(windows)]
use std::os::windows::io::AsRawSocket;
use std::sync::Arc;
use std::time::{Duration, Instant};
use std::{collections::HashMap, io, net::SocketAddr, rc::Rc, thread};
use mio::{net::UdpSocket, Events, Interest, Poll, Token};
use mio::{Registry, Waker};
use parking_lot::Mutex;
use packet::ip::ipv4::packet::IpV4Packet;
use packet::udp::udp::UdpPacket;
use std::io;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::Arc;
use std::time::Duration;
use tokio::net::UdpSocket;
use tokio::time::Instant;
/// 一个udp代理,作用是利用系统协议栈,将udp数据报解析出来再转发到目的地址
use crate::ip_proxy::ProxyHandler;
use crate::util::{Scheduler, StopManager};
const SERVER_VAL: usize = 0;
const SERVER: Token = Token(SERVER_VAL);
const NOTIFY_VAL: usize = 1;
const NOTIFY: Token = Token(NOTIFY_VAL);
// 开了ip代理后使用mstsc,mstsc会误以为在真实局域网,从而不维护udp心跳,导致断连,所以这里尽量长一点过期时间
const NAT_TIMEOUT: Duration = Duration::from_secs(20 * 60);
const NAT_FAST_TIMEOUT: Duration = Duration::from_secs(5 * 60);
const NAT_MAX: usize = 5_000;
#[derive(Clone)]
pub struct UdpProxy {
udp_proxy_port: u16,
udp_socket: Arc<UdpSocket>,
map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
port: u16,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
}
impl UdpProxy {
pub fn new(
udp_socket: UdpSocket,
map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
) -> io::Result<Self> {
let udp_socket = Arc::new(udp_socket);
let udp_proxy_port = udp_socket.local_addr()?.port();
Ok(Self {
udp_proxy_port,
udp_socket,
map,
})
}
pub fn udp_handler(&self) -> UdpHandler {
UdpHandler(self.udp_proxy_port, self.map.clone())
}
pub async fn start(self) {
let map = self.map;
let udp_socket = self.udp_socket;
let mut buf = [0u8; 65536];
let inner_map: Arc<DashMap<SocketAddrV4, (Arc<UdpSocket>, Arc<AtomicCell<Instant>>)>> =
Arc::new(DashMap::new0());
loop {
match udp_socket.recv_from(&mut buf).await {
Ok((len, sender_addr)) => match sender_addr {
SocketAddr::V4(sender_addr) => {
match start0(&buf[..len], sender_addr, &inner_map, &map, &udp_socket).await
{
Ok(_) => {}
Err(e) => {
log::warn!("udp代理异常:{:?},来源:{}", e, sender_addr);
}
}
}
SocketAddr::V6(_) => {}
},
Err(e) => {
log::warn!("udp代理异常:{:?}", e);
}
};
}
}
}
async fn start0(
buf: &[u8],
sender_addr: SocketAddrV4,
inner_map: &Arc<DashMap<SocketAddrV4, (Arc<UdpSocket>, Arc<AtomicCell<Instant>>)>>,
map: &Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
udp_socket: &Arc<UdpSocket>,
) -> io::Result<()> {
if let Some(entry) = inner_map.get(&sender_addr) {
entry.value().1.store(Instant::now());
let udp = entry.value().0.clone();
drop(entry);
udp.send(buf).await?;
} else if let Some(entry) = map.get(&sender_addr) {
let dest_addr = *entry.value();
drop(entry);
//先使用相同的端口,冲突了再随机端口
let peer_udp_socket = match UdpSocket::bind(format!("0.0.0.0:{}", sender_addr.port())).await
pub fn new(scheduler: Scheduler, stop_manager: StopManager) -> io::Result<Self> {
let nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>> =
Arc::new(Mutex::new(HashMap::with_capacity(16)));
let udp = UdpSocket::bind(format!("0.0.0.0:{}", 0).parse().unwrap())?;
let port = udp.local_addr()?.port();
{
Ok(udp) => udp,
Err(_) => UdpSocket::bind("0.0.0.0:0").await?,
};
peer_udp_socket.connect(dest_addr).await?;
peer_udp_socket.send(buf).await?;
let peer_udp_socket = Arc::new(peer_udp_socket);
let inner_map = inner_map.clone();
let time = Arc::new(AtomicCell::new(Instant::now()));
inner_map.insert(sender_addr, (peer_udp_socket.clone(), time.clone()));
let udp_socket = udp_socket.clone();
let map = map.clone();
tokio::spawn(async move {
let mut buf = [0u8; 65536];
loop {
match tokio::time::timeout(Duration::from_secs(600), peer_udp_socket.recv(&mut buf))
.await
{
Ok(rs) => match rs {
Ok(len) => match udp_socket.send_to(&buf[..len], sender_addr).await {
Ok(_) => {}
Err(e) => {
log::warn!(
"udp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
break;
}
},
Err(e) => {
log::warn!(
"udp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
break;
}
},
Err(_) => {
if time.load().elapsed() > Duration::from_secs(580) {
//超时关闭
log::warn!("udp代理超时关闭,来源:{},目标:{}", sender_addr, dest_addr);
break;
}
}
let nat_map = nat_map.clone();
thread::spawn(move || {
if let Err(e) = udp_proxy(udp, nat_map, scheduler, stop_manager) {
log::warn!("udp_proxy:{:?}", e);
}
}
inner_map.remove(&sender_addr);
map.remove(&sender_addr);
});
});
}
Ok(Self { port, nat_map })
}
Ok(())
}
#[derive(Clone)]
pub struct UdpHandler(u16, Arc<DashMap<SocketAddrV4, SocketAddrV4>>);
impl ProxyHandler for UdpHandler {
impl ProxyHandler for UdpProxy {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
@@ -152,12 +62,14 @@ impl ProxyHandler for UdpHandler {
let mut udp_packet = UdpPacket::new(source, destination, ipv4.payload_mut())?;
let source_port = udp_packet.source_port();
let dest_port = udp_packet.destination_port();
udp_packet.set_destination_port(self.0);
udp_packet.set_destination_port(self.port);
udp_packet.update_checksum();
ipv4.set_destination_ip(destination);
ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port);
self.1.insert(key, SocketAddrV4::new(dest_ip, dest_port));
self.nat_map
.lock()
.insert(key.into(), SocketAddrV4::new(dest_ip, dest_port).into());
Ok(false)
}
@@ -168,8 +80,7 @@ impl ProxyHandler for UdpHandler {
let udp_packet = UdpPacket::new(src_ip, dest_ip, ipv4.payload_mut())?;
SocketAddrV4::new(dest_ip, udp_packet.destination_port())
};
if let Some(entry) = self.1.get(&dest_addr) {
let source_addr = entry.value();
if let Some(source_addr) = self.nat_map.lock().get(&dest_addr) {
let source_ip = *source_addr.ip();
let mut udp_packet = UdpPacket::new(source_ip, dest_ip, ipv4.payload_mut())?;
udp_packet.set_source_port(source_addr.port());
@@ -180,3 +91,219 @@ impl ProxyHandler for UdpHandler {
Ok(())
}
}
fn udp_proxy(
mut udp: UdpSocket,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
scheduler: Scheduler,
stop_manager: StopManager,
) -> io::Result<()> {
let mut poll = Poll::new()?;
poll.registry()
.register(&mut udp, SERVER, Interest::READABLE)?;
let mut events = Events::with_capacity(32);
let mut buf = [0; 65536];
let mut token_map: HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)> =
HashMap::with_capacity(64);
let mut udp_map: HashMap<SocketAddrV4, (Rc<UdpSocket>, Instant)> = HashMap::with_capacity(64);
let mut timeout = false;
let waker = Arc::new(Waker::new(poll.registry(), NOTIFY)?);
let stop = waker.clone();
let _worker = stop_manager.add_listener("udp_proxy".into(), move || {
if let Err(e) = stop.wake() {
log::warn!("stop udp_proxy:{:?}", e);
}
})?;
loop {
let mut check = false;
if token_map.is_empty() {
poll.poll(&mut events, None)?;
} else {
//所有事件 50分钟超时
if let Err(e) = poll.poll(&mut events, Some(Duration::from_secs(50 * 60))) {
if e.kind() == io::ErrorKind::TimedOut || e.kind() == io::ErrorKind::WouldBlock {
token_map.clear();
udp_map.clear();
continue;
}
return Err(e);
}
}
for event in events.iter() {
match event.token() {
SERVER => server_handle(
poll.registry(),
&udp,
&nat_map,
&mut token_map,
&mut udp_map,
&mut buf,
),
NOTIFY => {
if stop_manager.is_stop() {
return Ok(());
}
check = true;
}
token => {
if let Err(e) = readable_handle(&udp, &mut token_map, &token, &mut buf) {
log::error!("发送目标失败:{:?}", e);
if let Some((_, src_addr, _)) = token_map.remove(&token) {
udp_map.remove(&src_addr);
}
}
}
}
}
if check {
//超时校验
if token_map.len() > NAT_MAX / 2 {
check_handle(&mut token_map, &mut udp_map, NAT_FAST_TIMEOUT)
} else {
check_handle(&mut token_map, &mut udp_map, NAT_TIMEOUT)
}
timeout = false;
}
if !token_map.is_empty() && !timeout {
//注册超时监听
timeout = true;
let waker = waker.clone();
scheduler.timeout(NAT_FAST_TIMEOUT, move |_| {
let _ = waker.wake();
});
}
}
}
fn check_handle(
token_map: &mut HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)>,
udp_map: &mut HashMap<SocketAddrV4, (Rc<UdpSocket>, Instant)>,
timeout: Duration,
) {
let mut remove_list = Vec::new();
for (token, (_, addr, time)) in token_map.iter() {
if time.elapsed() > timeout {
if let Some((_, time)) = udp_map.get(addr) {
if time.elapsed() > timeout {
//映射超时,需要移除
remove_list.push(*token);
}
}
}
}
for token in remove_list {
if let Some((_, src_addr, _)) = token_map.remove(&token) {
udp_map.remove(&src_addr);
}
}
}
fn server_handle(
registry: &Registry,
udp: &UdpSocket,
nat_map: &Mutex<HashMap<SocketAddrV4, SocketAddrV4>>,
token_map: &mut HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)>,
udp_map: &mut HashMap<SocketAddrV4, (Rc<UdpSocket>, Instant)>,
buf: &mut [u8],
) {
loop {
let (len, src_addr) = match udp.recv_from(buf) {
Ok((len, src_addr)) => match src_addr {
SocketAddr::V4(addr) => (len, addr),
SocketAddr::V6(_) => {
continue;
}
},
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::error!("接收数据失败:{:?}", e);
break;
}
};
if let Some((dest_udp, time)) = udp_map.get_mut(&src_addr) {
//发送失败就当丢包了
let _ = dest_udp.send(&buf[..len]);
*time = Instant::now();
} else if let Some(dest_addr) = nat_map.lock().get(&src_addr).cloned() {
if token_map.len() >= NAT_MAX {
log::error!(
"UDP NAT_MAX:src_addr={:?},dest_addr={:?}",
src_addr,
dest_addr
);
continue;
}
match udp_connect(src_addr.port(), dest_addr.into()) {
Ok((token_val, mut dest_udp)) => {
let token = Token(token_val);
if let Err(e) = registry.register(&mut dest_udp, token, Interest::READABLE) {
log::error!("register失败:{:?},addr={:?}", e, dest_addr);
continue;
}
if dest_udp.send(&buf[..len]).is_ok() {
let dest_udp = Rc::new(dest_udp);
token_map.insert(token, (dest_udp.clone(), src_addr, Instant::now()));
udp_map.insert(src_addr, (dest_udp, Instant::now()));
}
}
Err(e) => {
log::error!("绑定目标地址失败:{:?}", e);
continue;
}
};
}
}
}
/// 得到一个 fd不为SERVER_VAL或者NOTYFY_VAL的socket
fn udp_connect(src_port: u16, addr: SocketAddr) -> io::Result<(usize, UdpSocket)> {
loop {
let udp = if let Ok(udp) =
UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into())
{
udp
} else {
UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?
};
#[cfg(windows)]
let fd = udp.as_raw_socket() as usize;
#[cfg(unix)]
let fd = udp.as_raw_fd() as usize;
if fd == SERVER_VAL || fd == NOTIFY_VAL {
continue;
}
// 只接收目标的数据
udp.connect(addr)?;
return Ok((fd, udp));
}
}
fn readable_handle(
udp: &UdpSocket,
token_map: &mut HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)>,
token: &Token,
buf: &mut [u8],
) -> io::Result<()> {
if let Some((dest_udp, src_addr, time)) = token_map.get_mut(&token) {
loop {
let len = match dest_udp.recv(buf) {
Ok(rs) => rs,
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
};
if len == 0 {
return Err(io::Error::from(io::ErrorKind::UnexpectedEof));
}
let _ = udp.send_to(&buf[..len], (*src_addr).into());
}
*time = Instant::now();
}
Ok(())
}