[mio] 整理数据处理逻辑

This commit is contained in:
lubeilin
2024-02-29 22:27:15 +08:00
parent fe499d0476
commit 98bd713a91
17 changed files with 2244 additions and 225 deletions
+201
View File
@@ -0,0 +1,201 @@
#[cfg(feature = "server_encrypt")]
use rsa::RsaPublicKey;
use std::fmt::{Display, Formatter};
use std::io;
use std::net::{Ipv4Addr, SocketAddr};
#[derive(Debug)]
pub struct DeviceInfo {
pub name: String,
pub version: String,
}
impl Display for DeviceInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!("name={} ,version={}", self.name, self.version))
}
}
impl DeviceInfo {
pub fn new(name: String, version: String) -> Self {
return Self { name, version };
}
}
#[derive(Debug)]
pub struct ConnectInfo {
// 第几次连接,从1开始
pub count: usize,
// 服务端地址
pub address: SocketAddr,
}
impl Display for ConnectInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!("count={} ,address={}", self.count, self.address))
}
}
impl ConnectInfo {
pub fn new(count: usize, address: SocketAddr) -> Self {
Self { count, address }
}
}
#[derive(Debug)]
pub struct HandshakeInfo {
//服务端公钥
#[cfg(feature = "server_encrypt")]
pub public_key: Option<RsaPublicKey>,
//服务端指纹
#[cfg(feature = "server_encrypt")]
pub finger: Option<String>,
//服务端版本
pub version: String,
}
impl Display for HandshakeInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
#[cfg(feature = "server_encrypt")]
return match &self.finger {
None => f.write_str(&format!("no_secret server version={}", self.version)),
Some(finger) => f.write_str(&format!(
"finger={} ,server version={}",
finger, self.version
)),
};
#[cfg(not(feature = "server_encrypt"))]
f.write_str(&format!("server version={}", self.version))
}
}
#[cfg(feature = "server_encrypt")]
impl HandshakeInfo {
pub fn new(public_key: RsaPublicKey, finger: String, version: String) -> Self {
Self {
public_key: Some(public_key),
finger: Some(finger),
version,
}
}
pub fn new_no_secret(version: String) -> Self {
Self {
public_key: None,
finger: None,
version,
}
}
}
#[cfg(not(feature = "server_encrypt"))]
impl HandshakeInfo {
pub fn new_no_secret(version: String) -> Self {
Self { version }
}
}
#[derive(Debug)]
pub struct RegisterInfo {
//本机虚拟IP
pub virtual_ip: Ipv4Addr,
//子网掩码
pub virtual_netmask: Ipv4Addr,
//虚拟网关
pub virtual_gateway: Ipv4Addr,
}
impl Display for RegisterInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!(
"ip={} ,netmask={} ,gateway={}",
self.virtual_ip, self.virtual_netmask, self.virtual_gateway,
))
}
}
impl RegisterInfo {
pub fn new(virtual_ip: Ipv4Addr, virtual_netmask: Ipv4Addr, virtual_gateway: Ipv4Addr) -> Self {
Self {
virtual_ip,
virtual_netmask,
virtual_gateway,
}
}
}
#[derive(Debug)]
pub struct ErrorInfo {
pub code: ErrorType,
pub msg: Option<String>,
pub source: Option<io::Error>,
}
impl Display for ErrorInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!("ErrorType={:?} ", self.code))?;
if let Some(msg) = &self.msg {
f.write_str(&format!(",msg={:?} ", msg))?;
}
if let Some(source) = &self.source {
f.write_str(&format!(",source={:?} ", source))?;
}
Ok(())
}
}
impl ErrorInfo {
pub fn new(code: ErrorType) -> Self {
Self {
code,
msg: None,
source: None,
}
}
pub fn new_msg(code: ErrorType, msg: String) -> Self {
Self {
code,
msg: Some(msg),
source: None,
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum ErrorType {
TokenError,
Disconnect,
AddressExhausted,
IpAlreadyExists,
InvalidIp,
Unknown,
}
impl Into<u8> for ErrorType {
fn into(self) -> u8 {
match self {
ErrorType::TokenError => 1,
ErrorType::Disconnect => 2,
ErrorType::AddressExhausted => 3,
ErrorType::IpAlreadyExists => 4,
ErrorType::InvalidIp => 5,
ErrorType::Unknown => 6,
}
}
}
pub trait VntCallback: Clone + Send + Sync + 'static {
/// 创建网卡的信息
fn create_tun(&self, _info: DeviceInfo) {}
/// 连接
fn connect(&self, _info: ConnectInfo) {}
/// 握手,返回false则拒绝握手,可在此处检查服务端信息
fn handshake(&self, _info: HandshakeInfo) -> bool {
true
}
/// 注册,返回false则拒绝注册
fn register(&self, _info: RegisterInfo) -> bool {
true
}
/// 异常信息
fn error(&self, _info: ErrorInfo) {}
/// 服务停止
fn stop(&self) {}
}
+72
View File
@@ -0,0 +1,72 @@
use std::io;
#[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher;
use crate::handle::{GATEWAY_IP, SELF_IP};
use crate::proto::message::{HandshakeRequest, SecretHandshakeRequest};
use crate::protocol::body::RSA_ENCRYPTION_RESERVED;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
use protobuf::Message;
pub enum HandshakeEnum {
NotSecret,
KeyError,
Timeout,
ServerError(String),
Other(String),
}
/// 第一次握手数据
pub fn handshake_request_packet(secret: bool) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = HandshakeRequest::new();
request.secret = secret;
request.version = crate::VNT_VERSION.to_string();
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("handshake_request_packet {:?}", e),
)
})?;
let buf = vec![0u8; 12 + bytes.len()];
let mut net_packet = NetPacket::new(buf)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::HandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(net_packet)
}
/// 第二次加密握手
#[cfg(feature = "server_encrypt")]
pub fn secret_handshake_request_packet(
rsa_cipher: &RsaCipher,
token: String,
key: &[u8],
) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = SecretHandshakeRequest::new();
request.token = token;
request.key = key.to_vec();
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("secret_handshake_request_packet {:?}", e),
)
})?;
let mut net_packet = NetPacket::new0(
12 + bytes.len(),
vec![0u8; 12 + bytes.len() + RSA_ENCRYPTION_RESERVED],
)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::SecretHandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(rsa_cipher.encrypt(&mut net_packet)?)
}
+73
View File
@@ -0,0 +1,73 @@
use std::net::ToSocketAddrs;
use std::sync::Arc;
use std::time::Duration;
use crossbeam_utils::atomic::AtomicCell;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::{BaseConfigInfo, CurrentDeviceInfo};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{control_packet, NetPacket, Protocol, Version, MAX_TTL};
use crate::util::Scheduler;
pub fn addr_request(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
server_cipher: Cipher,
config: BaseConfigInfo,
) {
addr_request0(&context, &current_device_info, &server_cipher, &config);
// 9秒发送一次
let rs = scheduler.timeout(Duration::from_secs(9), |s| {
addr_request(s, context, current_device_info, server_cipher, config)
});
if !rs {
log::info!("定时任务停止");
}
}
pub fn addr_request0(
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
server_cipher: &Cipher,
config: &BaseConfigInfo,
) {
let mut current_dev = current_device.load();
// 探测服务端地址变化
if let Ok(mut addr) = config.server_addr.to_socket_addrs() {
if let Some(addr) = addr.next() {
if addr != current_dev.connect_server {
let mut tmp = current_dev.clone();
tmp.connect_server = addr;
let rs = current_device.compare_exchange(current_dev, tmp);
current_dev.connect_server = addr;
log::info!(
"服务端地址变化,旧地址:{},新地址:{},替换结果:{}",
current_dev.connect_server,
addr,
rs.is_ok()
);
}
}
}
if current_dev.connect_server.is_ipv4() {
// 如果连接的是ipv4服务,则探测公网端口
let gateway_ip = current_dev.virtual_gateway;
let src_ip = current_dev.virtual_ip;
let mut packet = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED]).unwrap();
packet.set_version(Version::V1);
packet.set_gateway_flag(true);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::AddrRequest.into());
packet.first_set_ttl(MAX_TTL);
packet.set_source(src_ip);
packet.set_destination(gateway_ip);
if let Err(e) = server_cipher.encrypt_ipv4(&mut packet) {
log::warn!("AddrRequest err={:?}", e)
} else {
context.try_send_all_main(packet.buffer(), current_dev.connect_server);
}
}
}
+239
View File
@@ -0,0 +1,239 @@
use std::io;
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::time::Duration;
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use rand::prelude::SliceRandom;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::PingPacket;
use crate::protocol::{control_packet, NetPacket, Protocol, Version, MAX_TTL};
use crate::util::Scheduler;
/// 定时发送心跳包
pub fn heartbeat(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
client_cipher: Cipher,
server_cipher: Cipher,
) {
heartbeat0(
&context,
&current_device_info.load(),
&device_list,
&client_cipher,
&server_cipher,
);
// 心跳包 3秒发送一次
let rs = scheduler.timeout(Duration::from_secs(3), |s| {
heartbeat(
s,
context,
current_device_info,
device_list,
client_cipher,
server_cipher,
)
});
if !rs {
log::info!("定时任务停止");
}
}
fn heartbeat0(
context: &Context,
current_device: &CurrentDeviceInfo,
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) {
let gateway_ip = current_device.virtual_gateway;
let src_ip = current_device.virtual_ip;
// 可能服务器ip发生变化,导致发送失败
let mut is_send_gateway = false;
match heartbeat_packet_server(device_list, server_cipher, src_ip, gateway_ip) {
Ok(net_packet) => {
if let Err(e) = context.send_default(net_packet.buffer(), current_device.connect_server)
{
log::warn!("heartbeat err={:?}", e)
} else {
is_send_gateway = true
}
}
Err(e) => {
log::error!("heartbeat_packet err={:?}", e);
}
}
for (dest_ip, routes) in context.route_table.route_table() {
let net_packet = if current_device.is_gateway(&dest_ip) {
if is_send_gateway {
continue;
}
heartbeat_packet_server(device_list, server_cipher, src_ip, gateway_ip)
} else {
heartbeat_packet_client(client_cipher, src_ip, dest_ip)
};
let net_packet = match net_packet {
Ok(net_packet) => net_packet,
Err(e) => {
log::error!("heartbeat_packet err={:?}", e);
continue;
}
};
for route in routes {
if let Err(e) = context.send_by_key(net_packet.buffer(), route.route_key()) {
log::warn!("heartbeat err={:?}", e)
}
}
}
let peer_list = { device_list.lock().1.clone() };
for peer in &peer_list {
if !peer.status.is_online() {
continue;
}
if current_device.is_gateway(&peer.virtual_ip) {
continue;
}
if context.route_table.route_one(&peer.virtual_ip).is_none() {
//路由为空,则向服务端地址发送
let net_packet = match heartbeat_packet_client(client_cipher, src_ip, peer.virtual_ip) {
Ok(net_packet) => net_packet,
Err(e) => {
log::error!("heartbeat_packet err={:?}", e);
continue;
}
};
if let Err(e) = context.send_default(net_packet.buffer(), current_device.connect_server)
{
log::error!("heartbeat_packet send_default err={:?}", e);
}
}
}
}
/// 客户端中继路径探测,延迟启动
pub fn client_relay(
scheduler: &Scheduler,
context: Context,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
client_cipher: Cipher,
) {
let rs = scheduler.timeout(Duration::from_secs(30), move |s| {
client_relay_(s, context, current_device, device_list, client_cipher)
});
if !rs {
log::info!("定时任务停止");
}
}
/// 客户端中继路径探测,每30秒探测一次
fn client_relay_(
scheduler: &Scheduler,
context: Context,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
client_cipher: Cipher,
) {
if let Err(e) = client_relay0(
&context,
&current_device.load(),
&device_list,
&client_cipher,
) {
log::error!("{:?}", e);
}
let rs = scheduler.timeout(Duration::from_secs(30), move |s| {
client_relay_(s, context, current_device, device_list, client_cipher)
});
if !rs {
log::info!("定时任务停止");
}
}
fn client_relay0(
context: &Context,
current_device: &CurrentDeviceInfo,
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
client_cipher: &Cipher,
) -> io::Result<()> {
let peer_list = { device_list.lock().1.clone() };
let mut routes = context.route_table.route_table_p2p();
for peer in &peer_list {
if peer.virtual_ip == current_device.virtual_ip {
continue;
}
if let Some(route) = context.route_table.route_one(&peer.virtual_ip) {
if route.is_p2p() && !context.first_latency() {
continue;
}
}
let client_packet =
heartbeat_packet_client(client_cipher, current_device.virtual_ip, peer.virtual_ip)?;
//随机发送到其他地址,看有没有客户端符合转发条件
routes.shuffle(&mut rand::thread_rng());
for (index, (ip, route)) in routes.iter().enumerate() {
if current_device.is_gateway(ip) {
continue;
}
if let Err(e) = context.send_by_key(client_packet.buffer(), route.route_key()) {
log::error!("{:?}", e);
}
if index >= 2 {
break;
}
}
}
Ok(())
}
/// 构建心跳包
fn heartbeat_packet(
src: Ipv4Addr,
dest: Ipv4Addr,
) -> io::Result<NetPacket<[u8; 12 + 4 + ENCRYPTION_RESERVED]>> {
let mut net_packet = NetPacket::new_encrypt([0u8; 12 + 4 + ENCRYPTION_RESERVED])?;
net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::Control);
net_packet.set_transport_protocol(control_packet::Protocol::Ping.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(src);
net_packet.set_destination(dest);
let mut ping = PingPacket::new(net_packet.payload_mut())?;
ping.set_time(crate::handle::now_time() as u16);
Ok(net_packet)
}
fn heartbeat_packet_client(
client_cipher: &Cipher,
src: Ipv4Addr,
dest: Ipv4Addr,
) -> io::Result<NetPacket<[u8; 12 + 4 + ENCRYPTION_RESERVED]>> {
let mut net_packet = heartbeat_packet(src, dest)?;
client_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
fn heartbeat_packet_server(
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
server_cipher: &Cipher,
src: Ipv4Addr,
dest: Ipv4Addr,
) -> io::Result<NetPacket<[u8; 12 + 4 + ENCRYPTION_RESERVED]>> {
let mut net_packet = heartbeat_packet(src, dest)?;
let mut ping = PingPacket::new(net_packet.payload_mut())?;
ping.set_epoch(device_list.lock().0);
net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
+133
View File
@@ -0,0 +1,133 @@
use crate::channel::context::Context;
use crate::channel::idle::{Idle, IdleType};
use crate::channel::sender::AcceptSocketSender;
use crate::handle::callback::{ConnectInfo, ErrorType};
use crate::handle::{handshaker, BaseConfigInfo, CurrentDeviceInfo};
use crate::util::Scheduler;
use crate::{ErrorInfo, VntCallback};
use crossbeam_utils::atomic::AtomicCell;
use mio::net::TcpStream;
use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
pub fn idle_route<Call: VntCallback>(
scheduler: &Scheduler,
idle: Idle,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
call: Call,
) {
let delay = idle_route0(&idle, &context, &current_device_info, &call);
let rs = scheduler.timeout(delay, move |s| {
idle_route(s, idle, context, current_device_info, call)
});
if !rs {
log::info!("定时任务停止");
}
}
pub fn idle_gateway<Call: VntCallback>(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
config: BaseConfigInfo,
tcp_socket_sender: AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
call: Call,
mut connect_count: usize,
) {
idle_gateway0(
&context,
&current_device_info,
&config,
&tcp_socket_sender,
&call,
&mut connect_count,
);
let rs = scheduler.timeout(Duration::from_secs(5), move |s| {
idle_gateway(
s,
context,
current_device_info,
config,
tcp_socket_sender,
call,
connect_count,
)
});
if !rs {
log::info!("定时任务停止");
}
}
fn idle_gateway0<Call: VntCallback>(
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
config: &BaseConfigInfo,
tcp_socket_sender: &AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
call: &Call,
connect_count: &mut usize,
) {
let cur = current_device.load();
if let Err(e) =
check_gateway_channel(context, cur, config, tcp_socket_sender, call, connect_count)
{
log::warn!("{:?}", e);
}
}
fn idle_route0<Call: VntCallback>(
idle: &Idle,
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
call: &Call,
) -> Duration {
let cur = current_device.load();
match idle.next_idle() {
IdleType::Timeout(ip, route) => {
context.route_table.remove_route(&ip, route);
if cur.is_gateway(&ip) {
//网关路由过期,则需要改变状态
let _ = context.change_status(current_device);
call.error(ErrorInfo::new(ErrorType::Disconnect));
}
Duration::from_millis(100)
}
IdleType::Sleep(duration) => duration,
IdleType::None => Duration::from_millis(3000),
}
}
fn check_gateway_channel<Call: VntCallback>(
context: &Context,
current_device: CurrentDeviceInfo,
config: &BaseConfigInfo,
tcp_socket_sender: &AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
call: &Call,
count: &mut usize,
) -> io::Result<()> {
let gateway_route = context
.route_table
.route_one(&current_device.virtual_gateway);
if gateway_route.is_none() {
*count += 1;
//需要重连
call.connect(ConnectInfo::new(*count, current_device.connect_server));
let request_packet = handshaker::handshake_request_packet(config.client_secret)?;
if let Err(e) = context.send_default(request_packet.buffer(), current_device.connect_server)
{
log::warn!("{:?}", e);
if context.is_main_tcp() {
//tcp需要重连
let tcp_stream = std::net::TcpStream::connect(current_device.connect_server)?;
tcp_stream.set_nonblocking(true)?;
if let Err(e) = tcp_socket_sender.try_add_socket((
TcpStream::from_std(tcp_stream),
current_device.connect_server,
Some(request_packet.into_buffer()),
)) {
log::warn!("{:?}", e)
}
}
}
}
Ok(())
}
+16
View File
@@ -0,0 +1,16 @@
mod heartbeat;
pub use heartbeat::client_relay;
pub use heartbeat::heartbeat;
mod re_nat_type;
pub use re_nat_type::retrieve_nat_type;
mod addr_request;
pub use addr_request::addr_request;
mod punch;
pub use punch::punch;
mod idle;
pub use idle::idle_gateway;
pub use idle::idle_route;
+185
View File
@@ -0,0 +1,185 @@
use std::net::Ipv4Addr;
use std::sync::mpsc::Receiver;
use std::sync::Arc;
use std::time::Duration;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use protobuf::Message;
use rand::prelude::SliceRandom;
use crate::channel::context::Context;
use crate::channel::punch::{NatInfo, Punch};
use crate::cipher::Cipher;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::nat::NatTest;
use crate::proto::message::{PunchInfo, PunchNatType};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{control_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL};
use crate::util::Scheduler;
pub fn punch(
scheduler: &Scheduler,
context: Context,
nat_test: NatTest,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
receiver: Receiver<(Ipv4Addr, NatInfo)>,
punch: Punch,
) {
punch_request(
scheduler,
context,
nat_test,
device_list,
current_device.clone(),
client_cipher.clone(),
0,
);
thread::spawn(move || {
punch_start(receiver, punch, current_device, client_cipher);
});
}
/// 接收打洞消息,配合对端打洞
fn punch_start(
receiver: Receiver<(Ipv4Addr, NatInfo)>,
mut punch: Punch,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) {
while let Ok((peer_ip, nat_info)) = receiver.recv() {
let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED]).unwrap();
packet.set_version(Version::V1);
packet.first_set_ttl(1);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::PunchRequest.into());
packet.set_source(current_device.load().virtual_ip());
packet.set_destination(peer_ip);
log::info!("发起打洞,目标:{:?},{:?}", peer_ip, nat_info);
if let Err(e) = client_cipher.encrypt_ipv4(&mut packet) {
log::error!("{:?}", e);
continue;
}
if let Err(e) = punch.punch(packet.buffer(), peer_ip, nat_info) {
log::warn!("{:?}", e)
}
}
}
/// 定时发起打洞请求
fn punch_request(
scheduler: &Scheduler,
context: Context,
nat_test: NatTest,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
count: usize,
) {
if let Err(e) = punch0(
&context,
&nat_test,
&device_list,
&current_device,
&client_cipher,
) {
log::warn!("{:?}", e)
}
let sleep_time = [3, 5, 7, 11, 13, 17, 19, 23, 29];
let secs = Duration::from_secs(sleep_time[count % sleep_time.len()]);
let rs = scheduler.timeout(secs, move |s| {
punch_request(
s,
context,
nat_test,
device_list,
current_device,
client_cipher,
count + 1,
);
});
if !rs {
log::info!("定时任务停止");
}
}
/// 随机对需要打洞的客户端发起打洞请求
fn punch0(
context: &Context,
nat_test: &NatTest,
device_list: &Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: &Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: &Cipher,
) -> io::Result<()> {
let current_device = current_device.load();
let nat_info = nat_test.nat_info();
let mut list = device_list.lock().clone().1;
list.shuffle(&mut rand::thread_rng());
let mut count = 0;
for info in list {
if !info.status.is_online() {
continue;
}
if info.virtual_ip <= current_device.virtual_ip {
continue;
}
if !context.route_table.need_punch(&info.virtual_ip) {
continue;
}
count += 1;
if count > 2 {
break;
}
let packet = punch_packet(
client_cipher,
current_device.virtual_ip(),
&nat_info,
info.virtual_ip,
)?;
context.send_default(packet.buffer(), current_device.connect_server)?;
}
Ok(())
}
fn punch_packet(
client_cipher: &Cipher,
virtual_ip: Ipv4Addr,
nat_info: &NatInfo,
dest: Ipv4Addr,
) -> io::Result<NetPacket<Vec<u8>>> {
let mut punch_reply = PunchInfo::new();
punch_reply.reply = false;
punch_reply.public_ip_list = nat_info
.public_ips
.iter()
.map(|ip| u32::from_be_bytes(ip.octets()))
.collect();
punch_reply.public_port = nat_info.public_ports.get(0).map_or(0, |v| *v as u32);
punch_reply.public_ports = nat_info.public_ports.iter().map(|e| *e as u32).collect();
punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.local_ip = u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED));
punch_reply.local_port = nat_info.udp_ports[0] as u32;
punch_reply.tcp_port = nat_info.tcp_port as u32;
punch_reply.udp_ports = nat_info.udp_ports.iter().map(|e| *e as u32).collect();
if let Some(ipv6) = nat_info.ipv6 {
punch_reply.ipv6_port = nat_info.udp_ports[0] as u32;
punch_reply.ipv6 = ipv6.octets().to_vec();
}
punch_reply.nat_type = protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
let bytes = punch_reply
.write_to_bytes()
.map_err(|e| io::Error::new(io::ErrorKind::Other, format!("punch_packet {:?}", e)))?;
let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::OtherTurn);
net_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip);
net_packet.set_destination(dest);
net_packet.set_payload(&bytes)?;
client_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
+45
View File
@@ -0,0 +1,45 @@
use std::thread;
use std::time::Duration;
use crate::channel::context::Context;
use crate::channel::sender::AcceptSocketSender;
use crate::nat;
use crate::nat::NatTest;
use crate::util::Scheduler;
/// 10分钟探测一次nat
pub fn retrieve_nat_type(
scheduler: &Scheduler,
context: Context,
nat_test: NatTest,
udp_socket_sender: AcceptSocketSender<Option<Vec<mio::net::UdpSocket>>>,
) {
retrieve_nat_type0(context.clone(), nat_test.clone(), udp_socket_sender.clone());
scheduler.timeout(Duration::from_secs(60 * 10), move |s| {
retrieve_nat_type(s, context, nat_test, udp_socket_sender)
});
}
fn retrieve_nat_type0(
context: Context,
nat_test: NatTest,
udp_socket_sender: AcceptSocketSender<Option<Vec<mio::net::UdpSocket>>>,
) {
thread::spawn(move || {
if nat_test.can_update() {
let nat_info = nat_test.nat_info();
let local_ipv4 = nat::local_ipv4();
let local_ipv6 = nat::local_ipv6();
let nat_info = nat_test.re_test(
nat_info.public_ports,
local_ipv4,
local_ipv6,
nat_info.udp_ports,
nat_info.tcp_port,
);
if let Err(e) = context.switch(nat_info.nat_type, &udp_socket_sender) {
log::warn!("{:?}", e);
}
}
});
}
+89 -12
View File
@@ -1,12 +1,15 @@
use std::net::{Ipv4Addr, SocketAddr};
pub mod handshake_handler;
pub mod heartbeat_handler;
pub mod punch_handler;
pub mod recv_handler;
pub mod registration_handler;
pub mod callback;
pub mod handshaker;
pub mod maintain;
pub mod recv_data;
pub mod registrar;
pub mod tun_tap;
const SELF_IP: Ipv4Addr = Ipv4Addr::new(0, 0, 0, 2);
const GATEWAY_IP: Ipv4Addr = Ipv4Addr::new(0, 0, 0, 1);
pub fn now_time() -> u64 {
let now = std::time::SystemTime::now();
if let Ok(timestamp) = now.duration_since(std::time::UNIX_EPOCH) {
@@ -41,12 +44,48 @@ impl PeerDeviceInfo {
}
}
#[derive(Clone, Debug)]
pub struct BaseConfigInfo {
pub name: String,
pub token: String,
pub ip: Option<Ipv4Addr>,
pub client_secret: bool,
pub device_id: String,
pub server_addr: String,
}
impl BaseConfigInfo {
pub fn new(
name: String,
token: String,
ip: Option<Ipv4Addr>,
client_secret: bool,
device_id: String,
server_addr: String,
) -> Self {
Self {
name,
token,
ip,
client_secret,
device_id,
server_addr,
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd)]
pub enum PeerDeviceStatus {
Online,
Offline,
}
impl PeerDeviceStatus {
pub fn is_online(&self) -> bool {
self == &PeerDeviceStatus::Online
}
}
impl Into<u8> for PeerDeviceStatus {
fn into(self) -> u8 {
match self {
@@ -73,27 +112,32 @@ pub enum ConnectStatus {
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub struct CurrentDeviceInfo {
virtual_ip: Ipv4Addr,
pub virtual_gateway: Ipv4Addr,
//本机虚拟IP
pub virtual_ip: Ipv4Addr,
//子网掩码
pub virtual_netmask: Ipv4Addr,
//虚拟网关
pub virtual_gateway: Ipv4Addr,
//网络地址
pub virtual_network: Ipv4Addr,
//直接广播地址
pub broadcast_address: Ipv4Addr,
pub broadcast_ip: Ipv4Addr,
//链接的服务器地址
pub connect_server: SocketAddr,
//连接状态
pub status: ConnectStatus,
}
impl CurrentDeviceInfo {
pub fn new(
virtual_ip: Ipv4Addr,
virtual_gateway: Ipv4Addr,
virtual_netmask: Ipv4Addr,
virtual_gateway: Ipv4Addr,
connect_server: SocketAddr,
) -> Self {
let broadcast_address = (!u32::from_be_bytes(virtual_netmask.octets()))
let broadcast_ip = (!u32::from_be_bytes(virtual_netmask.octets()))
| u32::from_be_bytes(virtual_gateway.octets());
let broadcast_address = Ipv4Addr::from(broadcast_address);
let broadcast_ip = Ipv4Addr::from(broadcast_ip);
let virtual_network = u32::from_be_bytes(virtual_netmask.octets())
& u32::from_be_bytes(virtual_gateway.octets());
let virtual_network = Ipv4Addr::from(virtual_network);
@@ -102,10 +146,40 @@ impl CurrentDeviceInfo {
virtual_netmask,
virtual_gateway,
virtual_network,
broadcast_address,
broadcast_ip,
connect_server,
status: ConnectStatus::Connecting,
}
}
pub fn new0(connect_server: SocketAddr) -> Self {
Self {
virtual_ip: Ipv4Addr::UNSPECIFIED,
virtual_gateway: Ipv4Addr::UNSPECIFIED,
virtual_netmask: Ipv4Addr::UNSPECIFIED,
virtual_network: Ipv4Addr::UNSPECIFIED,
broadcast_ip: Ipv4Addr::UNSPECIFIED,
connect_server,
status: ConnectStatus::Connecting,
}
}
pub fn update(
&mut self,
virtual_ip: Ipv4Addr,
virtual_netmask: Ipv4Addr,
virtual_gateway: Ipv4Addr,
) {
let broadcast_ip = (!u32::from_be_bytes(virtual_netmask.octets()))
| u32::from_be_bytes(virtual_gateway.octets());
let broadcast_ip = Ipv4Addr::from(broadcast_ip);
let virtual_network = u32::from_be_bytes(virtual_netmask.octets())
& u32::from_be_bytes(virtual_gateway.octets());
let virtual_network = Ipv4Addr::from(virtual_network);
self.virtual_ip = virtual_ip;
self.virtual_netmask = virtual_netmask;
self.virtual_gateway = virtual_gateway;
self.broadcast_ip = broadcast_ip;
self.virtual_network = virtual_network;
}
#[inline]
pub fn virtual_ip(&self) -> Ipv4Addr {
self.virtual_ip
@@ -114,4 +188,7 @@ impl CurrentDeviceInfo {
pub fn virtual_gateway(&self) -> Ipv4Addr {
self.virtual_gateway
}
pub fn is_gateway(&self, ip: &Ipv4Addr) -> bool {
&self.virtual_gateway == ip || ip == &GATEWAY_IP
}
}
+338
View File
@@ -0,0 +1,338 @@
use std::collections::HashMap;
use std::io;
use std::net::{Ipv4Addr, Ipv6Addr};
use std::sync::mpsc::SyncSender;
use std::sync::Arc;
use parking_lot::RwLock;
use protobuf::Message;
use packet::icmp::{icmp, Kind};
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use tun::device::IFace;
use tun::Device;
use crate::channel::context::Context;
use crate::channel::punch::NatInfo;
use crate::channel::{Route, RouteKey};
use crate::cipher::Cipher;
use crate::external_route::AllowExternalRoute;
use crate::handle::recv_data::PacketHandler;
use crate::handle::CurrentDeviceInfo;
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::{IpProxyMap, ProxyHandler};
use crate::nat::NatTest;
use crate::proto::message::{PunchInfo, PunchNatType};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::ControlPacket;
use crate::protocol::{
control_packet, ip_turn_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL,
};
/// 处理来源于客户端的包
#[derive(Clone)]
pub struct ClientPacketHandler {
device: Arc<Device>,
client_cipher: Cipher,
relay: bool,
punch_sender: SyncSender<(Ipv4Addr, NatInfo)>,
peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>>,
nat_test: NatTest,
route: AllowExternalRoute,
#[cfg(feature = "ip_proxy")]
ip_proxy_map: Option<IpProxyMap>,
}
impl ClientPacketHandler {
pub fn new(
device: Arc<Device>,
client_cipher: Cipher,
relay: bool,
punch_sender: SyncSender<(Ipv4Addr, NatInfo)>,
peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>>,
nat_test: NatTest,
route: AllowExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
) -> Self {
Self {
device,
client_cipher,
relay,
punch_sender,
peer_nat_info_map,
nat_test,
route,
#[cfg(feature = "ip_proxy")]
ip_proxy_map,
}
}
}
impl PacketHandler for ClientPacketHandler {
fn handle(
&self,
mut net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
context: &Context,
current_device: &CurrentDeviceInfo,
) -> io::Result<()> {
self.client_cipher.decrypt_ipv4(&mut net_packet)?;
match net_packet.protocol() {
Protocol::Service => {}
Protocol::Error => {}
Protocol::Control => {
self.control(context, current_device, net_packet, route_key)?;
}
Protocol::IpTurn => {
self.ip_turn(net_packet, context, current_device, route_key)?;
}
Protocol::OtherTurn => {
self.other_turn(context, current_device, net_packet, route_key)?;
}
Protocol::Unknown(_) => {}
}
Ok(())
}
}
impl ClientPacketHandler {
fn ip_turn(
&self,
mut net_packet: NetPacket<&mut [u8]>,
context: &Context,
current_device: &CurrentDeviceInfo,
route_key: RouteKey,
) -> io::Result<()> {
let destination = net_packet.destination();
let source = net_packet.source();
match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) {
ip_turn_packet::Protocol::Ipv4 => {
let mut ipv4 = IpV4Packet::new(net_packet.payload_mut())?;
match ipv4.protocol() {
ipv4::protocol::Protocol::Icmp => {
if ipv4.destination_ip() == destination {
let mut icmp_packet = icmp::IcmpPacket::new(ipv4.payload_mut())?;
if icmp_packet.kind() == Kind::EchoRequest {
//开启ping
icmp_packet.set_kind(Kind::EchoReply);
icmp_packet.update_checksum();
ipv4.set_source_ip(destination);
ipv4.set_destination_ip(source);
ipv4.update_checksum();
net_packet.set_source(destination);
net_packet.set_destination(source);
//不管加不加密,和接收到的数据长度都一致
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.send_by_key(net_packet.buffer(), route_key)?;
return Ok(());
}
}
}
_ => {}
}
// ip代理只关心实际目标
let real_dest = ipv4.destination_ip();
if real_dest != destination
&& !(real_dest.is_broadcast()
|| real_dest.is_multicast()
|| real_dest == current_device.broadcast_ip
|| real_dest.is_unspecified())
{
if !self.route.allow(&ipv4.destination_ip()) {
//拦截不符合的目标
return Ok(());
}
#[cfg(feature = "ip_proxy")]
if let Some(ip_proxy_map) = &self.ip_proxy_map {
if ip_proxy_map.recv_handle(&mut ipv4, source, destination)? {
return Ok(());
}
}
}
self.device.write(net_packet.payload())?;
}
ip_turn_packet::Protocol::Ipv4Broadcast => {
//客户端不帮忙转发广播包,所以不会出现这种类型的数据
}
ip_turn_packet::Protocol::Unknown(_) => {}
}
Ok(())
}
fn control(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
mut net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
let source = net_packet.source();
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PingPacket(_) => {
net_packet.set_transport_protocol(control_packet::Protocol::Pong.into());
net_packet.set_source(current_device.virtual_ip);
net_packet.set_destination(source);
net_packet.first_set_ttl(MAX_TTL);
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.send_by_key(net_packet.buffer(), route_key)?;
let route = Route::from(route_key, metric, 199);
context.route_table.add_route_if_absent(source, route);
}
ControlPacket::PongPacket(pong_packet) => {
let current_time = crate::handle::now_time() as u16;
if current_time < pong_packet.time() {
return Ok(());
}
let rt = (current_time - pong_packet.time()) as i64;
let route = Route::from(route_key, metric, rt);
context.route_table.add_route(source, route);
}
ControlPacket::PunchRequest => {
if self.relay {
return Ok(());
}
//回应
net_packet.set_transport_protocol(control_packet::Protocol::PunchResponse.into());
net_packet.set_source(current_device.virtual_ip);
net_packet.set_destination(source);
net_packet.first_set_ttl(1);
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.send_by_key(net_packet.buffer(), route_key)?;
let route = Route::from(route_key, 1, 199);
context.route_table.add_route_if_absent(source, route);
}
ControlPacket::PunchResponse => {
if self.relay {
return Ok(());
}
let route = Route::from(route_key, 1, 199);
context.route_table.add_route_if_absent(source, route);
}
ControlPacket::AddrRequest => match route_key.addr.ip() {
std::net::IpAddr::V4(ipv4) => {
let mut packet = NetPacket::new_encrypt([0; 12 + 6 + ENCRYPTION_RESERVED])?;
packet.set_version(Version::V1);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::AddrResponse.into());
packet.first_set_ttl(MAX_TTL);
packet.set_source(current_device.virtual_ip);
packet.set_destination(source);
let mut addr_packet = control_packet::AddrPacket::new(packet.payload_mut())?;
addr_packet.set_ipv4(ipv4);
addr_packet.set_port(route_key.addr.port());
self.client_cipher.encrypt_ipv4(&mut packet)?;
context.send_by_key(packet.buffer(), route_key)?;
}
std::net::IpAddr::V6(_) => {}
},
ControlPacket::AddrResponse(_) => {}
}
Ok(())
}
fn other_turn(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
if self.relay {
return Ok(());
}
let source = net_packet.source();
match other_turn_packet::Protocol::from(net_packet.transport_protocol()) {
other_turn_packet::Protocol::Punch => {
let mut punch_info =
PunchInfo::parse_from_bytes(net_packet.payload()).map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("PunchInfo {:?}", e))
})?;
let public_ips = punch_info
.public_ip_list
.iter()
.map(|v| Ipv4Addr::from(v.to_be_bytes()))
.collect();
let local_ipv4 = Some(Ipv4Addr::from(punch_info.local_ip.to_be_bytes()));
let tcp_port = punch_info.tcp_port as u16;
let ipv6 = if punch_info.ipv6.len() == 16 {
let ipv6: [u8; 16] = punch_info.ipv6.try_into().unwrap();
Some(Ipv6Addr::from(ipv6))
} else {
None
};
//兼容旧版本
if punch_info.public_ports.is_empty() {
punch_info.public_ports.push(punch_info.public_port);
}
//兼容旧版本
if punch_info.udp_ports.is_empty() {
punch_info.udp_ports.push(punch_info.local_port);
}
let peer_nat_info = NatInfo::new(
public_ips,
punch_info.public_ports.iter().map(|e| *e as u16).collect(),
punch_info.public_port_range as u16,
local_ipv4,
ipv6,
punch_info.udp_ports.iter().map(|e| *e as u16).collect(),
tcp_port,
punch_info.nat_type.enum_value_or_default().into(),
);
{
let peer_nat_info = peer_nat_info.clone();
self.peer_nat_info_map.write().insert(source, peer_nat_info);
}
if !punch_info.reply {
let mut punch_reply = PunchInfo::new();
punch_reply.reply = true;
let nat_info = self.nat_test.nat_info();
punch_reply.public_ip_list = nat_info
.public_ips
.iter()
.map(|ip| u32::from_be_bytes(ip.octets()))
.collect();
punch_reply.public_port = nat_info.public_ports.get(0).map_or(0, |v| *v as u32);
punch_reply.public_ports =
nat_info.public_ports.iter().map(|e| *e as u32).collect();
punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.tcp_port = nat_info.tcp_port as u32;
punch_reply.nat_type =
protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
punch_reply.local_ip =
u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED));
punch_reply.local_port = nat_info.udp_ports[0] as u32;
punch_reply.udp_ports = nat_info.udp_ports.iter().map(|e| *e as u32).collect();
if let Some(ipv6) = nat_info.ipv6() {
punch_reply.ipv6 = ipv6.octets().to_vec();
punch_reply.ipv6_port = nat_info.udp_ports[0] as u32;
}
let bytes = punch_reply.write_to_bytes().map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("punch_reply {:?}", e))
})?;
let mut punch_packet =
NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
punch_packet.set_version(Version::V1);
punch_packet.set_protocol(Protocol::OtherTurn);
punch_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into());
punch_packet.first_set_ttl(MAX_TTL);
punch_packet.set_source(current_device.virtual_ip());
punch_packet.set_destination(source);
punch_packet.set_payload(&bytes)?;
if self.punch(source, peer_nat_info) {
self.client_cipher.encrypt_ipv4(&mut punch_packet)?;
context.send_by_key(punch_packet.buffer(), route_key)?;
}
} else {
self.punch(source, peer_nat_info);
}
}
other_turn_packet::Protocol::Unknown(e) => {
log::warn!("不支持的转发协议 {:?},source:{:?}", e, source);
}
}
Ok(())
}
fn punch(&self, peer_ip: Ipv4Addr, peer_nat_info: NatInfo) -> bool {
self.punch_sender.try_send((peer_ip, peer_nat_info)).is_ok()
}
}
+152
View File
@@ -0,0 +1,152 @@
use std::collections::HashMap;
use std::net::Ipv4Addr;
use std::sync::mpsc::SyncSender;
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::{Mutex, RwLock};
use tun::Device;
use crate::channel::context::Context;
use crate::channel::handler::RecvChannelHandler;
use crate::channel::punch::NatInfo;
use crate::channel::RouteKey;
use crate::cipher::Cipher;
#[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher;
use crate::external_route::{AllowExternalRoute, ExternalRoute};
use crate::handle::callback::VntCallback;
use crate::handle::recv_data::client::ClientPacketHandler;
use crate::handle::recv_data::server::ServerPacketHandler;
use crate::handle::recv_data::turn::TurnPacketHandler;
use crate::handle::{BaseConfigInfo, CurrentDeviceInfo, PeerDeviceInfo, SELF_IP};
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::IpProxyMap;
use crate::nat::NatTest;
use crate::protocol::NetPacket;
use crate::util::U64Adder;
mod client;
mod server;
mod turn;
#[derive(Clone)]
pub struct RecvDataHandler<Call> {
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
turn: TurnPacketHandler,
client: ClientPacketHandler,
server: ServerPacketHandler<Call>,
counter: U64Adder,
}
impl<Call: VntCallback> RecvChannelHandler for RecvDataHandler<Call> {
fn handle(&mut self, buf: &mut [u8], route_key: RouteKey, context: &Context) {
if let Err(e) = self.handle0(buf, route_key, context) {
log::error!("[{}]-{:?}", thread::current().name().unwrap_or(""), e);
}
}
}
impl<Call: VntCallback> RecvDataHandler<Call> {
pub fn new(
#[cfg(feature = "server_encrypt")] rsa_cipher: Arc<Mutex<Option<RsaCipher>>>,
server_cipher: Cipher,
client_cipher: Cipher,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device: Arc<Device>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
config_info: BaseConfigInfo,
nat_test: NatTest,
callback: Call,
relay: bool,
punch_sender: SyncSender<(Ipv4Addr, NatInfo)>,
peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>>,
external_route: ExternalRoute,
route: AllowExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
counter: U64Adder,
) -> Self {
let server = ServerPacketHandler::new(
#[cfg(feature = "server_encrypt")]
rsa_cipher,
server_cipher,
current_device.clone(),
device.clone(),
device_list,
config_info,
nat_test.clone(),
callback,
external_route,
);
let client = ClientPacketHandler::new(
device.clone(),
client_cipher,
relay,
punch_sender,
peer_nat_info_map,
nat_test,
route,
#[cfg(feature = "ip_proxy")]
ip_proxy_map,
);
let turn = TurnPacketHandler::new();
Self {
current_device,
turn,
client,
server,
counter,
}
}
fn handle0(
&mut self,
buf: &mut [u8],
route_key: RouteKey,
context: &Context,
) -> io::Result<()> {
// 统计流量
self.counter.add(buf.len() as _);
let net_packet = NetPacket::new(buf)?;
if net_packet.ttl() == 0 || net_packet.source_ttl() < net_packet.ttl() {
return Ok(());
}
let current_device = self.current_device.load();
let dest = net_packet.destination();
let source = net_packet.source();
context.route_table.update_read_time(&source, &route_key);
if dest == current_device.virtual_ip
|| dest.is_broadcast()
|| dest.is_multicast()
|| dest == SELF_IP
|| dest.is_unspecified()
|| dest == current_device.broadcast_ip
{
//发给自己的包
if net_packet.is_gateway() {
//服务端-客户端包
self.server
.handle(net_packet, route_key, context, &current_device)
} else {
//客户端-客户端包
self.client
.handle(net_packet, route_key, context, &current_device)
}
} else {
//转发包
self.turn
.handle(net_packet, route_key, context, &current_device)
}
}
}
pub trait PacketHandler {
fn handle(
&self,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
context: &Context,
current_device: &CurrentDeviceInfo,
) -> io::Result<()>;
}
+437
View File
@@ -0,0 +1,437 @@
use std::io;
use std::net::Ipv4Addr;
use std::ops::Sub;
use std::sync::Arc;
use std::time::{Duration, Instant};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use protobuf::Message;
use packet::icmp::{icmp, Kind};
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use tun::device::IFace;
use tun::Device;
use crate::channel::context::Context;
use crate::channel::{Route, RouteKey};
use crate::cipher::Cipher;
#[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher;
use crate::external_route::ExternalRoute;
use crate::handle::callback::{ErrorInfo, ErrorType, HandshakeInfo, RegisterInfo, VntCallback};
use crate::handle::recv_data::PacketHandler;
use crate::handle::{
handshaker, registrar, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo,
GATEWAY_IP,
};
use crate::nat::NatTest;
use crate::proto;
use crate::proto::message::{DeviceList, HandshakeResponse, RegistrationResponse};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::ControlPacket;
use crate::protocol::error_packet::InErrorPacket;
use crate::protocol::{ip_turn_packet, service_packet, NetPacket, Protocol, Version, MAX_TTL};
/// 处理来源于服务端的包
#[derive(Clone)]
pub struct ServerPacketHandler<Call> {
#[cfg(feature = "server_encrypt")]
rsa_cipher: Arc<Mutex<Option<RsaCipher>>>,
server_cipher: Cipher,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device: Arc<Device>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
config_info: BaseConfigInfo,
nat_test: NatTest,
callback: Call,
time: Arc<AtomicCell<Instant>>,
route_record: Arc<Mutex<Vec<(Ipv4Addr, Ipv4Addr)>>>,
external_route: ExternalRoute,
}
impl<Call> ServerPacketHandler<Call> {
pub fn new(
#[cfg(feature = "server_encrypt")] rsa_cipher: Arc<Mutex<Option<RsaCipher>>>,
server_cipher: Cipher,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device: Arc<Device>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
config_info: BaseConfigInfo,
nat_test: NatTest,
callback: Call,
external_route: ExternalRoute,
) -> Self {
Self {
#[cfg(feature = "server_encrypt")]
rsa_cipher,
server_cipher,
current_device,
device,
device_list,
config_info,
nat_test,
callback,
time: Arc::new(AtomicCell::new(Instant::now().sub(Duration::from_secs(60)))),
route_record: Arc::new(Mutex::default()),
external_route,
}
}
}
impl<Call: VntCallback> PacketHandler for ServerPacketHandler<Call> {
fn handle(
&self,
mut net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
context: &Context,
current_device: &CurrentDeviceInfo,
) -> io::Result<()> {
if net_packet.protocol() == Protocol::Error
&& net_packet.transport_protocol()
== crate::protocol::error_packet::Protocol::NoKey.into()
{
//开启服务端加密的情况下,只有这个回应是不加密的,是服务端通知客户端上传密钥
#[cfg(feature = "server_encrypt")]
{
let mutex_guard = self.rsa_cipher.lock();
if let Some(rsa_cipher) = mutex_guard.as_ref() {
let last = self.time.load();
if last.elapsed() < Duration::from_secs(1)
|| self.time.compare_exchange(last, Instant::now()).is_err()
{
//短时间不重复上传服务端密钥
return Ok(());
}
if let Some(key) = self.server_cipher.key() {
log::warn!("上传密钥到服务端:{:?}", route_key);
let packet = handshaker::secret_handshake_request_packet(
rsa_cipher,
self.config_info.token.clone(),
key,
)?;
context.send_by_key(packet.buffer(), route_key)?;
}
}
}
return Ok(());
} else if net_packet.protocol() == Protocol::Service
&& net_packet.transport_protocol() == service_packet::Protocol::HandshakeResponse.into()
{
let response =
HandshakeResponse::parse_from_bytes(net_packet.payload()).map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("HandshakeResponse {:?}", e))
})?;
//如果开启了加密,则发送加密握手请求
#[cfg(feature = "server_encrypt")]
if let Some(key) = self.server_cipher.key() {
let rsa_cipher = RsaCipher::new(&response.public_key)?;
let handshake_info = HandshakeInfo::new(
rsa_cipher.public_key()?.clone(),
rsa_cipher.finger()?,
response.version,
);
log::warn!("加密握手请求:{:?}", handshake_info);
if self.callback.handshake(handshake_info) {
let packet = handshaker::secret_handshake_request_packet(
&rsa_cipher,
self.config_info.token.clone(),
key,
)?;
context.send_by_key(packet.buffer(), route_key)?;
self.rsa_cipher.lock().replace(rsa_cipher);
}
return Ok(());
}
let handshake_info = HandshakeInfo::new_no_secret(response.version);
if self.callback.handshake(handshake_info) {
//没有加密,则发送注册请求
self.register(current_device, context)?;
}
return Ok(());
}
//服务端数据解密
self.server_cipher.decrypt_ipv4(&mut net_packet)?;
match net_packet.protocol() {
Protocol::Service => {
self.service(context, current_device, net_packet, route_key)?;
}
Protocol::Error => {
self.error(context, current_device, net_packet, route_key)?;
}
Protocol::Control => {
self.control(context, current_device, net_packet, route_key)?;
}
Protocol::IpTurn => {
match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) {
ip_turn_packet::Protocol::Ipv4 => {
let ipv4 = IpV4Packet::new(net_packet.payload())?;
match ipv4.protocol() {
ipv4::protocol::Protocol::Icmp => {
if ipv4.destination_ip() == current_device.virtual_ip {
let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?;
if icmp_packet.kind() == Kind::EchoReply {
//网关ip ping的回应
self.device.write(net_packet.payload())?;
return Ok(());
}
}
}
_ => {}
}
}
ip_turn_packet::Protocol::Ipv4Broadcast => {}
ip_turn_packet::Protocol::Unknown(_) => {}
}
}
Protocol::OtherTurn => {}
Protocol::Unknown(_) => {}
}
Ok(())
}
}
impl<Call: VntCallback> ServerPacketHandler<Call> {
fn service(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
match service_packet::Protocol::from(net_packet.transport_protocol()) {
service_packet::Protocol::RegistrationResponse => {
let response = RegistrationResponse::parse_from_bytes(net_packet.payload())
.map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("RegistrationResponse {:?}", e),
)
})?;
let virtual_ip = Ipv4Addr::from(response.virtual_ip);
let virtual_netmask = Ipv4Addr::from(response.virtual_netmask);
let virtual_gateway = Ipv4Addr::from(response.virtual_gateway);
let virtual_network =
Ipv4Addr::from(response.virtual_ip & response.virtual_netmask);
let register_info = RegisterInfo::new(virtual_ip, virtual_netmask, virtual_gateway);
if self.callback.register(register_info) {
let route = Route::from(route_key, 1, 199);
context
.route_table
.add_route_if_absent(virtual_gateway, route);
let old = current_device;
let mut cur = *current_device;
loop {
let mut new_current_device = cur;
new_current_device.update(virtual_ip, virtual_netmask, virtual_gateway);
new_current_device.virtual_ip = virtual_ip;
new_current_device.virtual_netmask = virtual_netmask;
new_current_device.virtual_gateway = virtual_gateway;
if let Err(c) = self
.current_device
.compare_exchange(cur, new_current_device)
{
cur = c;
} else {
break;
}
}
let _ = context.change_status(&self.current_device);
let public_ip = response.public_ip.into();
let public_port = response.public_port as u16;
self.nat_test
.update_addr(route_key.index(), public_ip, public_port);
if old.virtual_ip != virtual_ip
|| old.virtual_gateway != virtual_gateway
|| old.virtual_netmask != virtual_netmask
{
if old.virtual_ip != Ipv4Addr::UNSPECIFIED {
log::info!("ip发生变化,old:{:?},response={:?}", old, response);
}
self.device.set_ip(virtual_ip, virtual_netmask)?;
let mut guard = self.route_record.lock();
for (dest, mask) in guard.drain(..) {
if let Err(e) = self.device.delete_route(dest, mask) {
log::warn!("删除路由失败 ={:?}", e);
}
}
if let Err(e) = self.device.add_route(virtual_network, virtual_netmask, 1) {
log::warn!("添加默认路由失败 ={:?}", e);
} else {
guard.push((virtual_network, virtual_netmask));
}
for (dest, mask) in self.external_route.to_route() {
if let Err(e) = self.device.add_route(dest, mask, 1) {
log::warn!("添加路由失败 ={:?}", e);
} else {
guard.push((dest, mask));
}
}
}
self.set_device_info_list(response.device_info_list, response.epoch as _);
}
}
service_packet::Protocol::RegistrationRequest => {
//不处理注册包
}
service_packet::Protocol::PollDeviceList => {}
service_packet::Protocol::PushDeviceList => {
let response = DeviceList::parse_from_bytes(net_packet.payload()).map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("PushDeviceList {:?}", e))
})?;
self.set_device_info_list(response.device_info_list, response.epoch as _);
}
service_packet::Protocol::HandshakeRequest => {}
service_packet::Protocol::HandshakeResponse => {}
service_packet::Protocol::SecretHandshakeRequest => {}
service_packet::Protocol::SecretHandshakeResponse => {
//加密握手结束,发送注册数据
self.register(current_device, context)?;
}
service_packet::Protocol::Unknown(e) => {
log::warn!("service_packet::Protocol::Unknown = {}", e);
}
}
Ok(())
}
fn set_device_info_list(&self, device_info_list: Vec<proto::message::DeviceInfo>, epoch: u16) {
let ip_list: Vec<PeerDeviceInfo> = device_info_list
.into_iter()
.map(|info| {
PeerDeviceInfo::new(
Ipv4Addr::from(info.virtual_ip),
info.name,
info.device_status as u8,
info.client_secret,
)
})
.collect();
let mut dev = self.device_list.lock();
//这里可能会收到旧的消息,但是随着时间推移总会收到新的
dev.0 = epoch;
dev.1 = ip_list;
}
fn register(&self, current_device: &CurrentDeviceInfo, context: &Context) -> io::Result<()> {
if current_device.status == ConnectStatus::Connected {
//已连接的不需要注册
return Ok(());
}
let token = self.config_info.token.clone();
let device_id = self.config_info.device_id.clone();
let name = self.config_info.name.clone();
let client_secret = self.config_info.client_secret;
let ip = self.config_info.ip;
let response = registrar::registration_request_packet(
&self.server_cipher,
token,
device_id,
name,
ip,
false,
false,
client_secret,
)?;
//注册请求只发送到默认通道
context.send_default(response.buffer(), current_device.connect_server)
}
fn error(
&self,
_context: &Context,
_current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
_route_key: RouteKey,
) -> io::Result<()> {
match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
InErrorPacket::TokenError => {
// token错误,可能是服务端设置了白名单
let err = ErrorInfo::new(ErrorType::TokenError);
self.callback.error(err);
}
InErrorPacket::Disconnect => {
let err = ErrorInfo::new(ErrorType::Disconnect);
self.callback.error(err);
//掉线epoch要归零
{
let mut dev = self.device_list.lock();
dev.0 = 0;
drop(dev);
}
// self.register(current_device, context, route_key)?;
}
InErrorPacket::AddressExhausted => {
// 地址用尽
let err = ErrorInfo::new(ErrorType::AddressExhausted);
self.callback.error(err);
}
InErrorPacket::OtherError(e) => {
let err = ErrorInfo::new_msg(ErrorType::Unknown, e.message()?);
self.callback.error(err);
}
InErrorPacket::IpAlreadyExists => {
log::error!("IpAlreadyExists");
let err = ErrorInfo::new(ErrorType::IpAlreadyExists);
self.callback.error(err);
}
InErrorPacket::InvalidIp => {
log::error!("InvalidIp");
let err = ErrorInfo::new(ErrorType::InvalidIp);
self.callback.error(err);
}
InErrorPacket::NoKey => {
//这个类型最开头已经处理过,这里忽略
}
}
Ok(())
}
fn control(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PongPacket(pong_packet) => {
let current_time = crate::handle::now_time() as u16;
if current_time < pong_packet.time() {
return Ok(());
}
let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
let rt = (current_time - pong_packet.time()) as i64;
let route = Route::from(route_key, metric, rt);
context.route_table.add_route(net_packet.source(), route);
let epoch = self.device_list.lock().0;
if pong_packet.epoch() != epoch {
//纪元不一致,可能有新客户端连接,向服务端拉取客户端列表
let mut poll_device = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED])?;
poll_device.set_source(current_device.virtual_ip);
poll_device.set_destination(GATEWAY_IP);
poll_device.set_version(Version::V1);
poll_device.set_gateway_flag(true);
poll_device.first_set_ttl(MAX_TTL);
poll_device.set_protocol(Protocol::Service);
poll_device
.set_transport_protocol(service_packet::Protocol::PollDeviceList.into());
self.server_cipher.encrypt_ipv4(&mut poll_device)?;
//发送到默认服务端即可
context.send_default(poll_device.buffer(), current_device.connect_server)?;
}
}
ControlPacket::AddrResponse(addr_packet) => {
//更新本地公网ipv4
self.nat_test.update_addr(
route_key.index(),
addr_packet.ipv4(),
addr_packet.port(),
);
}
_ => {}
}
Ok(())
}
}
+39
View File
@@ -0,0 +1,39 @@
use crate::channel::context::Context;
use crate::channel::RouteKey;
use crate::handle::recv_data::PacketHandler;
use crate::handle::CurrentDeviceInfo;
use crate::protocol::NetPacket;
/// 处理客户端中转包
#[derive(Clone)]
pub struct TurnPacketHandler {}
impl TurnPacketHandler {
pub fn new() -> Self {
Self {}
}
}
impl PacketHandler for TurnPacketHandler {
fn handle(
&self,
mut net_packet: NetPacket<&mut [u8]>,
_route_key: RouteKey,
context: &Context,
_current_device: &CurrentDeviceInfo,
) -> std::io::Result<()> {
// ttl减一
let ttl = net_packet.incr_ttl();
if ttl > 0 {
let destination = net_packet.destination();
if let Some(route) = context.route_table.route_one(&destination) {
if route.metric <= ttl {
context.send_by_key(net_packet.buffer(), route.route_key())?;
}
}
//其他没有路由的不转发
}
Ok(())
}
}
+49
View File
@@ -0,0 +1,49 @@
use std::io;
use std::net::Ipv4Addr;
use protobuf::Message;
use crate::cipher::Cipher;
use crate::handle::{GATEWAY_IP, SELF_IP};
use crate::proto::message::RegistrationRequest;
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
/// 注册数据
pub fn registration_request_packet(
server_cipher: &Cipher,
token: String,
device_id: String,
name: String,
ip: Option<Ipv4Addr>,
is_fast: bool,
allow_ip_change: bool,
client_secret: bool,
) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = RegistrationRequest::new();
request.token = token;
request.device_id = device_id;
request.name = name;
if let Some(ip) = ip {
request.virtual_ip = ip.into();
}
request.allow_ip_change = allow_ip_change;
request.is_fast = is_fast;
request.version = crate::VNT_VERSION.to_string();
request.client_secret = client_secret;
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("RegistrationRequest {:?}", e))
})?;
let buf = vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED];
let mut net_packet = NetPacket::new_encrypt(buf)?;
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::RegistrationRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
server_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
+24 -24
View File
@@ -1,30 +1,30 @@
#[derive(Clone)]
pub struct BufSenderGroup(
usize,
Vec<std::sync::mpsc::SyncSender<(Vec<u8>, usize, usize)>>,
);
use std::sync::mpsc::{sync_channel, Receiver, SendError, SyncSender};
pub struct BufReceiverGroup(pub Vec<std::sync::mpsc::Receiver<(Vec<u8>, usize, usize)>>);
impl BufSenderGroup {
pub fn send(&mut self, val: (Vec<u8>, usize, usize)) -> bool {
let index = self.0 % self.1.len();
self.0 = self.0.wrapping_add(1);
self.1[index].send(val).is_ok()
}
}
pub fn buf_channel_group(size: usize) -> (BufSenderGroup, BufReceiverGroup) {
let mut buf_sender_group = Vec::with_capacity(size);
let mut buf_receiver_group = Vec::with_capacity(size);
pub fn channel_group<T>(size: usize, bound: usize) -> (GroupSyncSender<T>, Vec<Receiver<T>>) {
let mut senders = Vec::with_capacity(size);
let mut receivers = Vec::with_capacity(size);
for _ in 0..size {
let (buf_sender, buf_receiver) =
std::sync::mpsc::sync_channel::<(Vec<u8>, usize, usize)>(1);
buf_sender_group.push(buf_sender);
buf_receiver_group.push(buf_receiver);
let (s, r) = sync_channel(bound);
senders.push(s);
receivers.push(r);
}
(
BufSenderGroup(0, buf_sender_group),
BufReceiverGroup(buf_receiver_group),
GroupSyncSender {
count: 0,
base: senders,
},
receivers,
)
}
pub struct GroupSyncSender<T> {
count: usize,
base: Vec<SyncSender<T>>,
}
impl<T> GroupSyncSender<T> {
pub fn send(&mut self, t: T) -> Result<(), SendError<T>> {
self.count += 1;
self.base[self.count % self.base.len()].send(t)
}
}
+26 -72
View File
@@ -1,17 +1,13 @@
use std::io;
use std::net::Ipv4Addr;
use std::sync::Arc;
use parking_lot::RwLock;
use crate::channel::context::Context;
use packet::ip::ipv4::packet::IpV4Packet;
use packet::ip::ipv4::protocol::Protocol;
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::error::*;
use crate::external_route::ExternalRoute;
use crate::handle::{check_dest, CurrentDeviceInfo};
use crate::igmp_server::{IgmpServer, Multicast};
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::{IpProxyMap, ProxyHandler};
use crate::protocol;
@@ -19,20 +15,17 @@ use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::ip_turn_packet::BroadcastPacket;
use crate::protocol::{ip_turn_packet, NetPacket, Version, MAX_TTL};
pub mod channel_group;
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
pub mod tap_handler;
mod channel_group;
pub mod tun_handler;
fn broadcast(
server_cipher: &Cipher,
multicast_members: Option<Arc<RwLock<Multicast>>>,
sender: &ChannelSender,
sender: &Context,
net_packet: &mut NetPacket<&mut [u8]>,
current_device: &CurrentDeviceInfo,
) -> Result<()> {
) -> io::Result<()> {
let mut peer_ips = Vec::with_capacity(8);
let vec = sender.route_table_one();
let vec = sender.route_table.route_table_one();
let mut relay_count = 0;
const MAX_COUNT: usize = 8;
for (peer_ip, route) in vec {
@@ -42,14 +35,9 @@ fn broadcast(
if peer_ips.len() == MAX_COUNT {
break;
}
if let Some(members) = &multicast_members {
if !members.read().is_send(&peer_ip) {
continue;
}
}
if route.is_p2p()
&& sender
.try_send_by_key(net_packet.buffer(), &route.route_key())
.send_by_key(net_packet.buffer(), route.route_key())
.is_ok()
{
peer_ips.push(peer_ip);
@@ -63,7 +51,7 @@ fn broadcast(
}
//转发到服务端的可选择广播,还要进行服务端加密
if peer_ips.is_empty() {
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
sender.send_default(net_packet.buffer(), current_device.connect_server)?;
} else {
let buf =
vec![0u8; 12 + 1 + peer_ips.len() * 4 + net_packet.data_len() + ENCRYPTION_RESERVED];
@@ -82,7 +70,7 @@ fn broadcast(
broadcast.set_address(&peer_ips)?;
broadcast.set_data(net_packet.buffer())?;
server_cipher.encrypt_ipv4(&mut server_packet)?;
sender.send_main(server_packet.buffer(), current_device.connect_server)?;
sender.send_default(server_packet.buffer(), current_device.connect_server)?;
}
Ok(())
}
@@ -92,16 +80,15 @@ fn broadcast(
///
#[inline]
pub fn base_handle(
sender: &ChannelSender,
context: &Context,
buf: &mut [u8],
data_len: usize, //数据总长度=12+ip包长度
igmp_server: &Option<IgmpServer>,
current_device: CurrentDeviceInfo,
ip_route: &Option<ExternalRoute>,
ip_route: &ExternalRoute,
#[cfg(feature = "ip_proxy")] proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> Result<()> {
) -> io::Result<()> {
let ipv4_packet = IpV4Packet::new(&buf[12..data_len])?;
let protocol = ipv4_packet.protocol();
let src_ip = ipv4_packet.source_ip();
@@ -117,52 +104,26 @@ pub fn base_handle(
if protocol == Protocol::Icmp {
net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?;
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
context.send_default(net_packet.buffer(), current_device.connect_server)?;
}
return Ok(());
}
if dest_ip.is_multicast() {
match protocol {
Protocol::Igmp => {
if igmp_server.is_some() {
//发送到服务端
net_packet.set_destination(current_device.virtual_gateway);
net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?;
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
}
}
Protocol::Udp => {
let multicast_members = if let Some(igmp_server) = igmp_server {
igmp_server.load(&dest_ip)
} else {
//当作广播处理
net_packet.set_destination(Ipv4Addr::BROADCAST);
None
};
//当作广播处理
net_packet.set_destination(Ipv4Addr::BROADCAST);
client_cipher.encrypt_ipv4(&mut net_packet)?;
broadcast(
server_cipher,
multicast_members,
sender,
&mut net_packet,
&current_device,
)?;
broadcast(server_cipher, context, &mut net_packet, &current_device)?;
}
_ => {}
}
return Ok(());
}
if dest_ip.is_broadcast() || current_device.broadcast_address == dest_ip {
if dest_ip.is_broadcast() || current_device.broadcast_ip == dest_ip {
// 广播 发送到直连目标
client_cipher.encrypt_ipv4(&mut net_packet)?;
broadcast(
server_cipher,
None,
sender,
&mut net_packet,
&current_device,
)?;
broadcast(server_cipher, context, &mut net_packet, &current_device)?;
return Ok(());
}
if !check_dest(
@@ -170,18 +131,14 @@ pub fn base_handle(
current_device.virtual_netmask,
current_device.virtual_network,
) {
if let Some(ip_route) = ip_route {
if let Some(r_dest_ip) = ip_route.route(&dest_ip) {
//路由的目标不能是自己
if r_dest_ip == src_ip {
return Ok(());
}
//需要修改目的地址
dest_ip = r_dest_ip;
net_packet.set_destination(r_dest_ip);
} else {
if let Some(r_dest_ip) = ip_route.route(&dest_ip) {
//路由的目标不能是自己
if r_dest_ip == src_ip {
return Ok(());
}
//需要修改目的地址
dest_ip = r_dest_ip;
net_packet.set_destination(r_dest_ip);
} else {
return Ok(());
}
@@ -193,11 +150,8 @@ pub fn base_handle(
}
client_cipher.encrypt_ipv4(&mut net_packet)?;
//优先发到直连到地址
if sender
.try_send_by_id(net_packet.buffer(), &dest_ip)
.is_err()
{
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
if context.send_by_id(net_packet.buffer(), &dest_ip).is_err() {
context.send_default(net_packet.buffer(), current_device.connect_server)?;
}
return Ok(());
}
+126 -117
View File
@@ -1,25 +1,26 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use packet::icmp::icmp::IcmpPacket;
use packet::icmp::Kind;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use tun::device::IFace;
use tun::Device;
use crate::error::*;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::external_route::ExternalRoute;
use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup};
use crate::handle::tun_tap::channel_group::{channel_group, GroupSyncSender};
use crate::handle::CurrentDeviceInfo;
use crate::igmp_server::IgmpServer;
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::IpProxyMap;
use crate::tun_tap_device::{DeviceReader, DeviceWriter};
fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> Result<()> {
use crate::util::StopManager;
fn icmp(device_writer: &Device, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> io::Result<()> {
if ipv4_packet.protocol() == ipv4::protocol::Protocol::Icmp {
let mut icmp = IcmpPacket::new(ipv4_packet.payload_mut())?;
if icmp.kind() == Kind::EchoRequest {
@@ -29,26 +30,28 @@ fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) ->
ipv4_packet.set_source_ip(ipv4_packet.destination_ip());
ipv4_packet.set_destination_ip(src);
ipv4_packet.update_checksum();
device_writer.write_ipv4_tun(ipv4_packet.buffer)?;
device_writer.write(ipv4_packet.buffer)?;
}
}
Ok(())
}
/// 接收tun数据,并且转发到udp上
#[inline]
fn handle(
sender: &ChannelSender,
context: &Context,
data: &mut [u8],
len: usize,
device_writer: &DeviceWriter,
igmp_server: &Option<IgmpServer>,
device_writer: &Device,
current_device: CurrentDeviceInfo,
ip_route: &Option<ExternalRoute>,
ip_route: &ExternalRoute,
#[cfg(feature = "ip_proxy")] proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> Result<()> {
) -> io::Result<()> {
if len > 12 && data[12] >> 4 != 4 {
//忽略非ipv4包
return Ok(());
}
let ipv4_packet = IpV4Packet::new(&mut data[12..len])?;
let src_ip = ipv4_packet.source_ip();
let dest_ip = ipv4_packet.destination_ip();
@@ -56,10 +59,9 @@ fn handle(
return icmp(&device_writer, ipv4_packet);
}
return crate::handle::tun_tap::base_handle(
sender,
context,
data,
len,
igmp_server,
current_device,
ip_route,
#[cfg(feature = "ip_proxy")]
@@ -70,143 +72,127 @@ fn handle(
}
pub fn start(
worker: VntWorker,
sender: ChannelSender,
device_reader: DeviceReader,
device_writer: DeviceWriter,
igmp_server: Option<IgmpServer>,
stop_manager: StopManager,
context: Context,
device: Arc<Device>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_route: Option<ExternalRoute>,
ip_route: ExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
parallel: usize,
) {
if parallel == 1 {
thread::Builder::new()
.name("tun_handler".into())
.spawn(move || {
if let Err(e) = start_simple(
&sender,
device_reader,
&device_writer,
igmp_server,
current_device,
ip_route,
#[cfg(feature = "ip_proxy")]
ip_proxy_map,
client_cipher,
server_cipher,
) {
log::warn!("stop:{}", e);
}
let _ = sender.close();
let _ = device_writer.close();
worker.stop_all();
})
.unwrap();
} else {
let (buf_sender, buf_receiver) = buf_channel_group(parallel);
for buf_receiver in buf_receiver.0 {
let sender = sender.clone();
let device_writer = device_writer.clone();
let igmp_server = igmp_server.clone();
up_counter: Arc<AtomicU64>,
) -> io::Result<()> {
let worker = {
let device = device.clone();
stop_manager.add_listener("tun_device".into(), move || {
if let Err(e) = device.shutdown() {
log::warn!("{:?}", e);
}
})?
};
if parallel > 1 {
let (sender, receivers) = channel_group::<(Vec<u8>, usize)>(parallel, 16);
for (index, receiver) in receivers.into_iter().enumerate() {
let context = context.clone();
let device = device.clone();
let current_device = current_device.clone();
let ip_route = ip_route.clone();
#[cfg(feature = "ip_proxy")]
let ip_proxy_map = ip_proxy_map.clone();
let client_cipher = client_cipher.clone();
let server_cipher = server_cipher.clone();
thread::spawn(move || {
while let Ok((mut buf, start, len)) = buf_receiver.recv() {
match handle(
&sender,
&mut buf[start..],
len,
&device_writer,
&igmp_server,
current_device.load(),
&ip_route,
#[cfg(feature = "ip_proxy")]
&ip_proxy_map,
&client_cipher,
&server_cipher,
) {
Ok(_) => {}
Err(e) => {
log::warn!("{:?}", e)
thread::Builder::new()
.name(format!("tun_handler_{}", index))
.spawn(move || {
while let Ok((mut buf, len)) = receiver.recv() {
#[cfg(not(target_os = "macos"))]
let start = 0;
#[cfg(target_os = "macos")]
let start = 4;
match handle(
&context,
&mut buf[start..],
len,
&device,
current_device.load(),
&ip_route,
#[cfg(feature = "ip_proxy")]
&ip_proxy_map,
&client_cipher,
&server_cipher,
) {
Ok(_) => {}
Err(e) => {
log::warn!("{:?}", e)
}
}
}
}
let _ = sender.close();
let _ = device_writer.close();
});
})?;
}
thread::Builder::new()
.name("tun_handler".into())
.spawn(move || {
if let Err(e) = start_(&sender, device_reader, buf_sender) {
if let Err(e) = start_multi(stop_manager, device, sender, &up_counter) {
log::warn!("stop:{}", e);
}
let _ = sender.close();
let _ = device_writer.close();
worker.stop_all();
})
.unwrap();
}
}
fn start_(
sender: &ChannelSender,
device_reader: DeviceReader,
mut buf_sender: BufSenderGroup,
) -> io::Result<()> {
loop {
let mut buf = vec![0; 4096];
buf[..12].fill(0);
if sender.is_close() {
return Ok(());
}
let start = 0;
let len = device_reader.read(&mut buf[12..])? + 12;
#[cfg(any(target_os = "macos"))]
let start = 4;
if !buf_sender.send((buf, start, len)) {
return Err(io::Error::new(
io::ErrorKind::Other,
"tun buf_sender发送失败",
));
}
})?;
} else {
thread::Builder::new()
.name("tun_handler".into())
.spawn(move || {
if let Err(e) = start_simple(
stop_manager,
&context,
device,
current_device,
ip_route,
#[cfg(feature = "ip_proxy")]
ip_proxy_map,
client_cipher,
server_cipher,
&up_counter,
) {
log::warn!("stop:{}", e);
}
worker.stop_all();
})?;
}
Ok(())
}
fn start_simple(
sender: &ChannelSender,
device_reader: DeviceReader,
device_writer: &DeviceWriter,
igmp_server: Option<IgmpServer>,
stop_manager: StopManager,
context: &Context,
device: Arc<Device>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_route: Option<ExternalRoute>,
ip_route: ExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
up_counter: &AtomicU64,
) -> io::Result<()> {
let mut buf = [0; 4096];
let mut buf = [0; 1024 * 16];
loop {
if sender.is_close() {
if stop_manager.is_stop() {
return Ok(());
}
buf[..12].fill(0);
let len = device_reader.read(&mut buf[12..])? + 12;
let len = device.read(&mut buf[12..])? + 12;
//单线程的
up_counter.store(
up_counter.load(Ordering::Relaxed) + len as u64,
Ordering::Relaxed,
);
#[cfg(any(target_os = "macos"))]
let mut buf = &mut buf[4..];
// buf是重复利用的,需要重置头部
buf[..12].fill(0);
match handle(
sender,
context,
&mut buf,
len,
device_writer,
&igmp_server,
&device,
current_device.load(),
&ip_route,
#[cfg(feature = "ip_proxy")]
@@ -221,3 +207,26 @@ fn start_simple(
}
}
}
fn start_multi(
stop_manager: StopManager,
device: Arc<Device>,
mut group_sync_sender: GroupSyncSender<(Vec<u8>, usize)>,
up_counter: &AtomicU64,
) -> io::Result<()> {
loop {
if stop_manager.is_stop() {
return Ok(());
}
let mut buf = vec![0; 1024 * 16];
let len = device.read(&mut buf[12..])? + 12;
//单线程的
up_counter.store(
up_counter.load(Ordering::Relaxed) + len as u64,
Ordering::Relaxed,
);
if group_sync_sender.send((buf, len)).is_err() {
return Ok(());
}
}
}