v2
This commit is contained in:
@@ -0,0 +1,94 @@
|
||||
use crate::context::config::Config;
|
||||
use crate::context::{AppState, NetworkAddr, PacketLossInfo, ServerNodeInfo, TrafficInfo};
|
||||
use crate::protocol::control_message::ClientSimpleInfo;
|
||||
use crate::tunnel_core::p2p::route_table::Route;
|
||||
use crate::tunnel_core::server::rpc::ServerRPC;
|
||||
use rust_p2p_core::nat::NatInfo;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct VntApi {
|
||||
app_state: AppState,
|
||||
server_rpc: ServerRPC,
|
||||
}
|
||||
|
||||
impl VntApi {
|
||||
pub(crate) fn new(app_state: AppState, server_rpc: ServerRPC) -> Self {
|
||||
Self {
|
||||
app_state,
|
||||
server_rpc,
|
||||
}
|
||||
}
|
||||
pub fn server_rpc(&self) -> &ServerRPC {
|
||||
&self.server_rpc
|
||||
}
|
||||
/// 获取启动配置
|
||||
pub fn get_config(&self) -> Option<Box<Config>> {
|
||||
self.app_state.get_config()
|
||||
}
|
||||
/// 获取所有客户端ip
|
||||
pub fn client_ips(&self) -> Vec<ClientSimpleInfo> {
|
||||
self.app_state.client_ips()
|
||||
}
|
||||
/// 判断目标IP是否直连
|
||||
pub fn is_direct(&self, ip: &Ipv4Addr) -> bool {
|
||||
self.app_state.route_table.p2p_num(ip) > 0
|
||||
}
|
||||
/// 查找路由
|
||||
pub fn find_route(&self, ip: &Ipv4Addr) -> Option<Route> {
|
||||
self.app_state.route_table.get_route_by_id(ip).ok()
|
||||
}
|
||||
pub fn get_rtt(&self, ip: &Ipv4Addr) -> Option<u32> {
|
||||
if let Some(route) = self.find_route(ip) {
|
||||
Some(route.rtt())
|
||||
} else {
|
||||
self.server_node_rtt(ip).map(|v| v * 2)
|
||||
}
|
||||
}
|
||||
/// 获取所有路由
|
||||
pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec<Route>)> {
|
||||
self.app_state.route_table.route_table()
|
||||
}
|
||||
/// 获取服务器节点
|
||||
pub fn server_node_list(&self) -> Vec<ServerNodeInfo> {
|
||||
self.app_state.server_info_collection.server_node_list()
|
||||
}
|
||||
pub fn server_node_rtt(&self, ip: &Ipv4Addr) -> Option<u32> {
|
||||
self.app_state.server_info_collection.get_server_rtt(ip)
|
||||
}
|
||||
/// 获取网络配置
|
||||
pub fn network(&self) -> Option<NetworkAddr> {
|
||||
self.app_state.get_network()
|
||||
}
|
||||
/// 获取当前的nat信息
|
||||
pub fn nat_info(&self) -> Option<NatInfo> {
|
||||
self.app_state.get_nat_info()
|
||||
}
|
||||
pub fn peer_nat_info(&self, ip: &Ipv4Addr) -> Option<NatInfo> {
|
||||
self.app_state.get_peer_info(ip).and_then(|v| v.nat_info)
|
||||
}
|
||||
pub fn packet_loss_info(&self, ip: &Ipv4Addr) -> Option<PacketLossInfo> {
|
||||
self.app_state.packet_loss_stats.get_loss_info(ip)
|
||||
}
|
||||
pub fn all_packet_loss_info(&self) -> Vec<PacketLossInfo> {
|
||||
self.app_state.packet_loss_stats.get_all_loss_info()
|
||||
}
|
||||
pub fn reset_packet_loss(&self, ip: &Ipv4Addr) {
|
||||
self.app_state.packet_loss_stats.reset(ip)
|
||||
}
|
||||
pub fn reset_all_packet_loss(&self) {
|
||||
self.app_state.packet_loss_stats.reset_all()
|
||||
}
|
||||
pub fn traffic_info(&self, ip: &Ipv4Addr) -> Option<TrafficInfo> {
|
||||
self.app_state.traffic_stats.get_traffic_info(ip)
|
||||
}
|
||||
pub fn all_traffic_info(&self) -> Vec<TrafficInfo> {
|
||||
self.app_state.traffic_stats.get_all_traffic_info()
|
||||
}
|
||||
pub fn reset_traffic(&self, ip: &Ipv4Addr) {
|
||||
self.app_state.traffic_stats.reset(ip)
|
||||
}
|
||||
pub fn reset_all_traffic(&self) {
|
||||
self.app_state.traffic_stats.reset_all()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use std::io;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct LZ4Compression {
|
||||
min_size: usize, // 只压缩大于此大小的数据包
|
||||
}
|
||||
|
||||
impl LZ4Compression {
|
||||
pub fn new() -> Self {
|
||||
Self::with_min_size(256)
|
||||
}
|
||||
pub fn with_min_size(min_size: usize) -> Self {
|
||||
Self { min_size }
|
||||
}
|
||||
|
||||
/// 压缩数据包,返回新的压缩后的数据包
|
||||
/// reserve: 尾部预留空间(用于后续加密等操作)
|
||||
pub fn compress(
|
||||
&self,
|
||||
pkt: NetPacket<TransmissionBytes>,
|
||||
reserve: usize,
|
||||
) -> io::Result<NetPacket<TransmissionBytes>> {
|
||||
let payload = pkt.payload();
|
||||
if payload.len() < self.min_size {
|
||||
return Ok(pkt);
|
||||
}
|
||||
let compressed = lz4_flex::compress_prepend_size(payload);
|
||||
if compressed.len() >= payload.len() {
|
||||
return Ok(pkt);
|
||||
}
|
||||
let total_len = HEAD_LENGTH + compressed.len();
|
||||
let mut buf = TransmissionBytes::zeroed_size(total_len, reserve);
|
||||
|
||||
buf[..HEAD_LENGTH].copy_from_slice(&pkt.buffer()[..HEAD_LENGTH]);
|
||||
buf[HEAD_LENGTH..total_len].copy_from_slice(&compressed);
|
||||
|
||||
let mut packet = NetPacket::new(buf)?;
|
||||
packet.set_compressed_flag(true);
|
||||
Ok(packet)
|
||||
}
|
||||
|
||||
/// 解压缩数据包,返回新的解压后的数据包
|
||||
/// reserve: 尾部预留空间(用于后续加密等操作)
|
||||
pub fn decompress(
|
||||
&self,
|
||||
pkt: NetPacket<TransmissionBytes>,
|
||||
) -> io::Result<NetPacket<TransmissionBytes>> {
|
||||
if !pkt.is_compressed() {
|
||||
return Ok(pkt);
|
||||
}
|
||||
let payload = pkt.payload();
|
||||
|
||||
let decompressed = lz4_flex::decompress_size_prepended(payload).map_err(|e| {
|
||||
io::Error::new(
|
||||
io::ErrorKind::InvalidData,
|
||||
format!("decompress failed: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
let total_len = HEAD_LENGTH + decompressed.len();
|
||||
let mut buf = TransmissionBytes::zeroed(total_len);
|
||||
|
||||
buf[..HEAD_LENGTH].copy_from_slice(&pkt.buffer()[..HEAD_LENGTH]);
|
||||
buf[HEAD_LENGTH..total_len].copy_from_slice(&decompressed);
|
||||
|
||||
let mut packet = NetPacket::new(buf)?;
|
||||
packet.set_compressed_flag(false);
|
||||
Ok(packet)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
|
||||
fn make_packet(data: &[u8]) -> NetPacket<TransmissionBytes> {
|
||||
let mut buf = TransmissionBytes::zeroed(HEAD_LENGTH + data.len());
|
||||
buf[HEAD_LENGTH..HEAD_LENGTH + data.len()].copy_from_slice(data);
|
||||
NetPacket::new(buf).unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lz4_compress_and_decompress() {
|
||||
let lz = LZ4Compression::with_min_size(10);
|
||||
|
||||
// --- 构造原始包 ---
|
||||
let payload = vec![1u8; 200];
|
||||
let original = make_packet(&payload);
|
||||
|
||||
// --- 压缩 ---
|
||||
let compressed = lz.compress(original, 0).unwrap();
|
||||
assert!(compressed.is_compressed());
|
||||
|
||||
// 压缩后的 payload 应变小
|
||||
assert!(
|
||||
compressed.payload().len() < payload.len(),
|
||||
"压缩后 payload 应该更小"
|
||||
);
|
||||
|
||||
// --- 解压 ---
|
||||
let decompressed = lz.decompress(compressed).unwrap();
|
||||
|
||||
// 标志应清除
|
||||
assert!(!decompressed.is_compressed());
|
||||
|
||||
// HEAD 不变
|
||||
assert_eq!(
|
||||
&decompressed.buffer()[..HEAD_LENGTH],
|
||||
&[0u8; HEAD_LENGTH][..],
|
||||
"HEAD 必须保持不变"
|
||||
);
|
||||
|
||||
// payload 必须等于原始 payload
|
||||
assert_eq!(decompressed.payload(), &payload[..]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_no_compress_when_small() {
|
||||
let lz = LZ4Compression::with_min_size(100);
|
||||
|
||||
let pkt = make_packet(&[7; 20]);
|
||||
let compressed = lz.compress(pkt, 0).unwrap();
|
||||
|
||||
assert!(!compressed.is_compressed(), "小包不应该被压缩");
|
||||
|
||||
assert_eq!(compressed.payload(), &[7; 20][..]);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
use crate::compression::lz4_compression::LZ4Compression;
|
||||
use crate::protocol::ip_packet_protocol::NetPacket;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use std::io;
|
||||
|
||||
mod lz4_compression;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct PacketCompression {
|
||||
compression: Option<LZ4Compression>,
|
||||
}
|
||||
|
||||
impl PacketCompression {
|
||||
pub(crate) fn new(enabled: bool) -> Self {
|
||||
Self {
|
||||
compression: if enabled {
|
||||
Some(LZ4Compression::new())
|
||||
} else {
|
||||
None
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn compress(
|
||||
&self,
|
||||
pkt: NetPacket<TransmissionBytes>,
|
||||
reserve: usize,
|
||||
) -> io::Result<NetPacket<TransmissionBytes>> {
|
||||
if let Some(compression) = self.compression.as_ref() {
|
||||
return compression.compress(pkt, reserve);
|
||||
}
|
||||
|
||||
Ok(pkt)
|
||||
}
|
||||
|
||||
pub(crate) fn decompress(
|
||||
&self,
|
||||
pkt: NetPacket<TransmissionBytes>,
|
||||
) -> io::Result<NetPacket<TransmissionBytes>> {
|
||||
if let Some(compression) = self.compression.as_ref() {
|
||||
return compression.decompress(pkt);
|
||||
}
|
||||
Ok(pkt)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::nat::NetInput;
|
||||
use crate::port_mapping::PortMapping;
|
||||
use crate::tls::verifier::CertValidationMode;
|
||||
use crate::tunnel_core::server::transport::config::{ConnectRegConfig, ProtocolAddress};
|
||||
use anyhow::bail;
|
||||
use ipnet::Ipv4Net;
|
||||
use std::collections::HashSet;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
pub const MAX_NETWORK_CODE_LEN: usize = 32;
|
||||
pub const MAX_DEVICE_ID_LEN: usize = 64;
|
||||
pub const MAX_NAME_LEN: usize = 128;
|
||||
pub const MAX_VERSION_LEN: usize = 32;
|
||||
pub const MAX_MTU: u16 = 1500;
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct Config {
|
||||
pub server_addr: Vec<ProtocolAddress>,
|
||||
pub cert_mode: CertValidationMode,
|
||||
pub network_code: String,
|
||||
pub device_id: String,
|
||||
pub device_name: String,
|
||||
pub tun_name: Option<String>,
|
||||
pub ip: Option<Ipv4Addr>,
|
||||
pub password: Option<String>,
|
||||
pub no_punch: bool,
|
||||
pub compress: bool,
|
||||
pub rtx: bool,
|
||||
pub fec: bool,
|
||||
pub input: Vec<NetInput>,
|
||||
pub output: Vec<Ipv4Net>,
|
||||
pub no_nat: bool,
|
||||
pub no_tun: bool,
|
||||
pub mtu: Option<u16>,
|
||||
pub port_mapping: Vec<PortMapping>,
|
||||
pub allow_port_mapping: bool,
|
||||
pub udp_stun: Vec<String>,
|
||||
pub tcp_stun: Vec<String>,
|
||||
}
|
||||
impl Config {
|
||||
pub fn check(&self) -> anyhow::Result<()> {
|
||||
if self.server_addr.is_empty() {
|
||||
bail!("服务器地址不能为空");
|
||||
}
|
||||
if self.server_addr.len() > 1 {
|
||||
let mut set = HashSet::new();
|
||||
|
||||
for a in self.server_addr.iter() {
|
||||
if !set.insert(a.address.as_str()) {
|
||||
bail!("服务器地址不能相同")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if self.network_code.len() > MAX_NETWORK_CODE_LEN {
|
||||
bail!(
|
||||
"network_code length exceeds {} characters (current: {})",
|
||||
MAX_NETWORK_CODE_LEN,
|
||||
self.network_code.len()
|
||||
)
|
||||
}
|
||||
|
||||
if self.device_id.len() > MAX_DEVICE_ID_LEN {
|
||||
bail!(
|
||||
"device_id length exceeds {} characters (current: {})",
|
||||
MAX_DEVICE_ID_LEN,
|
||||
self.device_id.len()
|
||||
)
|
||||
}
|
||||
|
||||
if self.device_name.len() > MAX_NAME_LEN {
|
||||
bail!(
|
||||
"name length exceeds {} characters (current: {})",
|
||||
MAX_NAME_LEN,
|
||||
self.device_name.len()
|
||||
)
|
||||
}
|
||||
if let Some(mtu) = self.mtu
|
||||
&& mtu > MAX_MTU
|
||||
{
|
||||
bail!("MTU is too large (Maximum mtu: {MAX_MTU})",)
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
pub fn key_sign(&self) -> Option<String> {
|
||||
self.password.as_ref().map(|p| PacketCrypto::key_sign(p))
|
||||
}
|
||||
pub(crate) fn to_connect_config(&self, index: usize) -> ConnectRegConfig {
|
||||
ConnectRegConfig {
|
||||
server_addr: self.server_addr[index].clone(),
|
||||
cert_mode: self.cert_mode.clone(),
|
||||
network_code: self.network_code.clone(),
|
||||
device_id: self.device_id.clone(),
|
||||
device_name: self.device_name.clone(),
|
||||
ip: self.ip,
|
||||
key_sign: self.key_sign(),
|
||||
ip_variable: self.ip.is_none(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,632 @@
|
||||
use crate::context::config::Config;
|
||||
use crate::context::nat::{MyNatInfo, PunchBackoff};
|
||||
use crate::nat::SubnetExternalRoute;
|
||||
use crate::protocol::client_message::PunchInfo;
|
||||
use crate::protocol::control_message::{ClientSimpleInfo, ClientSimpleInfoList};
|
||||
use crate::tunnel_core::p2p::route_table::RouteTable;
|
||||
use crate::tunnel_core::server::transport::config::ProtocolAddress;
|
||||
use ipnet::Ipv4Net;
|
||||
use parking_lot::{Mutex, RwLock};
|
||||
use rust_p2p_core::nat::NatInfo;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
#[derive(Default)]
|
||||
struct PingStats {
|
||||
sent: u64,
|
||||
received: u64,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct TrafficCounter {
|
||||
tx_bytes: u64,
|
||||
rx_bytes: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct TrafficStats {
|
||||
inner: Arc<RwLock<HashMap<Ipv4Addr, Arc<Mutex<TrafficCounter>>>>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub struct TrafficInfo {
|
||||
pub ip: Ipv4Addr,
|
||||
pub tx_bytes: u64,
|
||||
pub rx_bytes: u64,
|
||||
}
|
||||
|
||||
impl TrafficStats {
|
||||
fn get_or_create(&self, ip: Ipv4Addr) -> Arc<Mutex<TrafficCounter>> {
|
||||
{
|
||||
let read = self.inner.read();
|
||||
if let Some(counter) = read.get(&ip) {
|
||||
return counter.clone();
|
||||
}
|
||||
}
|
||||
let mut write = self.inner.write();
|
||||
write
|
||||
.entry(ip)
|
||||
.or_insert_with(|| Arc::new(Mutex::new(TrafficCounter::default())))
|
||||
.clone()
|
||||
}
|
||||
|
||||
pub fn record_tx(&self, ip: Ipv4Addr, bytes: u64) {
|
||||
let counter = self.get_or_create(ip);
|
||||
counter.lock().tx_bytes += bytes;
|
||||
}
|
||||
|
||||
pub fn record_rx(&self, ip: Ipv4Addr, bytes: u64) {
|
||||
let counter = self.get_or_create(ip);
|
||||
counter.lock().rx_bytes += bytes;
|
||||
}
|
||||
|
||||
pub fn get_traffic_info(&self, ip: &Ipv4Addr) -> Option<TrafficInfo> {
|
||||
let read = self.inner.read();
|
||||
read.get(ip).map(|counter| {
|
||||
let guard = counter.lock();
|
||||
TrafficInfo {
|
||||
ip: *ip,
|
||||
tx_bytes: guard.tx_bytes,
|
||||
rx_bytes: guard.rx_bytes,
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
pub fn get_all_traffic_info(&self) -> Vec<TrafficInfo> {
|
||||
let read = self.inner.read();
|
||||
read.iter()
|
||||
.map(|(ip, counter)| {
|
||||
let guard = counter.lock();
|
||||
TrafficInfo {
|
||||
ip: *ip,
|
||||
tx_bytes: guard.tx_bytes,
|
||||
rx_bytes: guard.rx_bytes,
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn reset(&self, ip: &Ipv4Addr) {
|
||||
let read = self.inner.read();
|
||||
if let Some(counter) = read.get(ip) {
|
||||
*counter.lock() = TrafficCounter::default();
|
||||
}
|
||||
}
|
||||
|
||||
pub fn reset_all(&self) {
|
||||
let read = self.inner.read();
|
||||
for counter in read.values() {
|
||||
*counter.lock() = TrafficCounter::default();
|
||||
}
|
||||
}
|
||||
|
||||
pub fn clear(&self) {
|
||||
self.inner.write().clear();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct PacketLossStats {
|
||||
inner: Arc<RwLock<HashMap<Ipv4Addr, Arc<Mutex<PingStats>>>>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub struct PacketLossInfo {
|
||||
pub ip: Ipv4Addr,
|
||||
pub sent: u64,
|
||||
pub received: u64,
|
||||
pub loss_rate: f64,
|
||||
}
|
||||
|
||||
impl PacketLossStats {
|
||||
fn get_or_create(&self, ip: Ipv4Addr) -> Arc<Mutex<PingStats>> {
|
||||
{
|
||||
let read = self.inner.read();
|
||||
if let Some(stats) = read.get(&ip) {
|
||||
return stats.clone();
|
||||
}
|
||||
}
|
||||
let mut write = self.inner.write();
|
||||
write
|
||||
.entry(ip)
|
||||
.or_insert_with(|| Arc::new(Mutex::new(PingStats::default())))
|
||||
.clone()
|
||||
}
|
||||
|
||||
pub fn record_sent(&self, ip: Ipv4Addr) {
|
||||
let stats = self.get_or_create(ip);
|
||||
stats.lock().sent += 1;
|
||||
}
|
||||
|
||||
pub fn record_received(&self, ip: Ipv4Addr) {
|
||||
let stats = self.get_or_create(ip);
|
||||
stats.lock().received += 1;
|
||||
}
|
||||
|
||||
pub fn get_loss_info(&self, ip: &Ipv4Addr) -> Option<PacketLossInfo> {
|
||||
let read = self.inner.read();
|
||||
read.get(ip).map(|stats| {
|
||||
let guard = stats.lock();
|
||||
let loss_rate = if guard.sent > 0 {
|
||||
1.0 - (guard.received as f64 / guard.sent as f64)
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
PacketLossInfo {
|
||||
ip: *ip,
|
||||
sent: guard.sent,
|
||||
received: guard.received,
|
||||
loss_rate,
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
pub fn get_all_loss_info(&self) -> Vec<PacketLossInfo> {
|
||||
let read = self.inner.read();
|
||||
read.iter()
|
||||
.map(|(ip, stats)| {
|
||||
let guard = stats.lock();
|
||||
let loss_rate = if guard.sent > 0 {
|
||||
1.0 - (guard.received as f64 / guard.sent as f64)
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
PacketLossInfo {
|
||||
ip: *ip,
|
||||
sent: guard.sent,
|
||||
received: guard.received,
|
||||
loss_rate,
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn reset(&self, ip: &Ipv4Addr) {
|
||||
let read = self.inner.read();
|
||||
if let Some(stats) = read.get(ip) {
|
||||
*stats.lock() = PingStats::default();
|
||||
}
|
||||
}
|
||||
|
||||
pub fn reset_all(&self) {
|
||||
let read = self.inner.read();
|
||||
for stats in read.values() {
|
||||
*stats.lock() = PingStats::default();
|
||||
}
|
||||
}
|
||||
|
||||
pub fn clear(&self) {
|
||||
self.inner.write().clear();
|
||||
}
|
||||
}
|
||||
|
||||
pub mod config;
|
||||
pub(crate) mod nat;
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub(crate) struct AppState {
|
||||
config: Arc<Mutex<Option<Box<Config>>>>,
|
||||
pub(crate) network: SharedNetworkAddr,
|
||||
pub(crate) server_info_collection: ServerInfoCollection,
|
||||
pub(crate) peer_map: PeerInfoMap,
|
||||
pub(crate) route_table: RouteTable,
|
||||
pub(crate) subnet_route: SubnetExternalRoute,
|
||||
pub(crate) nat_info: MyNatInfo,
|
||||
pub(crate) punch_backoff: PunchBackoff,
|
||||
pub(crate) packet_loss_stats: PacketLossStats,
|
||||
pub(crate) traffic_stats: TrafficStats,
|
||||
}
|
||||
#[derive(Clone, Default)]
|
||||
pub(crate) struct SharedNetworkAddr {
|
||||
inner: Arc<Mutex<Option<NetworkAddr>>>,
|
||||
}
|
||||
impl SharedNetworkAddr {
|
||||
pub fn network(&self) -> Option<Ipv4Net> {
|
||||
self.inner.lock().as_ref().map(|v| v.network())
|
||||
}
|
||||
pub fn ip(&self) -> Option<Ipv4Addr> {
|
||||
self.inner.lock().map(|v| v.ip)
|
||||
}
|
||||
pub fn get(&self) -> Option<NetworkAddr> {
|
||||
*self.inner.lock()
|
||||
}
|
||||
pub fn set(&self, addr: NetworkAddr) {
|
||||
*self.inner.lock() = Some(addr);
|
||||
}
|
||||
pub fn clear(&self) {
|
||||
*self.inner.lock() = None;
|
||||
}
|
||||
}
|
||||
|
||||
/// 网络路由封装,包含本地网络信息和子网路由
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct NetworkRoute {
|
||||
pub network: SharedNetworkAddr,
|
||||
pub subnet_route: SubnetExternalRoute,
|
||||
}
|
||||
|
||||
impl NetworkRoute {
|
||||
pub fn new(network: SharedNetworkAddr, subnet_route: SubnetExternalRoute) -> Self {
|
||||
Self {
|
||||
network,
|
||||
subnet_route,
|
||||
}
|
||||
}
|
||||
|
||||
/// 检查 IP 是否在本地网络或子网路由中
|
||||
pub fn network_contains(&self, ip: &Ipv4Addr) -> bool {
|
||||
if let Some(network) = self.network.network()
|
||||
&& network.contains(ip)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
self.subnet_route.route(ip).is_some()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub struct PeerClientInfo {
|
||||
pub nat_info: Option<NatInfo>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub(crate) struct PeerInfoMap {
|
||||
inner: Arc<Mutex<HashMap<Ipv4Addr, PeerClientInfo>>>,
|
||||
}
|
||||
|
||||
impl PeerInfoMap {
|
||||
pub fn get(&self, ip: &Ipv4Addr) -> Option<PeerClientInfo> {
|
||||
self.inner.lock().get(ip).cloned()
|
||||
}
|
||||
|
||||
pub fn update_nat_info(&self, ip: Ipv4Addr, nat_info: NatInfo) {
|
||||
let mut guard = self.inner.lock();
|
||||
if let Some(v) = guard.get_mut(&ip) {
|
||||
v.nat_info = Some(nat_info);
|
||||
return;
|
||||
}
|
||||
guard.insert(
|
||||
ip,
|
||||
PeerClientInfo {
|
||||
nat_info: Some(nat_info),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
pub fn clear(&self) {
|
||||
self.inner.lock().clear();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub(crate) struct ServerInfoCollection {
|
||||
client_simple_list: Arc<RwLock<Vec<ClientSimpleInfo>>>,
|
||||
server_node_map: Arc<RwLock<HashMap<u32, ServerNodeInfo>>>,
|
||||
}
|
||||
#[derive(Clone, Default)]
|
||||
pub struct ServerNodeInfo {
|
||||
pub server_id: u32,
|
||||
pub server_addr: ProtocolAddress,
|
||||
pub connected: bool,
|
||||
pub rtt: Option<u32>,
|
||||
pub data_version: u64,
|
||||
pub client_map: HashMap<Ipv4Addr, ClientSimpleInfo>,
|
||||
pub last_connected_time: Option<i64>,
|
||||
pub disconnected_time: Option<i64>,
|
||||
pub server_version: Option<String>,
|
||||
}
|
||||
impl ServerInfoCollection {
|
||||
pub fn server_client_ip_map(&self) -> HashMap<u32, (Vec<Ipv4Addr>, u32)> {
|
||||
self.server_node_map
|
||||
.read()
|
||||
.iter()
|
||||
.map(|(k, v)| {
|
||||
(
|
||||
*k,
|
||||
(
|
||||
v.client_map
|
||||
.iter()
|
||||
.filter(|(_, v)| v.online)
|
||||
.map(|(k, _)| *k)
|
||||
.collect(),
|
||||
v.rtt.unwrap_or(500),
|
||||
),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
pub fn server_node_list(&self) -> Vec<ServerNodeInfo> {
|
||||
self.server_node_map.read().values().cloned().collect()
|
||||
}
|
||||
pub fn update_server(&self, addr: Vec<(u32, ProtocolAddress)>) {
|
||||
let mut server_node_map_guard = self.server_node_map.write();
|
||||
let mut client_simple_list_guard = self.client_simple_list.write();
|
||||
server_node_map_guard.clear();
|
||||
client_simple_list_guard.clear();
|
||||
for (server_id, server_addr) in addr {
|
||||
let server_node = ServerNodeInfo {
|
||||
server_id,
|
||||
server_addr,
|
||||
..Default::default()
|
||||
};
|
||||
server_node_map_guard.insert(server_id, server_node);
|
||||
}
|
||||
}
|
||||
pub fn find_connected_server(&self, server_ids: &[u32]) -> Option<u32> {
|
||||
let map = self.server_node_map.read();
|
||||
|
||||
server_ids
|
||||
.iter()
|
||||
.filter_map(|id| {
|
||||
let server = map.get(id)?;
|
||||
|
||||
if !server.connected {
|
||||
return None;
|
||||
}
|
||||
|
||||
let rtt = server.rtt.unwrap_or(u32::MAX);
|
||||
|
||||
Some((*id, rtt))
|
||||
})
|
||||
.min_by_key(|(_, rtt)| *rtt)
|
||||
.map(|(id, _)| id)
|
||||
}
|
||||
pub fn find_ip_to_server(&self, server_ids: &[u32], ip: &Ipv4Addr) -> Option<u32> {
|
||||
let map = self.server_node_map.read();
|
||||
|
||||
server_ids
|
||||
.iter()
|
||||
.filter_map(|id| {
|
||||
let server = map.get(id)?;
|
||||
|
||||
if !server.connected {
|
||||
return None;
|
||||
}
|
||||
|
||||
let client = server.client_map.get(ip)?;
|
||||
if !client.online {
|
||||
return None;
|
||||
}
|
||||
|
||||
let rtt = server.rtt.unwrap_or(u32::MAX);
|
||||
|
||||
Some((*id, rtt))
|
||||
})
|
||||
.min_by_key(|(_, rtt)| *rtt)
|
||||
.map(|(id, _)| id)
|
||||
}
|
||||
pub fn client_online_ips(&self) -> Vec<Ipv4Addr> {
|
||||
self.client_simple_list
|
||||
.read()
|
||||
.iter()
|
||||
.filter(|v| v.online)
|
||||
.map(|c| c.ip)
|
||||
.collect()
|
||||
}
|
||||
pub fn client_ips(&self) -> Vec<ClientSimpleInfo> {
|
||||
self.client_simple_list.read().clone()
|
||||
}
|
||||
pub fn data_version(&self, server_id: u32) -> u64 {
|
||||
self.server_node_map
|
||||
.read()
|
||||
.get(&server_id)
|
||||
.map(|v| v.data_version)
|
||||
.unwrap_or(0)
|
||||
}
|
||||
pub fn update_client_simple_list(
|
||||
&self,
|
||||
server_id: u32,
|
||||
self_ip: Ipv4Addr,
|
||||
client_simple_list: ClientSimpleInfoList,
|
||||
now: i64,
|
||||
) {
|
||||
let mut guard = self.server_node_map.write();
|
||||
let server_node = guard.entry(server_id).or_default();
|
||||
if now > client_simple_list.time {
|
||||
server_node.rtt = Some((now - client_simple_list.time) as u32);
|
||||
}
|
||||
server_node.data_version = client_simple_list.data_version;
|
||||
let map: HashMap<Ipv4Addr, ClientSimpleInfo> = client_simple_list
|
||||
.list
|
||||
.into_iter()
|
||||
.filter(|v| v.ip != self_ip)
|
||||
.map(|info| (info.ip, info))
|
||||
.collect();
|
||||
|
||||
if client_simple_list.is_all {
|
||||
server_node.client_map = map;
|
||||
} else {
|
||||
server_node.client_map.extend(map);
|
||||
}
|
||||
let mut client_simple_map = HashMap::<Ipv4Addr, ClientSimpleInfo>::new();
|
||||
for (_, server_node) in guard.iter() {
|
||||
for (_, x) in server_node.client_map.iter() {
|
||||
if let Some(v) = client_simple_map.get_mut(&x.ip) {
|
||||
if x.online {
|
||||
v.online = true;
|
||||
}
|
||||
} else {
|
||||
client_simple_map.insert(x.ip, x.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut guard = self.client_simple_list.write();
|
||||
*guard = client_simple_map.into_values().collect()
|
||||
}
|
||||
pub fn set_server_connected(&self, server_id: u32, val: bool) -> bool {
|
||||
let mut mutex_guard = self.server_node_map.write();
|
||||
let server_node = mutex_guard.entry(server_id).or_default();
|
||||
let old = server_node.connected;
|
||||
server_node.connected = val;
|
||||
old
|
||||
}
|
||||
pub fn is_any_server_connected(&self, server_ids: Option<&[u32]>) -> bool {
|
||||
let guard = self.server_node_map.read();
|
||||
if let Some(server_ids) = server_ids {
|
||||
for id in server_ids {
|
||||
if guard.get(id).map(|v| v.connected).unwrap_or(false) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for (_, server_node) in guard.iter() {
|
||||
if server_node.connected {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
pub fn is_server_connected(&self, server_id: u32) -> bool {
|
||||
self.server_node_map
|
||||
.read()
|
||||
.get(&server_id)
|
||||
.map(|v| v.connected)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
pub fn set_last_connected_time(&self, server_id: u32, last_connected_time: Option<i64>) {
|
||||
self.server_node_map
|
||||
.write()
|
||||
.entry(server_id)
|
||||
.or_default()
|
||||
.last_connected_time = last_connected_time;
|
||||
}
|
||||
pub fn set_disconnected_time(&self, server_id: u32, last_connected_time: Option<i64>) {
|
||||
self.server_node_map
|
||||
.write()
|
||||
.entry(server_id)
|
||||
.or_default()
|
||||
.disconnected_time = last_connected_time;
|
||||
}
|
||||
pub fn set_server_rtt(&self, server_id: u32, rtt: u32) {
|
||||
if let Some(v) = self.server_node_map.write().get_mut(&server_id) {
|
||||
v.rtt = Some(rtt);
|
||||
}
|
||||
}
|
||||
pub fn set_server_version(&self, server_id: u32, version: String) {
|
||||
if let Some(v) = self.server_node_map.write().get_mut(&server_id) {
|
||||
v.server_version = Some(version);
|
||||
}
|
||||
}
|
||||
pub fn get_server_rtt(&self, ip: &Ipv4Addr) -> Option<u32> {
|
||||
let server_node_map_guard = self.server_node_map.read();
|
||||
for (_, server_node) in server_node_map_guard.iter() {
|
||||
if !server_node.connected {
|
||||
continue;
|
||||
}
|
||||
if let Some(v) = server_node.client_map.get(ip)
|
||||
&& v.online
|
||||
{
|
||||
return server_node.rtt;
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
pub fn exists_online_client_ip(&self, ip: &Ipv4Addr) -> bool {
|
||||
self.client_simple_list
|
||||
.read()
|
||||
.iter()
|
||||
.any(|v| v.ip == *ip && v.online)
|
||||
}
|
||||
pub fn clear(&self) {
|
||||
self.client_simple_list.write().clear();
|
||||
let mut guard = self.server_node_map.write();
|
||||
for server_node in guard.values_mut() {
|
||||
server_node.connected = false;
|
||||
server_node.rtt = None;
|
||||
server_node.data_version = 0;
|
||||
server_node.client_map.clear();
|
||||
server_node.last_connected_time = None;
|
||||
server_node.disconnected_time = None;
|
||||
}
|
||||
}
|
||||
}
|
||||
#[derive(Copy, Clone, Debug)]
|
||||
pub struct NetworkAddr {
|
||||
pub gateway: Ipv4Addr,
|
||||
pub broadcast: Ipv4Addr,
|
||||
pub ip: Ipv4Addr,
|
||||
pub prefix_len: u8,
|
||||
}
|
||||
impl NetworkAddr {
|
||||
pub fn network(&self) -> Ipv4Net {
|
||||
Ipv4Net::new_assert(self.ip, self.prefix_len)
|
||||
}
|
||||
}
|
||||
|
||||
impl AppState {
|
||||
pub fn stop_network(&self) {
|
||||
self.network.clear();
|
||||
self.server_info_collection.clear();
|
||||
self.peer_map.clear();
|
||||
// route_table 来自外部 crate,会在任务停止后自动失效
|
||||
self.nat_info.clear();
|
||||
self.punch_backoff.clear();
|
||||
self.packet_loss_stats.clear();
|
||||
self.traffic_stats.clear();
|
||||
}
|
||||
|
||||
fn network(&self) -> Option<Ipv4Net> {
|
||||
self.network.network()
|
||||
}
|
||||
pub fn get_network(&self) -> Option<NetworkAddr> {
|
||||
self.network.get()
|
||||
}
|
||||
fn network_contains(&self, ip: &Ipv4Addr) -> bool {
|
||||
let Some(network) = self.network() else {
|
||||
return false;
|
||||
};
|
||||
if network.contains(ip) {
|
||||
return true;
|
||||
}
|
||||
self.subnet_route.route(ip).is_some()
|
||||
}
|
||||
pub fn client_ips(&self) -> Vec<ClientSimpleInfo> {
|
||||
self.server_info_collection.client_ips()
|
||||
}
|
||||
pub fn get_peer_info(&self, ip: &Ipv4Addr) -> Option<PeerClientInfo> {
|
||||
self.peer_map.get(ip)
|
||||
}
|
||||
|
||||
pub fn set_config(&self, config: Box<Config>) {
|
||||
*self.config.lock() = Some(config);
|
||||
}
|
||||
|
||||
pub fn get_config(&self) -> Option<Box<Config>> {
|
||||
self.config.lock().clone()
|
||||
}
|
||||
pub(crate) fn udp_stun(&self) -> Vec<String> {
|
||||
self.config
|
||||
.lock()
|
||||
.as_ref()
|
||||
.map(|v| v.udp_stun.clone())
|
||||
.unwrap_or_default()
|
||||
}
|
||||
pub(crate) fn tcp_stun(&self) -> Vec<String> {
|
||||
self.config
|
||||
.lock()
|
||||
.as_ref()
|
||||
.map(|v| v.tcp_stun.clone())
|
||||
.unwrap_or_default()
|
||||
}
|
||||
}
|
||||
impl AppState {
|
||||
pub fn get_punch_info(&self) -> Option<PunchInfo> {
|
||||
self.nat_info.get().map(|info| PunchInfo {
|
||||
nat_info: self.filter_ip(info),
|
||||
})
|
||||
}
|
||||
pub fn get_nat_info(&self) -> Option<NatInfo> {
|
||||
self.nat_info.get().map(|info| self.filter_ip(info))
|
||||
}
|
||||
pub fn filter_ip(&self, mut info: NatInfo) -> NatInfo {
|
||||
if self.network_contains(&info.local_ipv4) {
|
||||
info.local_ipv4 = Ipv4Addr::UNSPECIFIED;
|
||||
}
|
||||
info.local_ipv4s.retain(|ip| !self.network_contains(ip));
|
||||
info
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
use parking_lot::RwLock;
|
||||
use rust_p2p_core::nat::NatInfo;
|
||||
use rust_p2p_core::route::Index;
|
||||
use rust_p2p_core::tunnel::udp::UDPIndex;
|
||||
use std::collections::HashMap;
|
||||
use std::net::{Ipv4Addr, SocketAddr};
|
||||
use std::sync::Arc;
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct MyNatInfo {
|
||||
nat_info: Arc<RwLock<Option<NatInfo>>>,
|
||||
}
|
||||
impl MyNatInfo {
|
||||
pub fn get(&self) -> Option<NatInfo> {
|
||||
self.nat_info.read().clone()
|
||||
}
|
||||
pub fn update_public_addr(&self, index: Index, addr: SocketAddr) {
|
||||
let (ip, port) = if let Some(r) = mapping_addr(addr) {
|
||||
r
|
||||
} else {
|
||||
return;
|
||||
};
|
||||
log::debug!("public_addr:{},{},index={index:?}", ip, port);
|
||||
let mut nat_info = self.nat_info.write();
|
||||
let Some(nat_info) = nat_info.as_mut() else {
|
||||
return;
|
||||
};
|
||||
if rust_p2p_core::extend::addr::is_ipv4_global(&ip) {
|
||||
if !nat_info.public_ips.contains(&ip) {
|
||||
nat_info.public_ips.push(ip);
|
||||
}
|
||||
match index {
|
||||
Index::Udp(index) => {
|
||||
let index = match index {
|
||||
UDPIndex::MainV4(index) => index,
|
||||
UDPIndex::MainV6(index) => index,
|
||||
UDPIndex::SubV4(_) => return,
|
||||
};
|
||||
if let Some(p) = nat_info.public_udp_ports.get_mut(index) {
|
||||
*p = port;
|
||||
}
|
||||
}
|
||||
Index::Tcp(_) => {
|
||||
nat_info.public_tcp_port = port;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
} else {
|
||||
log::debug!("not public addr: {addr:?}")
|
||||
}
|
||||
}
|
||||
pub fn update_tcp_public_addr(&self, addr: SocketAddr) {
|
||||
let SocketAddr::V4(addr) = addr else {
|
||||
return;
|
||||
};
|
||||
let ip = *addr.ip();
|
||||
let port = addr.port();
|
||||
log::info!("tcp_public_addr, {}:{}", ip, port);
|
||||
let mut nat_info = self.nat_info.write();
|
||||
let Some(nat_info) = nat_info.as_mut() else {
|
||||
return;
|
||||
};
|
||||
if ip.is_unspecified() && port == 0 {
|
||||
nat_info.public_tcp_port = 0;
|
||||
return;
|
||||
}
|
||||
if rust_p2p_core::extend::addr::is_ipv4_global(&ip) {
|
||||
if !nat_info.public_ips.contains(&ip) {
|
||||
nat_info.public_ips.push(ip);
|
||||
}
|
||||
nat_info.public_tcp_port = port;
|
||||
} else {
|
||||
log::debug!("not public addr: {addr:?}")
|
||||
}
|
||||
}
|
||||
pub fn replace_nat_info(&self, nat_info: NatInfo) {
|
||||
self.nat_info.write().replace(nat_info);
|
||||
}
|
||||
pub fn clear(&self) {
|
||||
*self.nat_info.write() = None;
|
||||
}
|
||||
}
|
||||
fn mapping_addr(addr: SocketAddr) -> Option<(Ipv4Addr, u16)> {
|
||||
match addr {
|
||||
SocketAddr::V4(addr) => Some((*addr.ip(), addr.port())),
|
||||
SocketAddr::V6(addr) => addr.ip().to_ipv4_mapped().map(|ip| (ip, addr.port())),
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Copy, Clone, Debug)]
|
||||
pub struct PunchState {
|
||||
pub count: i64,
|
||||
pub last_ts: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct PunchBackoff {
|
||||
inner: Arc<RwLock<HashMap<Ipv4Addr, PunchState>>>,
|
||||
}
|
||||
|
||||
impl PunchBackoff {
|
||||
const MAX_BACKOFF_MS: i64 = 3_600_000; // 1h
|
||||
const BASE_MS: i64 = 3000;
|
||||
|
||||
fn now() -> i64 {
|
||||
crate::utils::time::now_ts_ms()
|
||||
}
|
||||
|
||||
pub fn record(&self, ip: Ipv4Addr) {
|
||||
let mut map = self.inner.write();
|
||||
let entry = map.entry(ip).or_insert(PunchState {
|
||||
count: 0,
|
||||
last_ts: 0,
|
||||
});
|
||||
entry.count += 1;
|
||||
entry.last_ts = Self::now();
|
||||
}
|
||||
#[allow(dead_code)]
|
||||
pub fn reset(&self, ip: Ipv4Addr) {
|
||||
self.inner.write().remove(&ip);
|
||||
}
|
||||
#[allow(dead_code)]
|
||||
pub fn get(&self, ip: &Ipv4Addr) -> Option<PunchState> {
|
||||
self.inner.read().get(ip).copied()
|
||||
}
|
||||
|
||||
pub fn should_punch(&self, ip: Ipv4Addr) -> bool {
|
||||
let map = self.inner.read();
|
||||
let Some(state) = map.get(&ip) else {
|
||||
return true;
|
||||
};
|
||||
|
||||
let now = Self::now();
|
||||
let elapsed = now - state.last_ts;
|
||||
|
||||
let mut backoff = Self::BASE_MS * state.count;
|
||||
if backoff > Self::MAX_BACKOFF_MS {
|
||||
backoff = Self::MAX_BACKOFF_MS;
|
||||
}
|
||||
|
||||
elapsed >= backoff
|
||||
}
|
||||
#[allow(dead_code)]
|
||||
pub fn clear(&self) {
|
||||
self.inner.write().clear();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,353 @@
|
||||
use crate::api::VntApi;
|
||||
use crate::compression::PacketCompression;
|
||||
use crate::context::config::Config;
|
||||
use crate::context::{AppState, NetworkAddr, NetworkRoute};
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::enhanced_tunnel::enhanced_ipv4_tunnel;
|
||||
use crate::enhanced_tunnel::inbound::EnhancedInbound;
|
||||
use crate::enhanced_tunnel::outbound::EnhancedOutbound;
|
||||
use crate::fec::{FecDecoder, FecEncoder};
|
||||
use crate::nat::internal_nat::{InternalNatInbound, PortMappingManager};
|
||||
use crate::nat::{AllowSubnetExternalRoute, SubnetExternalRoute};
|
||||
use crate::tun::enhanced_tun::EnhancedTunInbound;
|
||||
use crate::tun::{DeviceConfig, DeviceIOManager, TunDataInbound, TunReceiver, tun_channel};
|
||||
use crate::tunnel_core::outbound::{BasicOutbound, HybridOutbound};
|
||||
use crate::tunnel_core::p2p::inbound::{P2pInboundConfig, P2pInboundHandler};
|
||||
use crate::tunnel_core::p2p::transport::punch::NatPuncher;
|
||||
use crate::tunnel_core::p2p::transport::task::init_tunnel;
|
||||
use crate::tunnel_core::server::connection_manager::{
|
||||
InboundHandlerConfig, ServerTurnManager, coordinated_registration, create_server_tunnel,
|
||||
};
|
||||
use crate::tunnel_core::server::rpc::ServerRPC;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::bail;
|
||||
use ipnet::Ipv4Net;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
pub const DEFAULT_MTU: u16 = 1380;
|
||||
|
||||
/// Context for deferred registration
|
||||
struct RegistrationContext {
|
||||
server_managers: Vec<ServerTurnManager>,
|
||||
subnet_external_route: SubnetExternalRoute,
|
||||
puncher: NatPuncher,
|
||||
packet_crypto: PacketCrypto,
|
||||
packet_compression: PacketCompression,
|
||||
enhanced_inbound: EnhancedInbound,
|
||||
fec_decoder: FecDecoder,
|
||||
}
|
||||
|
||||
pub struct NetworkManager {
|
||||
config: Box<Config>,
|
||||
app_state: AppState,
|
||||
task_group: TaskGroup,
|
||||
device_io_manager: DeviceIOManager,
|
||||
enhanced_outbound: Option<EnhancedOutbound>,
|
||||
server_rpc: ServerRPC,
|
||||
tun_receiver: Option<TunReceiver>,
|
||||
registration_context: Option<Box<RegistrationContext>>,
|
||||
}
|
||||
|
||||
impl NetworkManager {
|
||||
pub async fn create_network(
|
||||
config: Box<Config>,
|
||||
task_group: TaskGroup,
|
||||
) -> anyhow::Result<NetworkManager> {
|
||||
let app_state = AppState::default();
|
||||
config.check()?;
|
||||
let mtu = config.mtu.unwrap_or(DEFAULT_MTU);
|
||||
let packet_crypto = PacketCrypto::new_from_str(config.password.as_deref());
|
||||
let packet_compression = PacketCompression::new(config.compress);
|
||||
let (server_manager_list, tunnel_to_server, server_rpc) =
|
||||
create_server_tunnel(app_state.clone(), &config, packet_crypto.clone());
|
||||
let device_io_manager = DeviceIOManager::new(task_group.clone());
|
||||
let allow_subnet = AllowSubnetExternalRoute::new(config.output.clone());
|
||||
|
||||
let (puncher, p2p_socket, p2p_task) = if !config.no_punch {
|
||||
let (puncher, p2p_socket_manager, p2p_task) = init_tunnel(
|
||||
task_group.clone(),
|
||||
app_state.clone(),
|
||||
tunnel_to_server.clone(),
|
||||
packet_crypto.clone(),
|
||||
)
|
||||
.await?;
|
||||
|
||||
(Some(puncher), Some(p2p_socket_manager), Some(p2p_task))
|
||||
} else {
|
||||
(None, None, None)
|
||||
};
|
||||
let puncher = NatPuncher::new(
|
||||
app_state.network.clone(),
|
||||
app_state.punch_backoff.clone(),
|
||||
puncher,
|
||||
packet_crypto.clone(),
|
||||
);
|
||||
let subnet_external_route = app_state.subnet_route.clone();
|
||||
subnet_external_route.set_route_table(config.input.clone());
|
||||
|
||||
let fec_decoder = FecDecoder::new();
|
||||
let basic_outbound = BasicOutbound::new(
|
||||
tunnel_to_server.clone(),
|
||||
p2p_socket.clone(),
|
||||
packet_crypto.clone(),
|
||||
);
|
||||
let fec_encoder = if config.fec {
|
||||
Some(FecEncoder::new(&task_group, basic_outbound.clone()))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let hybrid_outbound = HybridOutbound::new(
|
||||
app_state.network.clone(),
|
||||
app_state.server_info_collection.clone(),
|
||||
app_state.traffic_stats.clone(),
|
||||
basic_outbound,
|
||||
packet_compression.clone(),
|
||||
subnet_external_route.clone(),
|
||||
fec_encoder,
|
||||
);
|
||||
let port_mapping_manager = PortMappingManager::new(
|
||||
config.no_tun,
|
||||
config.allow_port_mapping,
|
||||
app_state.network.clone(),
|
||||
);
|
||||
let internal_nat_inbound = if config.no_nat && !config.no_tun {
|
||||
None
|
||||
} else {
|
||||
let nat_inbound = InternalNatInbound::create(
|
||||
&task_group,
|
||||
mtu,
|
||||
hybrid_outbound.clone(),
|
||||
allow_subnet.clone(),
|
||||
app_state.network.clone(),
|
||||
config.no_tun,
|
||||
)
|
||||
.await?;
|
||||
Some(nat_inbound)
|
||||
};
|
||||
|
||||
let (enhanced_tun_inbound, tun_receiver) = if config.no_tun {
|
||||
(
|
||||
EnhancedTunInbound::Nat(
|
||||
internal_nat_inbound
|
||||
.clone()
|
||||
.expect("internal_nat_inbound must be Some when no_tun is true"),
|
||||
),
|
||||
None,
|
||||
)
|
||||
} else {
|
||||
let (tun_inbound, tun_receiver) = tun_channel();
|
||||
let tun_data_sender = TunDataInbound::new(tun_inbound, allow_subnet.clone());
|
||||
(EnhancedTunInbound::Tun(tun_data_sender), Some(tun_receiver))
|
||||
};
|
||||
|
||||
let (enhanced_inbound, enhanced_outbound) = enhanced_ipv4_tunnel(
|
||||
app_state.clone(),
|
||||
task_group.clone(),
|
||||
enhanced_tun_inbound,
|
||||
crate::enhanced_tunnel::TunnelConfig {
|
||||
mtu,
|
||||
password: config.password.clone(),
|
||||
open_quic_client: config.rtx,
|
||||
port_mapping: config.port_mapping.clone(),
|
||||
},
|
||||
crate::enhanced_tunnel::TunnelComponents {
|
||||
hybrid_outbound: hybrid_outbound.clone(),
|
||||
external_route: subnet_external_route.clone(),
|
||||
internal_nat_inbound,
|
||||
port_mapping_manager,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
|
||||
if let Some(p2p_task) = p2p_task {
|
||||
let handler = P2pInboundHandler::new(P2pInboundConfig {
|
||||
network_route: NetworkRoute::new(
|
||||
app_state.network.clone(),
|
||||
subnet_external_route.clone(),
|
||||
),
|
||||
route_table: app_state.route_table.clone(),
|
||||
packet_loss_stats: app_state.packet_loss_stats.clone(),
|
||||
packet_crypto: packet_crypto.clone(),
|
||||
packet_compression: packet_compression.clone(),
|
||||
enhanced_inbound: enhanced_inbound.clone(),
|
||||
fec_decoder: fec_decoder.clone(),
|
||||
});
|
||||
p2p_task.start(handler);
|
||||
}
|
||||
|
||||
let registration_context = Box::new(RegistrationContext {
|
||||
server_managers: server_manager_list,
|
||||
subnet_external_route,
|
||||
puncher,
|
||||
packet_crypto,
|
||||
packet_compression,
|
||||
enhanced_inbound,
|
||||
fec_decoder,
|
||||
});
|
||||
|
||||
app_state.set_config(config.clone());
|
||||
Ok(Self {
|
||||
config,
|
||||
app_state,
|
||||
task_group,
|
||||
device_io_manager,
|
||||
enhanced_outbound,
|
||||
server_rpc,
|
||||
tun_receiver,
|
||||
registration_context: Some(registration_context),
|
||||
})
|
||||
}
|
||||
|
||||
/// Register with server(s) and start data handling tasks.
|
||||
/// This method can only be called once.
|
||||
/// Returns the registration response on success.
|
||||
pub async fn register(&mut self) -> anyhow::Result<NetworkAddr> {
|
||||
let Some(mut ctx) = self.registration_context.take() else {
|
||||
bail!("register can only be called once");
|
||||
};
|
||||
|
||||
let is_multi_server = ctx.server_managers.len() > 1;
|
||||
|
||||
let reg_response = if is_multi_server {
|
||||
// Multi-server: coordinated pre-registration
|
||||
log::info!(
|
||||
"Multi-server mode: performing coordinated registration for {} servers",
|
||||
ctx.server_managers.len()
|
||||
);
|
||||
let reg_response = coordinated_registration(&mut ctx.server_managers).await?;
|
||||
log::info!(
|
||||
"Coordinated registration completed, IP: {}, prefix_len: {}",
|
||||
reg_response.ip,
|
||||
reg_response.prefix_len
|
||||
);
|
||||
reg_response
|
||||
} else {
|
||||
// Single-server: normal registration
|
||||
log::info!("Single-server mode: performing normal registration");
|
||||
let response = ctx.server_managers[0]
|
||||
.connect_and_reg(crate::protocol::control_message::RegistrationMode::Normal)
|
||||
.await?;
|
||||
match response {
|
||||
crate::protocol::control_message::ResponseMessage::Reg(reg) => {
|
||||
log::info!(
|
||||
"Registration completed, IP: {}, prefix_len: {}",
|
||||
reg.ip,
|
||||
reg.prefix_len
|
||||
);
|
||||
reg
|
||||
}
|
||||
crate::protocol::control_message::ResponseMessage::Error(e) => {
|
||||
bail!("Registration failed: {}", e.message);
|
||||
}
|
||||
crate::protocol::control_message::ResponseMessage::ConfirmReg(_) => {
|
||||
bail!("Unexpected ConfirmReg response");
|
||||
}
|
||||
}
|
||||
};
|
||||
let network_addr = NetworkAddr {
|
||||
gateway: reg_response.gateway,
|
||||
broadcast: Ipv4Net::new(reg_response.ip, reg_response.prefix_len)?.broadcast(),
|
||||
ip: reg_response.ip,
|
||||
prefix_len: reg_response.prefix_len,
|
||||
};
|
||||
self.app_state.network.set(network_addr);
|
||||
|
||||
// 保存服务器版本信息
|
||||
if !reg_response.server_version.is_empty() {
|
||||
for (index, _) in ctx.server_managers.iter().enumerate() {
|
||||
self.app_state.server_info_collection.set_server_version(
|
||||
index as u32,
|
||||
reg_response.server_version.clone(),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Start data handling tasks for all servers
|
||||
for turn_manager in ctx.server_managers {
|
||||
let handler_config = Box::new(InboundHandlerConfig {
|
||||
network_route: NetworkRoute::new(
|
||||
self.app_state.network.clone(),
|
||||
ctx.subnet_external_route.clone(),
|
||||
),
|
||||
server_info: self.app_state.server_info_collection.clone(),
|
||||
nat_info: self.app_state.nat_info.clone(),
|
||||
peer_map: self.app_state.peer_map.clone(),
|
||||
punch_backoff: self.app_state.punch_backoff.clone(),
|
||||
puncher: ctx.puncher.clone(),
|
||||
packet_crypto: ctx.packet_crypto.clone(),
|
||||
packet_compression: ctx.packet_compression.clone(),
|
||||
enhanced_inbound: ctx.enhanced_inbound.clone(),
|
||||
fec_decoder: ctx.fec_decoder.clone(),
|
||||
});
|
||||
turn_manager.data_handle_task_connected(&self.task_group, handler_config, network_addr);
|
||||
}
|
||||
|
||||
Ok(network_addr)
|
||||
}
|
||||
|
||||
pub fn is_no_tun(&self) -> bool {
|
||||
self.config.no_tun
|
||||
}
|
||||
|
||||
pub async fn start_tun(&mut self) -> anyhow::Result<()> {
|
||||
let Some(receiver) = self.tun_receiver.take() else {
|
||||
bail!("start_tun can only be called once");
|
||||
};
|
||||
let Some(enhanced_outbound) = self.enhanced_outbound.take() else {
|
||||
bail!("start_tun can only be called once");
|
||||
};
|
||||
let mut config = DeviceConfig::default();
|
||||
config = config.set_mtu(self.config.mtu.unwrap_or(DEFAULT_MTU));
|
||||
if let Some(tun_name) = self.config.tun_name.clone() {
|
||||
config = config.set_tun_name(tun_name);
|
||||
}
|
||||
self.device_io_manager
|
||||
.start_task(config, receiver, enhanced_outbound)
|
||||
.await
|
||||
}
|
||||
#[cfg(unix)]
|
||||
pub async fn start_tun_fd(&mut self, tun_fd: Option<i32>) -> anyhow::Result<()> {
|
||||
let Some(receiver) = self.tun_receiver.take() else {
|
||||
bail!("start_tun_fd can only be called once");
|
||||
};
|
||||
let Some(enhanced_outbound) = self.enhanced_outbound.take() else {
|
||||
bail!("start_tun_fd can only be called once");
|
||||
};
|
||||
let mut config = DeviceConfig::default();
|
||||
if let Some(tun_fd) = tun_fd {
|
||||
config = config.set_tun_fd(tun_fd);
|
||||
}
|
||||
if let Some(tun_name) = self.config.tun_name.clone() {
|
||||
config = config.set_tun_name(tun_name);
|
||||
}
|
||||
self.device_io_manager
|
||||
.start_task(config, receiver, enhanced_outbound)
|
||||
.await
|
||||
}
|
||||
#[cfg(not(target_os = "android"))]
|
||||
pub async fn set_network_ip(&self, ip: Ipv4Addr, prefix_len: u8) -> anyhow::Result<()> {
|
||||
self.device_io_manager.set_network(ip, prefix_len).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn stop_network(&mut self) {
|
||||
self.task_group.stop();
|
||||
self.app_state.stop_network();
|
||||
}
|
||||
#[cfg(not(target_os = "android"))]
|
||||
pub async fn tun_if_index(&self) -> anyhow::Result<u32> {
|
||||
self.device_io_manager.tun_if_index().await
|
||||
}
|
||||
pub async fn wait_all_stopped(&mut self) {
|
||||
self.task_group.wait_all_stopped().await;
|
||||
}
|
||||
pub fn vnt_api(&self) -> VntApi {
|
||||
VntApi::new(self.app_state.clone(), self.server_rpc.clone())
|
||||
}
|
||||
}
|
||||
impl Drop for NetworkManager {
|
||||
fn drop(&mut self) {
|
||||
self.stop_network();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, NetPacket};
|
||||
use ring::aead::{Aad, CHACHA20_POLY1305, LessSafeKey, Nonce, UnboundKey};
|
||||
use std::io;
|
||||
|
||||
pub const TAG_LEN: usize = 16;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PacketCrypto {
|
||||
key: LessSafeKey,
|
||||
}
|
||||
|
||||
impl PacketCrypto {
|
||||
pub fn key_sign(s: &str) -> String {
|
||||
use ring::digest::{Context, SHA256};
|
||||
|
||||
const PREFIX: &[u8] = b"KEY-BEGIN";
|
||||
const SUFFIX: &[u8] = b"KEY-END";
|
||||
|
||||
let mut ctx = Context::new(&SHA256);
|
||||
ctx.update(PREFIX);
|
||||
ctx.update(s.as_bytes());
|
||||
ctx.update(SUFFIX);
|
||||
let digest = ctx.finish();
|
||||
let mut key_bytes = [0u8; 16];
|
||||
key_bytes.copy_from_slice(&digest.as_ref()[..16]);
|
||||
key_bytes
|
||||
.iter()
|
||||
.map(|b| format!("{:02x}", b))
|
||||
.collect::<String>()
|
||||
}
|
||||
pub fn new(key_bytes: [u8; 32]) -> Self {
|
||||
let unbound = UnboundKey::new(&CHACHA20_POLY1305, &key_bytes).unwrap();
|
||||
let key = LessSafeKey::new(unbound);
|
||||
Self { key }
|
||||
}
|
||||
pub fn new_from_str(s: &str) -> Self {
|
||||
let hash = ring::digest::digest(&ring::digest::SHA256, s.as_bytes());
|
||||
let mut key_bytes = [0u8; 32];
|
||||
key_bytes.copy_from_slice(hash.as_ref());
|
||||
Self::new(key_bytes)
|
||||
}
|
||||
/// 根据包头生成 12 字节 nonce
|
||||
pub fn make_nonce<B: AsRef<[u8]>>(&self, pkt: &NetPacket<B>) -> io::Result<[u8; 12]> {
|
||||
let buf = pkt.buffer();
|
||||
|
||||
if buf.len() < HEAD_LENGTH {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"buffer too small",
|
||||
));
|
||||
}
|
||||
let msg_type = buf[0];
|
||||
let seq = &buf[4..8];
|
||||
let src = &buf[8..12];
|
||||
let dst = &buf[12..16];
|
||||
|
||||
let mut nonce12 = [0u8; 12];
|
||||
nonce12[0..4].copy_from_slice(seq);
|
||||
nonce12[4..8].copy_from_slice(dst);
|
||||
nonce12[8..12].copy_from_slice(src);
|
||||
nonce12[0] = msg_type;
|
||||
|
||||
Ok(nonce12)
|
||||
}
|
||||
|
||||
/// 原地加密(in-place)
|
||||
/// payload 后需要预留16字节用于存放 tag
|
||||
pub fn encrypt_in_place<B: AsRef<[u8]> + AsMut<[u8]>>(
|
||||
&self,
|
||||
pkt: &mut NetPacket<B>,
|
||||
) -> io::Result<()> {
|
||||
let nonce = Nonce::assume_unique_for_key(self.make_nonce(pkt)?);
|
||||
|
||||
let payload = pkt.payload_mut();
|
||||
let payload_len = payload.len() - TAG_LEN; // 实际 payload 长度(不含 tag 预留空间)
|
||||
|
||||
// 只加密实际的 payload 部分
|
||||
let tag = self
|
||||
.key
|
||||
.seal_in_place_separate_tag(nonce, Aad::empty(), &mut payload[..payload_len])
|
||||
.map_err(|_| io::Error::other("encrypt failed"))?;
|
||||
|
||||
// 将 tag 写入 payload 后的预留空间
|
||||
payload[payload_len..payload_len + TAG_LEN].copy_from_slice(tag.as_ref());
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 原地解密(in-place)
|
||||
pub fn decrypt_in_place<B: AsRef<[u8]> + AsMut<[u8]>>(
|
||||
&self,
|
||||
pkt: &mut NetPacket<B>,
|
||||
) -> io::Result<usize> {
|
||||
let nonce = Nonce::assume_unique_for_key(self.make_nonce(pkt)?);
|
||||
|
||||
let payload_with_tag = pkt.payload_mut();
|
||||
|
||||
let plaintext = self
|
||||
.key
|
||||
.open_in_place(nonce, Aad::empty(), payload_with_tag)
|
||||
.map_err(|_| io::Error::other("decrypt failed"))?;
|
||||
Ok(plaintext.len())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use bytes::BytesMut;
|
||||
|
||||
// 用于构造一个简单的 NetPacket,包含头 16 字节 + payload + 16 字节 TAG 预留
|
||||
fn build_test_packet(payload_len: usize) -> NetPacket<BytesMut> {
|
||||
// 16 字节 head + payload + 16 字节预留 TAG
|
||||
let total_len = HEAD_LENGTH + payload_len + TAG_LEN;
|
||||
let mut buf = BytesMut::zeroed(total_len);
|
||||
|
||||
// 构造一个头(16 字节)
|
||||
buf[0] = 4; // MsgType::Ping
|
||||
buf[4..8].copy_from_slice(&12345u32.to_be_bytes());
|
||||
buf[8..12].copy_from_slice(&111u32.to_be_bytes());
|
||||
buf[12..16].copy_from_slice(&222u32.to_be_bytes());
|
||||
|
||||
// 构造 payload(明文)
|
||||
let payload_plain = &mut buf[HEAD_LENGTH..HEAD_LENGTH + payload_len];
|
||||
for (i, p) in payload_plain.iter_mut().enumerate() {
|
||||
*p = (i as u8) ^ 0xAB;
|
||||
}
|
||||
|
||||
NetPacket::new(buf).unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_encrypt_decrypt_in_place() {
|
||||
let key = [7u8; 32];
|
||||
let crypto = PacketCrypto::new(key);
|
||||
|
||||
let payload_len = 20;
|
||||
let mut pkt = build_test_packet(payload_len);
|
||||
|
||||
// 备份原 payload
|
||||
let original_payload: Vec<u8> =
|
||||
pkt.buffer()[HEAD_LENGTH..HEAD_LENGTH + payload_len].to_vec();
|
||||
|
||||
// 加密
|
||||
crypto.encrypt_in_place(&mut pkt).expect("encrypt failed");
|
||||
|
||||
let encrypted_buf = pkt.buffer();
|
||||
let tag_start = HEAD_LENGTH + payload_len;
|
||||
let tag_end = tag_start + TAG_LEN;
|
||||
|
||||
// TAG 不应该是全 0
|
||||
assert_ne!(&encrypted_buf[tag_start..tag_end], &[0u8; TAG_LEN]);
|
||||
|
||||
// payload 已被加密,不等于明文
|
||||
assert_ne!(
|
||||
&encrypted_buf[HEAD_LENGTH..HEAD_LENGTH + payload_len],
|
||||
&original_payload[..]
|
||||
);
|
||||
|
||||
// 解密
|
||||
crypto.decrypt_in_place(&mut pkt).expect("decrypt failed");
|
||||
|
||||
let decrypted_buf = pkt.buffer();
|
||||
let decrypted_payload = &decrypted_buf[HEAD_LENGTH..HEAD_LENGTH + payload_len];
|
||||
|
||||
// 解密后与原文一致
|
||||
assert_eq!(decrypted_payload, &original_payload[..]);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
use crate::crypto::chacha20_poly1305::TAG_LEN;
|
||||
use crate::protocol::ip_packet_protocol::NetPacket;
|
||||
use std::io;
|
||||
use std::sync::Arc;
|
||||
|
||||
mod chacha20_poly1305;
|
||||
|
||||
use crate::protocol::transmission::{ExtendEnd, ShrinkEnd};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct PacketCrypto {
|
||||
crypto: Option<Arc<chacha20_poly1305::PacketCrypto>>,
|
||||
}
|
||||
impl PacketCrypto {
|
||||
pub(crate) fn key_sign(s: &str) -> String {
|
||||
chacha20_poly1305::PacketCrypto::key_sign(s)
|
||||
}
|
||||
|
||||
pub(crate) fn new_from_str(s: Option<&str>) -> Self {
|
||||
Self {
|
||||
crypto: s.map(chacha20_poly1305::PacketCrypto::new_from_str).map(Arc::new),
|
||||
}
|
||||
}
|
||||
pub(crate) fn encrypt_reserve(&self) -> usize {
|
||||
if self.crypto.is_some() { TAG_LEN } else { 0 }
|
||||
}
|
||||
pub(crate) fn encrypt_in_place<B: AsRef<[u8]> + AsMut<[u8]> + ExtendEnd>(
|
||||
&self,
|
||||
pkt: &mut NetPacket<B>,
|
||||
) -> io::Result<()> {
|
||||
if let Some(crypto) = self.crypto.as_ref() {
|
||||
pkt.source_buf_mut().extend_end(TAG_LEN);
|
||||
return crypto.encrypt_in_place(pkt);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
pub(crate) fn decrypt_in_place<B: AsRef<[u8]> + AsMut<[u8]> + ShrinkEnd>(
|
||||
&self,
|
||||
pkt: &mut NetPacket<B>,
|
||||
) -> io::Result<()> {
|
||||
if let Some(crypto) = self.crypto.as_ref() {
|
||||
let _ = crypto.decrypt_in_place(pkt)?;
|
||||
pkt.source_buf_mut().shrink_end(TAG_LEN);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
use crate::context::{NetworkAddr, TrafficStats};
|
||||
use crate::enhanced_tunnel::quic_over::quic_inbound::EnhancedQuicInbound;
|
||||
use crate::nat::internal_nat::InternalNatInbound;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tun::enhanced_tun::EnhancedTunInbound;
|
||||
use anyhow::{Context, bail};
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct EnhancedInbound {
|
||||
tun_data_inbound: EnhancedTunInbound,
|
||||
quic_inbound: EnhancedQuicInbound,
|
||||
internal_nat_inbound: Option<InternalNatInbound>,
|
||||
traffic_stats: TrafficStats,
|
||||
}
|
||||
|
||||
impl EnhancedInbound {
|
||||
pub fn new(
|
||||
tun_data_inbound: EnhancedTunInbound,
|
||||
quic_inbound: EnhancedQuicInbound,
|
||||
internal_nat_inbound: Option<InternalNatInbound>,
|
||||
traffic_stats: TrafficStats,
|
||||
) -> Self {
|
||||
Self {
|
||||
tun_data_inbound,
|
||||
quic_inbound,
|
||||
internal_nat_inbound,
|
||||
traffic_stats,
|
||||
}
|
||||
}
|
||||
pub async fn inbound(
|
||||
&self,
|
||||
network_addr: &NetworkAddr,
|
||||
msg_type: MsgType,
|
||||
src: Ipv4Addr,
|
||||
packet: NetPacket<TransmissionBytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut buf = packet.into_buffer();
|
||||
self.traffic_stats.record_rx(src, buf.len() as u64);
|
||||
buf.advance_head(HEAD_LENGTH)?;
|
||||
|
||||
match msg_type {
|
||||
MsgType::Turn => {
|
||||
if let Some(internal_nat_inbound) = self.internal_nat_inbound.as_ref() {
|
||||
let Some(ipv4) = Ipv4Packet::new(&buf) else {
|
||||
bail!("EnhancedInbound not ipv4")
|
||||
};
|
||||
let dest = ipv4.get_destination();
|
||||
if dest != network_addr.ip && !network_addr.network().contains(&dest) {
|
||||
internal_nat_inbound.send(&buf, network_addr).await?;
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
self.tun_data_inbound.inbound(buf, network_addr).await?;
|
||||
}
|
||||
MsgType::Broadcast | MsgType::ExcludeBroadcast => {
|
||||
self.tun_data_inbound.inbound(buf, network_addr).await?;
|
||||
}
|
||||
MsgType::Quic => {
|
||||
let payload = buf.into_bytes().freeze();
|
||||
self.quic_inbound
|
||||
.inbound(payload, src)
|
||||
.await
|
||||
.context("inbound quic")?;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
use crate::context::AppState;
|
||||
use crate::enhanced_tunnel::inbound::EnhancedInbound;
|
||||
use crate::enhanced_tunnel::outbound::EnhancedOutbound;
|
||||
use crate::nat::SubnetExternalRoute;
|
||||
use crate::nat::internal_nat::{InternalNatInbound, PortMappingManager};
|
||||
use crate::port_mapping::PortMapping;
|
||||
use crate::tun::enhanced_tun::EnhancedTunInbound;
|
||||
use crate::tunnel_core::outbound::HybridOutbound;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
|
||||
pub(crate) mod quic_over;
|
||||
|
||||
pub(crate) mod inbound;
|
||||
pub(crate) mod outbound;
|
||||
|
||||
pub(crate) struct TunnelConfig {
|
||||
pub mtu: u16,
|
||||
pub password: Option<String>,
|
||||
pub open_quic_client: bool,
|
||||
pub port_mapping: Vec<PortMapping>,
|
||||
}
|
||||
|
||||
pub(crate) struct TunnelComponents {
|
||||
pub hybrid_outbound: HybridOutbound,
|
||||
pub external_route: SubnetExternalRoute,
|
||||
pub internal_nat_inbound: Option<InternalNatInbound>,
|
||||
pub port_mapping_manager: PortMappingManager,
|
||||
}
|
||||
|
||||
pub(crate) async fn enhanced_ipv4_tunnel(
|
||||
app_state: AppState,
|
||||
task_group: TaskGroup,
|
||||
tun_data_sender: EnhancedTunInbound,
|
||||
config: TunnelConfig,
|
||||
components: TunnelComponents,
|
||||
) -> anyhow::Result<(EnhancedInbound, Option<EnhancedOutbound>)> {
|
||||
let password = config.password.unwrap_or_else(|| "password".to_string());
|
||||
let tun = match &tun_data_sender {
|
||||
EnhancedTunInbound::Tun(tun) => Some(tun.clone()),
|
||||
EnhancedTunInbound::Nat(_) => None,
|
||||
};
|
||||
let (inbound, outbound) = quic_over::boot::quic_tunnel_start(
|
||||
app_state.clone(),
|
||||
task_group,
|
||||
tun,
|
||||
quic_over::boot::QuicTunnelConfig {
|
||||
mtu: config.mtu,
|
||||
password,
|
||||
open_quic_client: config.open_quic_client,
|
||||
port_mapping: config.port_mapping,
|
||||
},
|
||||
quic_over::boot::QuicTunnelComponents {
|
||||
hybrid_outbound: components.hybrid_outbound.clone(),
|
||||
external_route: components.external_route,
|
||||
internal_nat_manager: components.internal_nat_inbound.clone(),
|
||||
port_mapping_manager: components.port_mapping_manager,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
let enhanced_inbound = EnhancedInbound::new(
|
||||
tun_data_sender,
|
||||
inbound,
|
||||
components.internal_nat_inbound,
|
||||
app_state.traffic_stats.clone(),
|
||||
);
|
||||
|
||||
let enhanced_outbound = outbound.map(|outbound| {
|
||||
EnhancedOutbound::new(
|
||||
app_state.network.clone(),
|
||||
outbound,
|
||||
components.hybrid_outbound,
|
||||
)
|
||||
});
|
||||
Ok((enhanced_inbound, enhanced_outbound))
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
use crate::context::SharedNetworkAddr;
|
||||
use crate::enhanced_tunnel::quic_over::quic_outbound::EnhancedQuicOutbound;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::outbound::HybridOutbound;
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
|
||||
pub struct EnhancedOutbound {
|
||||
network: SharedNetworkAddr,
|
||||
enhanced_quic_outbound: EnhancedQuicOutbound,
|
||||
hybrid_outbound: HybridOutbound,
|
||||
}
|
||||
|
||||
impl EnhancedOutbound {
|
||||
pub fn new(
|
||||
network: SharedNetworkAddr,
|
||||
enhanced_quic_outbound: EnhancedQuicOutbound,
|
||||
hybrid_outbound: HybridOutbound,
|
||||
) -> Self {
|
||||
Self {
|
||||
network,
|
||||
enhanced_quic_outbound,
|
||||
hybrid_outbound,
|
||||
}
|
||||
}
|
||||
pub async fn ipv4_outbound(&self, data: TransmissionBytes) {
|
||||
if data.is_empty() || data[0] >> 4 != 4 {
|
||||
return;
|
||||
}
|
||||
if let Err(e) = self.ipv4_outbound_impl(data).await {
|
||||
log::warn!("EnhancedOutbound error: {:?}", e);
|
||||
}
|
||||
}
|
||||
async fn ipv4_outbound_impl(&self, data: TransmissionBytes) -> anyhow::Result<()> {
|
||||
let Some(ipv4) = Ipv4Packet::new(data.as_ref()) else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(net) = self.network.get() else {
|
||||
return Ok(());
|
||||
};
|
||||
let src = ipv4.get_source();
|
||||
|
||||
let dest = ipv4.get_destination();
|
||||
if dest == src || dest.is_unspecified() {
|
||||
return Ok(());
|
||||
}
|
||||
if dest == net.gateway {
|
||||
// 发送到网关
|
||||
return self.hybrid_outbound.ipv4_gateway_outbound(net, data).await;
|
||||
}
|
||||
if dest.is_multicast() || dest == net.broadcast || dest.is_broadcast() {
|
||||
// 广播
|
||||
return self
|
||||
.hybrid_outbound
|
||||
.ipv4_broadcast_outbound(net, data)
|
||||
.await;
|
||||
}
|
||||
if self
|
||||
.enhanced_quic_outbound
|
||||
.outbound(&net, data.as_ref())
|
||||
.await
|
||||
{
|
||||
// 使用quic 通道传输
|
||||
return Ok(());
|
||||
}
|
||||
// 使用通用通道传输
|
||||
self.hybrid_outbound.ipv4_outbound(net, data).await
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
use crate::context::AppState;
|
||||
use crate::enhanced_tunnel::quic_over::enhanced_io::enhanced_inbound::{
|
||||
QuicDataInbound, create_enhanced_inbound,
|
||||
};
|
||||
use crate::enhanced_tunnel::quic_over::enhanced_io::enhanced_outbound::create_enhanced_outbound;
|
||||
use crate::enhanced_tunnel::quic_over::enhanced_io::socket::ExtendedQuicSocket;
|
||||
use crate::enhanced_tunnel::quic_over::quic_client::QuicTunnelClient;
|
||||
use crate::enhanced_tunnel::quic_over::quic_inbound::EnhancedQuicInbound;
|
||||
use crate::enhanced_tunnel::quic_over::quic_outbound::EnhancedQuicOutbound;
|
||||
use crate::enhanced_tunnel::quic_over::{quic_client, quic_server};
|
||||
use crate::nat::SubnetExternalRoute;
|
||||
use crate::nat::internal_nat::{InternalNatInbound, PortMappingManager};
|
||||
use crate::port_mapping::PortMapping;
|
||||
use crate::tls;
|
||||
use crate::tun::TunDataInbound;
|
||||
use crate::tunnel_core::outbound::HybridOutbound;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::Context;
|
||||
use quinn::congestion::BbrConfig;
|
||||
use quinn::crypto::rustls::QuicServerConfig;
|
||||
use quinn::{ClientConfig, Endpoint, EndpointConfig, TransportConfig, default_runtime};
|
||||
use rustls::ServerConfig;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::io;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tcp_ip::{IpStackConfig, IpStackRecv};
|
||||
|
||||
pub(crate) struct QuicTunnelConfig {
|
||||
pub mtu: u16,
|
||||
pub password: String,
|
||||
pub open_quic_client: bool,
|
||||
pub port_mapping: Vec<PortMapping>,
|
||||
}
|
||||
|
||||
pub(crate) struct QuicTunnelComponents {
|
||||
pub hybrid_outbound: HybridOutbound,
|
||||
pub external_route: SubnetExternalRoute,
|
||||
pub internal_nat_manager: Option<InternalNatInbound>,
|
||||
pub port_mapping_manager: PortMappingManager,
|
||||
}
|
||||
|
||||
pub(crate) async fn quic_tunnel_start(
|
||||
app_state: AppState,
|
||||
task_group: TaskGroup,
|
||||
tun_data_sender: Option<TunDataInbound>,
|
||||
config: QuicTunnelConfig,
|
||||
components: QuicTunnelComponents,
|
||||
) -> anyhow::Result<(EnhancedQuicInbound, Option<EnhancedQuicOutbound>)> {
|
||||
let ip_stack_config = IpStackConfig {
|
||||
mtu: config.mtu,
|
||||
..Default::default()
|
||||
};
|
||||
let (ip_stack, ip_socket, quic_outbound) = if let Some(tun_data_sender) = tun_data_sender {
|
||||
let (ip_stack, ip_stack_send, ip_stack_recv) = tcp_ip::ip_stack(ip_stack_config)?;
|
||||
let ip_socket = tcp_ip::ip::IpSocket::bind_all(None, ip_stack.clone()).await?;
|
||||
let ip_socket = Arc::new(ip_socket);
|
||||
task_group.spawn(ip_stack_recv_task(
|
||||
ip_stack_recv,
|
||||
app_state.clone(),
|
||||
tun_data_sender,
|
||||
));
|
||||
let quic_outbound =
|
||||
EnhancedQuicOutbound::new(config.open_quic_client, ip_stack_send, ip_stack.clone());
|
||||
|
||||
(Some(ip_stack), Some(ip_socket), Some(quic_outbound))
|
||||
} else {
|
||||
(None, None, None)
|
||||
};
|
||||
|
||||
let (inbound, endpoint) = create_quic_endpoint(
|
||||
config.password,
|
||||
task_group.clone(),
|
||||
components.hybrid_outbound,
|
||||
)
|
||||
.await?;
|
||||
quic_server::server_listen(
|
||||
&task_group,
|
||||
endpoint.clone(),
|
||||
ip_socket.clone(),
|
||||
ip_stack.clone(),
|
||||
components.internal_nat_manager,
|
||||
components.port_mapping_manager,
|
||||
)
|
||||
.await;
|
||||
if config.open_quic_client {
|
||||
let quic_client =
|
||||
QuicTunnelClient::new(app_state.clone(), endpoint, components.external_route);
|
||||
|
||||
// 客户端使用指纹验证
|
||||
if let (Some(ip_stack), Some(ip_socket)) = (ip_stack, ip_socket) {
|
||||
quic_client::create_client(
|
||||
quic_client.clone(),
|
||||
task_group.clone(),
|
||||
ip_stack.clone(),
|
||||
ip_socket,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
if !config.port_mapping.is_empty() {
|
||||
crate::port_mapping::port_mapping_start(&task_group, config.port_mapping, quic_client)
|
||||
.await?;
|
||||
}
|
||||
} else if !config.port_mapping.is_empty() {
|
||||
let quic_client =
|
||||
QuicTunnelClient::new(app_state.clone(), endpoint, components.external_route);
|
||||
|
||||
crate::port_mapping::port_mapping_start(&task_group, config.port_mapping, quic_client)
|
||||
.await?;
|
||||
}
|
||||
|
||||
let quic_inbound = EnhancedQuicInbound::new(inbound);
|
||||
Ok((quic_inbound, quic_outbound))
|
||||
}
|
||||
async fn create_quic_endpoint(
|
||||
password: String,
|
||||
task_group: TaskGroup,
|
||||
hybrid_outbound: HybridOutbound,
|
||||
) -> anyhow::Result<(QuicDataInbound, Endpoint)> {
|
||||
let (cert, private_key) = crate::tls::cert::generate_deterministic_cert(&password)?;
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(cert.as_ref());
|
||||
let calculated_hash: [u8; 32] = hasher.finalize().into();
|
||||
log::info!("QUIC Cert Fingerprint: {}", hex::encode(calculated_hash));
|
||||
|
||||
let outbound = create_enhanced_outbound(task_group.clone(), hybrid_outbound).await;
|
||||
let (inbound, inbound_receiver) = create_enhanced_inbound();
|
||||
let socket = ExtendedQuicSocket::new(inbound_receiver, outbound);
|
||||
|
||||
let server_config = ServerConfig::builder()
|
||||
.with_no_client_auth()
|
||||
.with_single_cert(vec![cert], private_key)
|
||||
.context("TLS config error")?;
|
||||
|
||||
let server_crypto = QuicServerConfig::try_from(server_config)
|
||||
.map_err(|e| anyhow::anyhow!("QUIC TLS config error: {:?}", e))?;
|
||||
let server_config = quinn::ServerConfig::with_crypto(Arc::new(server_crypto));
|
||||
// 替换运行时
|
||||
let runtime = default_runtime().ok_or_else(|| io::Error::other("no async runtime found"))?;
|
||||
let fingerprint_verifier = tls::verifier::FingerprintVerifier::new(calculated_hash);
|
||||
|
||||
let client_config = rustls::ClientConfig::builder()
|
||||
.dangerous()
|
||||
.with_custom_certificate_verifier(Arc::new(fingerprint_verifier))
|
||||
.with_no_client_auth();
|
||||
let mut client_config = ClientConfig::new(Arc::new(
|
||||
quinn::crypto::rustls::QuicClientConfig::try_from(client_config)
|
||||
.context("Failed to create QUIC client config")?,
|
||||
));
|
||||
client_config.transport_config(build_transport_config());
|
||||
let mut endpoint_config = EndpointConfig::default();
|
||||
endpoint_config.max_udp_payload_size(1300)?;
|
||||
let mut endpoint = quinn::Endpoint::new_with_abstract_socket(
|
||||
endpoint_config,
|
||||
Some(server_config),
|
||||
Arc::new(socket),
|
||||
runtime,
|
||||
)
|
||||
.context("quic server create failed")?;
|
||||
endpoint.set_default_client_config(client_config);
|
||||
Ok((inbound, endpoint))
|
||||
}
|
||||
|
||||
fn build_transport_config() -> Arc<TransportConfig> {
|
||||
let mut transport = TransportConfig::default();
|
||||
transport.congestion_controller_factory(Arc::new(BbrConfig::default()));
|
||||
transport.keep_alive_interval(Some(Duration::from_secs(5)));
|
||||
|
||||
transport.max_idle_timeout(Some(Duration::from_secs(10).try_into().unwrap()));
|
||||
|
||||
Arc::new(transport)
|
||||
}
|
||||
|
||||
async fn ip_stack_recv_task(
|
||||
mut ip_stack_recv: IpStackRecv,
|
||||
app_state: AppState,
|
||||
tun_data_sender: TunDataInbound,
|
||||
) {
|
||||
let mut buf = vec![0u8; 1500];
|
||||
loop {
|
||||
let len = match ip_stack_recv.recv(&mut buf).await {
|
||||
Ok(len) => len,
|
||||
Err(e) => {
|
||||
log::error!("IP stack recv error: {:?}", e);
|
||||
break;
|
||||
}
|
||||
};
|
||||
let Some(net) = app_state.get_network() else {
|
||||
log::error!("not network");
|
||||
break;
|
||||
};
|
||||
match tun_data_sender.send((&buf[..len]).into(), &net).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
log::error!("IP stack send error: {:?}", e);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
use anyhow::anyhow;
|
||||
use bytes::Bytes;
|
||||
use parking_lot::Mutex;
|
||||
use quinn::udp::RecvMeta;
|
||||
use std::fmt::{Debug, Formatter};
|
||||
use std::io::IoSliceMut;
|
||||
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
use tokio::sync::mpsc::{Receiver, Sender};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct QuicInnerInboundReceiver {
|
||||
receiver: Arc<Mutex<Receiver<(Bytes, Ipv4Addr)>>>,
|
||||
}
|
||||
#[derive(Clone)]
|
||||
pub struct QuicDataInbound {
|
||||
sender: Sender<(Bytes, Ipv4Addr)>,
|
||||
}
|
||||
impl QuicDataInbound {
|
||||
pub async fn send(&self, data: Bytes, addr: Ipv4Addr) -> anyhow::Result<()> {
|
||||
self.sender
|
||||
.send((data, addr))
|
||||
.await
|
||||
.map_err(|_e| anyhow!("quic data inbound error"))
|
||||
}
|
||||
}
|
||||
impl Debug for QuicInnerInboundReceiver {
|
||||
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("EnhancedInbound").finish()
|
||||
}
|
||||
}
|
||||
pub fn create_enhanced_inbound() -> (QuicDataInbound, QuicInnerInboundReceiver) {
|
||||
let (sender, receiver) = tokio::sync::mpsc::channel(256);
|
||||
(
|
||||
QuicDataInbound { sender },
|
||||
QuicInnerInboundReceiver::new(receiver),
|
||||
)
|
||||
}
|
||||
impl QuicInnerInboundReceiver {
|
||||
pub fn new(receiver: Receiver<(Bytes, Ipv4Addr)>) -> Self {
|
||||
Self {
|
||||
receiver: Arc::new(Mutex::new(receiver)),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn poll_recv(
|
||||
&self,
|
||||
cx: &mut Context,
|
||||
bufs: &mut [IoSliceMut<'_>],
|
||||
meta: &mut [RecvMeta],
|
||||
) -> Poll<std::io::Result<usize>> {
|
||||
let mut guard = self.receiver.lock();
|
||||
let rs = guard.poll_recv(cx);
|
||||
drop(guard);
|
||||
match rs {
|
||||
Poll::Ready(Some((buf, ip))) => {
|
||||
let (buf_mut, meta) = match (bufs.get_mut(0), meta.get_mut(0)) {
|
||||
(Some(b), Some(m)) => (b, m),
|
||||
_ => {
|
||||
return Poll::Ready(Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidInput,
|
||||
"no buffer available",
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
if buf_mut.len() < buf.len() {
|
||||
return Poll::Ready(Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidInput,
|
||||
format!(
|
||||
"buffer too small: need {}, got {}",
|
||||
buf.len(),
|
||||
buf_mut.len()
|
||||
),
|
||||
)));
|
||||
}
|
||||
|
||||
buf_mut[..buf.len()].copy_from_slice(&buf);
|
||||
|
||||
meta.len = buf.len();
|
||||
meta.stride = buf.len();
|
||||
meta.addr = SocketAddr::V4(SocketAddrV4::new(ip, 10000));
|
||||
Poll::Ready(Ok(1))
|
||||
}
|
||||
Poll::Ready(None) => Poll::Ready(Err(std::io::Error::new(
|
||||
std::io::ErrorKind::BrokenPipe,
|
||||
"inbound channel closed",
|
||||
))),
|
||||
Poll::Pending => Poll::Pending,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::outbound::HybridOutbound;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use quinn::UdpPoller;
|
||||
use std::fmt::{Debug, Formatter};
|
||||
use std::io;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
use tokio::sync::mpsc::{Sender, error::TrySendError};
|
||||
use tokio_util::sync::PollSender;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct QuicInnerOutbound {
|
||||
sender: Sender<(Ipv4Addr, NetPacket<TransmissionBytes>)>,
|
||||
}
|
||||
|
||||
pub async fn create_enhanced_outbound(
|
||||
task_group: TaskGroup,
|
||||
hybrid_outbound: HybridOutbound,
|
||||
) -> QuicInnerOutbound {
|
||||
let (s, mut r) = tokio::sync::mpsc::channel(256);
|
||||
|
||||
task_group.spawn(async move {
|
||||
while let Some((dst, packet)) = r.recv().await {
|
||||
if let Err(e) = hybrid_outbound.outbound_raw(dst, packet).await {
|
||||
log::debug!("outbound error: {e:?}, dst={dst}");
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
QuicInnerOutbound { sender: s }
|
||||
}
|
||||
|
||||
impl QuicInnerOutbound {
|
||||
pub fn try_outbound(&self, buf: &[u8], dest: Ipv4Addr) -> io::Result<()> {
|
||||
let send = match self.sender.try_reserve() {
|
||||
Ok(send) => send,
|
||||
Err(TrySendError::Full(_)) => {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::WouldBlock,
|
||||
"outbound channel full",
|
||||
));
|
||||
}
|
||||
Err(TrySendError::Closed(_)) => {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::BrokenPipe,
|
||||
"outbound channel closed",
|
||||
));
|
||||
}
|
||||
};
|
||||
let mut packet = NetPacket::new(TransmissionBytes::zeroed(HEAD_LENGTH + buf.len()))?;
|
||||
packet.set_ttl(5);
|
||||
packet.set_msg_type(MsgType::Quic);
|
||||
packet.set_dest_id(dest.into());
|
||||
packet.set_payload(buf)?;
|
||||
send.send((dest, packet));
|
||||
Ok(())
|
||||
}
|
||||
pub fn create_io_poller(&self) -> Pin<Box<dyn UdpPoller>> {
|
||||
Box::pin(EnhancedOutboundPoller {
|
||||
sender: PollSender::new(self.sender.clone()),
|
||||
})
|
||||
}
|
||||
}
|
||||
pub struct EnhancedOutboundPoller {
|
||||
sender: PollSender<(Ipv4Addr, NetPacket<TransmissionBytes>)>,
|
||||
}
|
||||
impl Debug for EnhancedOutboundPoller {
|
||||
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("EnhancedOutboundPoller").finish()
|
||||
}
|
||||
}
|
||||
impl UdpPoller for EnhancedOutboundPoller {
|
||||
fn poll_writable(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
|
||||
match self.sender.poll_reserve(cx) {
|
||||
Poll::Ready(Ok(_)) => {
|
||||
self.sender.abort_send();
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
Poll::Ready(Err(_e)) => Poll::Ready(Err(io::Error::new(
|
||||
io::ErrorKind::BrokenPipe,
|
||||
"outbound channel closed",
|
||||
))),
|
||||
Poll::Pending => Poll::Pending,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
pub mod enhanced_inbound;
|
||||
pub mod enhanced_outbound;
|
||||
pub mod socket;
|
||||
@@ -0,0 +1,55 @@
|
||||
use crate::enhanced_tunnel::quic_over::enhanced_io::enhanced_inbound::QuicInnerInboundReceiver;
|
||||
use crate::enhanced_tunnel::quic_over::enhanced_io::enhanced_outbound::QuicInnerOutbound;
|
||||
use quinn::udp::{RecvMeta, Transmit};
|
||||
use quinn::{AsyncUdpSocket, UdpPoller};
|
||||
use std::fmt::{Debug, Formatter};
|
||||
use std::io::IoSliceMut;
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
pub struct ExtendedQuicSocket {
|
||||
inbound: QuicInnerInboundReceiver,
|
||||
outbound: QuicInnerOutbound,
|
||||
}
|
||||
impl Debug for ExtendedQuicSocket {
|
||||
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("QuicSocket").finish()
|
||||
}
|
||||
}
|
||||
impl ExtendedQuicSocket {
|
||||
pub fn new(inbound: QuicInnerInboundReceiver, outbound: QuicInnerOutbound) -> Self {
|
||||
Self { inbound, outbound }
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncUdpSocket for ExtendedQuicSocket {
|
||||
fn create_io_poller(self: Arc<Self>) -> Pin<Box<dyn UdpPoller>> {
|
||||
self.outbound.create_io_poller()
|
||||
}
|
||||
|
||||
fn try_send(&self, transmit: &Transmit) -> std::io::Result<()> {
|
||||
let IpAddr::V4(dest) = transmit.destination.ip() else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
self.outbound.try_outbound(transmit.contents, dest)
|
||||
}
|
||||
|
||||
fn poll_recv(
|
||||
&self,
|
||||
cx: &mut Context,
|
||||
bufs: &mut [IoSliceMut<'_>],
|
||||
meta: &mut [RecvMeta],
|
||||
) -> Poll<std::io::Result<usize>> {
|
||||
self.inbound.poll_recv(cx, bufs, meta)
|
||||
}
|
||||
|
||||
fn local_addr(&self) -> std::io::Result<SocketAddr> {
|
||||
Ok(SocketAddr::V4(SocketAddrV4::new(
|
||||
Ipv4Addr::new(127, 0, 0, 1),
|
||||
10000,
|
||||
)))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
mod enhanced_io;
|
||||
pub(crate) mod quic_client;
|
||||
pub(crate) mod quic_inbound;
|
||||
pub(crate) mod quic_outbound;
|
||||
mod quic_server;
|
||||
|
||||
pub(crate) mod boot;
|
||||
@@ -0,0 +1,309 @@
|
||||
use crate::context::AppState;
|
||||
|
||||
use crate::nat::SubnetExternalRoute;
|
||||
use crate::protocol::client_message::{
|
||||
IpProxyHandshake, QuicProxyHandshake, TcpProxyHandshake, quic_proxy_handshake,
|
||||
};
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::{Context, bail};
|
||||
use bytes::Bytes;
|
||||
use futures::SinkExt;
|
||||
use parking_lot::Mutex;
|
||||
use pnet_packet::ip::IpNextHeaderProtocol;
|
||||
use prost::Message;
|
||||
use quinn::{Connection, Endpoint, RecvStream, SendStream};
|
||||
use std::collections::HashMap;
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
use std::sync::Arc;
|
||||
use tcp_ip::IpStack;
|
||||
use tcp_ip::ip::IpSocket;
|
||||
use tcp_ip::tcp::TcpStream;
|
||||
use tokio::io::AsyncWriteExt;
|
||||
use tokio::sync::OnceCell;
|
||||
use tokio::sync::mpsc::Sender;
|
||||
use tokio::sync::mpsc::error::TrySendError;
|
||||
use tokio_util::codec::{FramedWrite, LengthDelimitedCodec};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct QuicTunnelClient {
|
||||
app_state: AppState,
|
||||
endpoint: Endpoint,
|
||||
connection_map: Arc<Mutex<HashMap<Ipv4Addr, Arc<OnceCell<Connection>>>>>,
|
||||
external_route: SubnetExternalRoute,
|
||||
}
|
||||
|
||||
impl QuicTunnelClient {
|
||||
pub fn new(
|
||||
app_state: AppState,
|
||||
endpoint: Endpoint,
|
||||
external_route: SubnetExternalRoute,
|
||||
) -> QuicTunnelClient {
|
||||
Self {
|
||||
app_state,
|
||||
endpoint,
|
||||
connection_map: Arc::new(Default::default()),
|
||||
external_route,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn open_bi(&self, mut dest: Ipv4Addr) -> anyhow::Result<(SendStream, RecvStream)> {
|
||||
let Some(net) = self.app_state.get_network() else {
|
||||
bail!("no network found");
|
||||
};
|
||||
if !net.network().contains(&dest) {
|
||||
if let Some(v) = self.external_route.route(&dest) {
|
||||
dest = v;
|
||||
} else {
|
||||
bail!("invalid route found:{dest}");
|
||||
}
|
||||
}
|
||||
let mut count = 0;
|
||||
loop {
|
||||
count += 1;
|
||||
let cell = self
|
||||
.connection_map
|
||||
.lock()
|
||||
.entry(dest)
|
||||
.or_insert_with(|| Arc::new(OnceCell::new()))
|
||||
.clone();
|
||||
let connection = cell
|
||||
.get_or_try_init(|| async {
|
||||
self.endpoint
|
||||
.connect(SocketAddr::new(dest.into(), 10000), "localhost")?
|
||||
.await
|
||||
.context("connect failed")
|
||||
})
|
||||
.await?;
|
||||
|
||||
return match connection.open_bi().await {
|
||||
Ok(rs) => Ok(rs),
|
||||
Err(e) => {
|
||||
self.connection_map.lock().remove(&dest);
|
||||
if count == 1 {
|
||||
continue;
|
||||
}
|
||||
Err(e.into())
|
||||
}
|
||||
};
|
||||
}
|
||||
}
|
||||
pub async fn open_uni(&self, mut dest: Ipv4Addr) -> anyhow::Result<SendStream> {
|
||||
let Some(net) = self.app_state.get_network() else {
|
||||
bail!("no network found");
|
||||
};
|
||||
if !net.network().contains(&dest) {
|
||||
if let Some(v) = self.external_route.route(&dest) {
|
||||
dest = v;
|
||||
} else {
|
||||
bail!("invalid route found:{dest}");
|
||||
}
|
||||
}
|
||||
let mut count = 0;
|
||||
loop {
|
||||
count += 1;
|
||||
let cell = self
|
||||
.connection_map
|
||||
.lock()
|
||||
.entry(dest)
|
||||
.or_insert_with(|| Arc::new(OnceCell::new()))
|
||||
.clone();
|
||||
let connection = cell
|
||||
.get_or_try_init(|| async {
|
||||
self.endpoint
|
||||
.connect(SocketAddr::new(dest.into(), 10000), "localhost")?
|
||||
.await
|
||||
.context("connect failed")
|
||||
})
|
||||
.await?;
|
||||
return match connection.open_uni().await {
|
||||
Ok(rs) => Ok(rs),
|
||||
Err(e) => {
|
||||
self.connection_map.lock().remove(&dest);
|
||||
if count == 1 {
|
||||
continue;
|
||||
}
|
||||
Err(e.into())
|
||||
}
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
pub(crate) async fn send_handshake(
|
||||
send_stream: &mut SendStream,
|
||||
handshake: QuicProxyHandshake,
|
||||
) -> anyhow::Result<()> {
|
||||
let handshake = handshake.encode_to_vec();
|
||||
send_stream.write_u16(handshake.len() as u16).await?;
|
||||
send_stream.write_all(&handshake).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn create_client(
|
||||
quic_client: QuicTunnelClient,
|
||||
task_group: TaskGroup,
|
||||
ip_stack: IpStack,
|
||||
ip_socket: Arc<IpSocket>,
|
||||
) {
|
||||
task_group.spawn(tcp_listen(
|
||||
task_group.clone(),
|
||||
ip_stack.clone(),
|
||||
quic_client.clone(),
|
||||
));
|
||||
task_group.spawn(ip_listen(task_group.clone(), ip_socket, quic_client));
|
||||
}
|
||||
|
||||
async fn tcp_listen(
|
||||
task_group: TaskGroup,
|
||||
ip_stack: IpStack,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) {
|
||||
if let Err(e) = tcp_listen_impl(task_group, ip_stack, quic_tunnel_client).await {
|
||||
log::error!("tcp_listen {e:?}");
|
||||
}
|
||||
}
|
||||
|
||||
async fn tcp_listen_impl(
|
||||
task_group: TaskGroup,
|
||||
ip_stack: IpStack,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut listener = tcp_ip::tcp::TcpListener::bind_all(ip_stack).await?;
|
||||
loop {
|
||||
let (tcp_stream, addr) = listener.accept().await?;
|
||||
let quic_tunnel_client = quic_tunnel_client.clone();
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = tcp_stream_handle(tcp_stream, quic_tunnel_client).await {
|
||||
log::error!("TCP stream handle failed with error: {e:?},addr={addr}");
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async fn tcp_stream_handle(
|
||||
tcp_stream: TcpStream,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
// 连接方向是反过来的,因为自己充当目标做了tcp卸载
|
||||
let SocketAddr::V4(peer_addr) = tcp_stream.local_addr()? else {
|
||||
bail!("invalid IP address");
|
||||
};
|
||||
let SocketAddr::V4(local_addr) = tcp_stream.peer_addr()? else {
|
||||
bail!("invalid IP address");
|
||||
};
|
||||
log::debug!("connect TCP stream {}->{}", local_addr, peer_addr);
|
||||
|
||||
let (mut send_stream, mut recv_stream) = quic_tunnel_client.open_bi(*peer_addr.ip()).await?;
|
||||
let handshake = QuicProxyHandshake {
|
||||
handshake: Some(quic_proxy_handshake::Handshake::Tcp(TcpProxyHandshake {
|
||||
src_ip: (*local_addr.ip()).into(),
|
||||
src_port: local_addr.port().into(),
|
||||
dst_ip: (*peer_addr.ip()).into(),
|
||||
dst_port: peer_addr.port().into(),
|
||||
})),
|
||||
};
|
||||
send_handshake(&mut send_stream, handshake).await?;
|
||||
let (mut tcp_w, mut tcp_r) = tcp_stream.split()?;
|
||||
tokio::select! {
|
||||
_ = tokio::io::copy(&mut recv_stream, &mut tcp_w) => {},
|
||||
_ = tokio::io::copy(&mut tcp_r, &mut send_stream) => {},
|
||||
}
|
||||
log::debug!("disconnect TCP stream {}->{}", local_addr, peer_addr);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn ip_listen(
|
||||
task_group: TaskGroup,
|
||||
ip_socket: Arc<IpSocket>,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) {
|
||||
if let Err(e) = ip_listen_impl(task_group, ip_socket, quic_tunnel_client).await {
|
||||
log::error!("ip_listen {e:?}");
|
||||
}
|
||||
}
|
||||
#[derive(Eq, PartialEq, Hash, Copy, Clone, Debug)]
|
||||
struct IpKey {
|
||||
protocol: IpNextHeaderProtocol,
|
||||
src: Ipv4Addr,
|
||||
dest: Ipv4Addr,
|
||||
}
|
||||
async fn ip_listen_impl(
|
||||
task_group: TaskGroup,
|
||||
ip_socket: Arc<IpSocket>,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut buf = vec![0u8; 65536];
|
||||
let dest_map = Arc::new(Mutex::new(HashMap::<IpKey, Sender<Bytes>>::new()));
|
||||
|
||||
loop {
|
||||
let (len, protocol, src, dest) = ip_socket.recv_protocol_from_to(&mut buf).await?;
|
||||
let (IpAddr::V4(src), IpAddr::V4(dest)) = (src, dest) else {
|
||||
continue;
|
||||
};
|
||||
let key = IpKey {
|
||||
protocol,
|
||||
src,
|
||||
dest,
|
||||
};
|
||||
let bytes = Bytes::copy_from_slice(&buf[..len]);
|
||||
|
||||
let tx = {
|
||||
let mut map = dest_map.lock();
|
||||
if let Some(tx) = map.get(&key) {
|
||||
tx.clone()
|
||||
} else {
|
||||
let (tx, rx) = tokio::sync::mpsc::channel::<Bytes>(128);
|
||||
|
||||
spawn_dest_sender(task_group.clone(), key, rx, quic_tunnel_client.clone());
|
||||
|
||||
map.insert(key, tx.clone());
|
||||
tx
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(err) = tx.try_send(bytes) {
|
||||
match err {
|
||||
TrySendError::Full(_) => {}
|
||||
TrySendError::Closed(_) => {
|
||||
let mut map = dest_map.lock();
|
||||
map.remove(&key);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn spawn_dest_sender(
|
||||
task_group: TaskGroup,
|
||||
key: IpKey,
|
||||
mut rx: tokio::sync::mpsc::Receiver<Bytes>,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) {
|
||||
log::info!("send ip({}) packet {}->{}", key.protocol, key.src, key.dest);
|
||||
task_group.spawn(async move {
|
||||
let result = async {
|
||||
let mut send_stream = quic_tunnel_client.open_uni(key.dest).await?;
|
||||
|
||||
let handshake = QuicProxyHandshake {
|
||||
handshake: Some(quic_proxy_handshake::Handshake::Ip(IpProxyHandshake {
|
||||
ip_next_header_protocol: key.protocol.0 as _,
|
||||
src_ip: key.src.into(),
|
||||
dst_ip: key.dest.into(),
|
||||
})),
|
||||
};
|
||||
send_handshake(&mut send_stream, handshake).await?;
|
||||
|
||||
let mut framed = FramedWrite::new(send_stream, LengthDelimitedCodec::new());
|
||||
|
||||
while let Some(pkt) = rx.recv().await {
|
||||
framed.send(pkt).await?;
|
||||
}
|
||||
|
||||
Ok::<(), anyhow::Error>(())
|
||||
}
|
||||
.await;
|
||||
|
||||
if let Err(e) = result {
|
||||
log::error!("key {:?} sender task exit: {:?}", key, e);
|
||||
}
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
use crate::enhanced_tunnel::quic_over::enhanced_io::enhanced_inbound::QuicDataInbound;
|
||||
use bytes::Bytes;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct EnhancedQuicInbound {
|
||||
quic_data_inbound: QuicDataInbound,
|
||||
}
|
||||
|
||||
impl EnhancedQuicInbound {
|
||||
pub fn new(quic_data_inbound: QuicDataInbound) -> Self {
|
||||
Self { quic_data_inbound }
|
||||
}
|
||||
pub async fn inbound(&self, data: Bytes, src: Ipv4Addr) -> anyhow::Result<()> {
|
||||
self.quic_data_inbound.send(data, src).await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
use crate::context::NetworkAddr;
|
||||
use pnet_packet::Packet;
|
||||
use pnet_packet::ip::IpNextHeaderProtocols;
|
||||
use pnet_packet::ipv4::{Ipv4Flags, Ipv4Packet};
|
||||
use pnet_packet::tcp::TcpFlags::{ACK, SYN};
|
||||
use pnet_packet::tcp::TcpPacket;
|
||||
use std::net::SocketAddr;
|
||||
use tcp_ip::{IpStack, IpStackSend};
|
||||
|
||||
pub struct EnhancedQuicOutbound {
|
||||
open_quic_client: bool,
|
||||
ip_stack_send: IpStackSend,
|
||||
ip_stack: IpStack,
|
||||
}
|
||||
|
||||
impl EnhancedQuicOutbound {
|
||||
pub fn new(open_quic_client: bool, ip_stack_send: IpStackSend, ip_stack: IpStack) -> Self {
|
||||
Self {
|
||||
open_quic_client,
|
||||
ip_stack_send,
|
||||
ip_stack,
|
||||
}
|
||||
}
|
||||
pub async fn outbound(&self, _net: &NetworkAddr, data: &[u8]) -> bool {
|
||||
let Some(ipv4) = Ipv4Packet::new(data) else {
|
||||
return true;
|
||||
};
|
||||
|
||||
if self.open_quic_client {
|
||||
// 针对tcp 如果不是从IpStack建立的连接,则不使用IpStack解析
|
||||
if ipv4.get_next_level_protocol() == IpNextHeaderProtocols::Tcp {
|
||||
let more_fragments =
|
||||
ipv4.get_flags() & Ipv4Flags::MoreFragments == Ipv4Flags::MoreFragments;
|
||||
let offset = ipv4.get_fragment_offset();
|
||||
let segmented = more_fragments || offset > 0;
|
||||
if !segmented {
|
||||
let Some(tcp) = TcpPacket::new(ipv4.payload()) else {
|
||||
return true;
|
||||
};
|
||||
// 不是第一个包
|
||||
if !(tcp.get_flags() & SYN == SYN && tcp.get_flags() & ACK != ACK) {
|
||||
let local_addr =
|
||||
SocketAddr::new(ipv4.get_source().into(), tcp.get_source());
|
||||
let peer_addr =
|
||||
SocketAddr::new(ipv4.get_destination().into(), tcp.get_destination());
|
||||
// 在IpStack中找不到连接
|
||||
if !self
|
||||
.ip_stack
|
||||
.has_tcp_connection(local_addr, peer_addr)
|
||||
.unwrap_or(false)
|
||||
&& !self
|
||||
.ip_stack
|
||||
.has_tcp_connection(peer_addr, local_addr)
|
||||
.unwrap_or(false)
|
||||
&& !self
|
||||
.ip_stack
|
||||
.has_tcp_half_open(peer_addr, local_addr)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = self.ip_stack_send.send_ip_packet(data).await;
|
||||
return true;
|
||||
}
|
||||
// 判断tcp流
|
||||
if ipv4.get_next_level_protocol() == IpNextHeaderProtocols::Tcp {
|
||||
let more_fragments =
|
||||
ipv4.get_flags() & Ipv4Flags::MoreFragments == Ipv4Flags::MoreFragments;
|
||||
let offset = ipv4.get_fragment_offset();
|
||||
let segmented = more_fragments || offset > 0;
|
||||
if !segmented && let Some(tcp) = TcpPacket::new(ipv4.payload()) {
|
||||
// 如果对端使用IpStack连接了自己,则也需要原路回复
|
||||
// 这是连接回复,所以方向是和流方向相反的
|
||||
let peer_addr = SocketAddr::new(ipv4.get_source().into(), tcp.get_source());
|
||||
let local_addr =
|
||||
SocketAddr::new(ipv4.get_destination().into(), tcp.get_destination());
|
||||
|
||||
if self
|
||||
.ip_stack
|
||||
.has_tcp_connection(local_addr, peer_addr)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
_ = self.ip_stack_send.send_ip_packet(data).await;
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,245 @@
|
||||
use crate::nat::internal_nat::{InternalNatInbound, PortMappingManager};
|
||||
use crate::protocol::client_message::QuicProxyHandshake;
|
||||
use crate::protocol::client_message::quic_proxy_handshake::Handshake;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::{Context, bail};
|
||||
use futures::StreamExt;
|
||||
use pnet_packet::ip::IpNextHeaderProtocol;
|
||||
use prost::Message;
|
||||
use quinn::{Connection, Endpoint, RecvStream, SendStream};
|
||||
use std::net::{Ipv4Addr, SocketAddr};
|
||||
use std::sync::Arc;
|
||||
use tcp_ip::IpStack;
|
||||
use tcp_ip::ip::IpSocket;
|
||||
use tokio::io::AsyncReadExt;
|
||||
use tokio_util::codec::{FramedRead, LengthDelimitedCodec};
|
||||
|
||||
pub async fn server_listen(
|
||||
task_group: &TaskGroup,
|
||||
endpoint: Endpoint,
|
||||
ip_socket: Option<Arc<IpSocket>>,
|
||||
ip_stack: Option<IpStack>,
|
||||
internal_nat_manager: Option<InternalNatInbound>,
|
||||
port_mapping_manager: PortMappingManager,
|
||||
) {
|
||||
task_group.spawn(quic_endpoint_accept(
|
||||
ip_stack,
|
||||
task_group.clone(),
|
||||
endpoint,
|
||||
ip_socket,
|
||||
internal_nat_manager,
|
||||
port_mapping_manager,
|
||||
));
|
||||
}
|
||||
|
||||
async fn quic_endpoint_accept(
|
||||
ip_stack: Option<IpStack>,
|
||||
task_group: TaskGroup,
|
||||
endpoint: Endpoint,
|
||||
ip_socket: Option<Arc<IpSocket>>,
|
||||
internal_nat_manager: Option<InternalNatInbound>,
|
||||
port_mapping_manager: PortMappingManager,
|
||||
) {
|
||||
while let Some(connecting) = endpoint.accept().await {
|
||||
let remote_addr = connecting.remote_address();
|
||||
let task_group_clone = task_group.clone();
|
||||
let ip_socket = ip_socket.clone();
|
||||
let ip_stack = ip_stack.clone();
|
||||
let internal_nat_manager = internal_nat_manager.clone();
|
||||
let port_mapping_manager = port_mapping_manager.clone();
|
||||
task_group.spawn(async move {
|
||||
match connecting.await {
|
||||
Ok(connection) => {
|
||||
log::info!("QUIC connection: {}", remote_addr);
|
||||
if let Err(e) = quic_accept(
|
||||
ip_stack,
|
||||
task_group_clone,
|
||||
connection,
|
||||
ip_socket,
|
||||
internal_nat_manager,
|
||||
port_mapping_manager,
|
||||
)
|
||||
.await
|
||||
{
|
||||
log::info!("quic close: {remote_addr},{e:?}",);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
log::error!("connect: {:?},remote_addr={remote_addr}", e);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
log::warn!("quic server closed");
|
||||
}
|
||||
|
||||
async fn quic_accept(
|
||||
ip_stack: Option<IpStack>,
|
||||
task_group_clone: TaskGroup,
|
||||
connection: Connection,
|
||||
ip_socket: Option<Arc<IpSocket>>,
|
||||
internal_nat_manager: Option<InternalNatInbound>,
|
||||
port_mapping_manager: PortMappingManager,
|
||||
) -> anyhow::Result<()> {
|
||||
loop {
|
||||
tokio::select! {
|
||||
rs = connection.accept_bi()=>{
|
||||
let (send_stream, recv_stream) = rs?;
|
||||
let ip_stack = ip_stack.clone();
|
||||
let internal_nat_manager = internal_nat_manager.clone();
|
||||
let port_mapping_manager = port_mapping_manager.clone();
|
||||
task_group_clone.spawn(async move {
|
||||
if let Err(e) = quic_stream_bi_handle(ip_stack,send_stream, recv_stream,&internal_nat_manager,port_mapping_manager).await{
|
||||
log::error!("quic_stream_bi_handle: {e:?}");
|
||||
}
|
||||
});
|
||||
}
|
||||
rs = connection.accept_uni()=>{
|
||||
let recv_stream = rs?;
|
||||
let ip_socket = ip_socket.clone();
|
||||
let internal_nat_manager = internal_nat_manager.clone();
|
||||
task_group_clone.spawn(async move {
|
||||
if let Err(e) = quic_stream_uni_handle(recv_stream, ip_socket,&internal_nat_manager).await{
|
||||
log::error!("quic_stream_uni_handle: {e:?}");
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn quic_stream_bi_handle(
|
||||
ip_stack: Option<IpStack>,
|
||||
mut send_stream: SendStream,
|
||||
mut recv_stream: RecvStream,
|
||||
internal_nat_manager: &Option<InternalNatInbound>,
|
||||
port_mapping_manager: PortMappingManager,
|
||||
) -> anyhow::Result<()> {
|
||||
let handshake = recv_handshake(&mut recv_stream).await?;
|
||||
let Some(handshake) = handshake.handshake else {
|
||||
return Ok(());
|
||||
};
|
||||
match handshake {
|
||||
Handshake::Tcp(handshake) => {
|
||||
let src = SocketAddr::new(
|
||||
Ipv4Addr::from(handshake.src_ip).into(),
|
||||
handshake.src_port as _,
|
||||
);
|
||||
let dst_ip = Ipv4Addr::from(handshake.dst_ip);
|
||||
let dst = SocketAddr::new(dst_ip.into(), handshake.dst_port as _);
|
||||
if src == dst {
|
||||
bail!("tcp handshake failed, ip: {}", src);
|
||||
}
|
||||
log::debug!("accept TCP stream {src}->{dst}");
|
||||
// 如果不是网段内的,并且启用了内置nat,则直接转发
|
||||
if let Some(internal_nat_manager) = internal_nat_manager {
|
||||
if internal_nat_manager.use_nat(&dst_ip) {
|
||||
internal_nat_manager
|
||||
.tcp_nat(recv_stream, send_stream, dst_ip, dst.port())
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
if internal_nat_manager.no_tun() {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
if let Some(ip_stack) = ip_stack {
|
||||
let stream = tcp_ip::tcp::TcpStream::bind(ip_stack, src)?
|
||||
.connect_to(dst)
|
||||
.await?;
|
||||
let (mut tcp_w, mut tcp_r) = stream.split()?;
|
||||
tokio::select! {
|
||||
_ = tokio::io::copy(&mut recv_stream, &mut tcp_w) => {},
|
||||
_ = tokio::io::copy(&mut tcp_r, &mut send_stream) => {},
|
||||
}
|
||||
log::debug!("accept close TCP stream {src}->{dst}");
|
||||
}
|
||||
}
|
||||
Handshake::Ip(_) => {}
|
||||
Handshake::TcpPortMapping(handshake) => {
|
||||
port_mapping_manager
|
||||
.tcp_mapping(
|
||||
recv_stream,
|
||||
send_stream,
|
||||
handshake.dst_host,
|
||||
handshake.dst_port as _,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
Handshake::UdpPortMapping(handshake) => {
|
||||
port_mapping_manager
|
||||
.udp_mapping(
|
||||
recv_stream,
|
||||
send_stream,
|
||||
handshake.dst_host,
|
||||
handshake.dst_port as _,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn recv_handshake(recv_stream: &mut RecvStream) -> anyhow::Result<QuicProxyHandshake> {
|
||||
let len = recv_stream.read_u16().await?;
|
||||
let mut buf = vec![0u8; len as usize];
|
||||
recv_stream.read_exact(&mut buf).await?;
|
||||
let handshake = QuicProxyHandshake::decode(&buf[..])?;
|
||||
Ok(handshake)
|
||||
}
|
||||
async fn quic_stream_uni_handle(
|
||||
mut recv_stream: RecvStream,
|
||||
ip_socket: Option<Arc<IpSocket>>,
|
||||
internal_nat_manager: &Option<InternalNatInbound>,
|
||||
) -> anyhow::Result<()> {
|
||||
let handshake = recv_handshake(&mut recv_stream).await?;
|
||||
let Some(handshake) = handshake.handshake else {
|
||||
return Ok(());
|
||||
};
|
||||
match handshake {
|
||||
Handshake::Tcp(_) => {}
|
||||
Handshake::Ip(handshake) => {
|
||||
let ip_next_header_protocol =
|
||||
IpNextHeaderProtocol::new(handshake.ip_next_header_protocol as _);
|
||||
let src_ip = Ipv4Addr::from(handshake.src_ip);
|
||||
let dest_ip = Ipv4Addr::from(handshake.dst_ip);
|
||||
log::debug!("recv IP({ip_next_header_protocol}) packet {src_ip}->{dest_ip}");
|
||||
let mut framed_read = FramedRead::new(recv_stream, LengthDelimitedCodec::new());
|
||||
// 如果不是网段内的,并且启用了内置nat,则直接转发
|
||||
if let Some(internal_nat_manager) = internal_nat_manager {
|
||||
if internal_nat_manager.use_nat(&dest_ip) {
|
||||
loop {
|
||||
let buf = framed_read
|
||||
.next()
|
||||
.await
|
||||
.context("receive quic stream failed")??;
|
||||
internal_nat_manager
|
||||
.send_ipv4_payload(ip_next_header_protocol, src_ip, dest_ip, buf)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
if internal_nat_manager.no_tun() {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
let Some(ip_socket) = ip_socket else {
|
||||
return Ok(());
|
||||
};
|
||||
let src_ip = src_ip.into();
|
||||
let dest_ip = dest_ip.into();
|
||||
loop {
|
||||
let buf = framed_read
|
||||
.next()
|
||||
.await
|
||||
.context("receive quic stream failed")??;
|
||||
|
||||
ip_socket
|
||||
.send_protocol_from_to(&buf, ip_next_header_protocol, src_ip, dest_ip)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
Handshake::TcpPortMapping(_) => {}
|
||||
Handshake::UdpPortMapping(_) => {}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,302 @@
|
||||
use crate::fec::encoder::FecPacket;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use anyhow::{Result, bail};
|
||||
use prost::Message;
|
||||
use reed_solomon_erasure::galois_8::ReedSolomon;
|
||||
use std::collections::HashMap;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
const GROUP_TIMEOUT: Duration = Duration::from_secs(3);
|
||||
const MAX_GROUPS: usize = 1000;
|
||||
const MAX_NUM: usize = 50;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct FecDecoder {
|
||||
inner: Arc<parking_lot::Mutex<FecDecoderInner>>,
|
||||
}
|
||||
|
||||
struct FecDecoderInner {
|
||||
groups: HashMap<(Ipv4Addr, u64), FecGroup>,
|
||||
last_cleanup: Instant,
|
||||
}
|
||||
|
||||
struct FecGroup {
|
||||
data_shards: usize,
|
||||
parity_shards: usize,
|
||||
received_original_count: usize,
|
||||
received_shards: Vec<Option<Vec<u8>>>,
|
||||
last_update: Instant,
|
||||
}
|
||||
impl FecGroup {
|
||||
fn is_done(&self) -> bool {
|
||||
self.data_shards != 0 && self.received_original_count == self.data_shards
|
||||
}
|
||||
fn done(&mut self) {
|
||||
self.received_original_count = self.data_shards;
|
||||
self.received_shards = vec![];
|
||||
}
|
||||
}
|
||||
impl Default for FecGroup {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
data_shards: 0,
|
||||
parity_shards: 0,
|
||||
received_original_count: 0,
|
||||
received_shards: Vec::with_capacity(16),
|
||||
last_update: Instant::now(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl FecDecoder {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
inner: Arc::new(parking_lot::Mutex::new(FecDecoderInner {
|
||||
groups: HashMap::new(),
|
||||
last_cleanup: Instant::now(),
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
/// 接收FEC包并尝试恢复丢失的包
|
||||
pub fn receive(
|
||||
&self,
|
||||
net_packet: NetPacket<TransmissionBytes>,
|
||||
) -> Result<Option<Vec<NetPacket<TransmissionBytes>>>> {
|
||||
let mut inner = self.inner.lock();
|
||||
let src_ip = Ipv4Addr::from(net_packet.src_id());
|
||||
let fec_packet = FecPacket::decode(net_packet.payload())?;
|
||||
|
||||
let group_id = fec_packet.group_id;
|
||||
let packet_index = fec_packet.packet_index as usize;
|
||||
let payload = fec_packet.payload;
|
||||
|
||||
if packet_index > MAX_NUM {
|
||||
log::warn!(
|
||||
"packet_index overflow, src={src_ip},group_id={group_id}, packet_index={packet_index}",
|
||||
);
|
||||
bail!("packet_index overflow {src_ip}");
|
||||
}
|
||||
let mut packet = None;
|
||||
let group = inner.groups.entry((src_ip, group_id)).or_default();
|
||||
if group.is_done() {
|
||||
return Ok(None);
|
||||
}
|
||||
if group
|
||||
.received_shards
|
||||
.get(packet_index)
|
||||
.is_some_and(|v| v.is_some())
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
if let Some(parity_data) = fec_packet.parity_data {
|
||||
let data_shards = parity_data.data_shards as usize;
|
||||
let parity_shards = parity_data.parity_shards as usize;
|
||||
if data_shards > MAX_NUM {
|
||||
bail!("data_shards overflow {src_ip}");
|
||||
}
|
||||
if parity_shards > MAX_NUM {
|
||||
bail!("parity_shards overflow {src_ip}");
|
||||
}
|
||||
|
||||
if data_shards + parity_shards <= packet_index {
|
||||
log::warn!(
|
||||
"packet_index overflow in parity, src={},group_id={}, packet_index={}, total_shards={}",
|
||||
src_ip,
|
||||
group_id,
|
||||
packet_index,
|
||||
data_shards + parity_shards
|
||||
);
|
||||
bail!("packet_index overflow {src_ip}");
|
||||
}
|
||||
if group.data_shards != 0 && group.data_shards != data_shards {
|
||||
bail!("group data_shards!=data_shards {src_ip}");
|
||||
}
|
||||
if group.parity_shards != 0 && group.parity_shards != parity_shards {
|
||||
bail!("group parity_shards!=parity_shards {src_ip}");
|
||||
}
|
||||
group.data_shards = data_shards;
|
||||
group.parity_shards = parity_shards;
|
||||
if group.received_shards.len() < data_shards + parity_shards {
|
||||
group
|
||||
.received_shards
|
||||
.resize(data_shards + parity_shards, None);
|
||||
}
|
||||
|
||||
group.received_shards[packet_index] = Some(payload);
|
||||
} else {
|
||||
let buffer = TransmissionBytes::zeroed(HEAD_LENGTH + payload.len());
|
||||
let mut result_packet = NetPacket::new(buffer)?;
|
||||
result_packet.head_mut().copy_from_slice(net_packet.head());
|
||||
result_packet.set_fec_flag(false);
|
||||
result_packet.set_payload(&payload)?;
|
||||
packet = Some(result_packet);
|
||||
|
||||
if group.received_shards.len() <= packet_index {
|
||||
group.received_shards.resize(packet_index + 1, None);
|
||||
}
|
||||
|
||||
// 保存FEC数据: [type_byte, flags_byte, payload_len(u16), payload...]
|
||||
let type_byte = net_packet.head()[0];
|
||||
let flags_byte = net_packet.head()[2];
|
||||
let mut batch_data = vec![0u8; 4 + payload.len()];
|
||||
batch_data[0] = type_byte;
|
||||
batch_data[1] = flags_byte;
|
||||
batch_data[2..4].copy_from_slice(&(payload.len() as u16).to_be_bytes());
|
||||
batch_data[4..].copy_from_slice(&payload);
|
||||
group.received_shards[packet_index] = Some(batch_data);
|
||||
group.received_original_count += 1;
|
||||
}
|
||||
group.last_update = Instant::now();
|
||||
|
||||
if group.is_done() {
|
||||
group.done();
|
||||
if inner.last_cleanup.elapsed() > Duration::from_secs(1) {
|
||||
Self::cleanup_old_groups(&mut inner.groups);
|
||||
inner.last_cleanup = Instant::now();
|
||||
}
|
||||
return Ok(packet.map(|v| vec![v]));
|
||||
}
|
||||
|
||||
let result = Self::try_decode(group, (src_ip, group_id), &net_packet)?;
|
||||
|
||||
if inner.last_cleanup.elapsed() > Duration::from_secs(1) {
|
||||
Self::cleanup_old_groups(&mut inner.groups);
|
||||
inner.last_cleanup = Instant::now();
|
||||
}
|
||||
match (packet, result) {
|
||||
(Some(packet), Some(mut result)) => {
|
||||
result.push(packet);
|
||||
Ok(Some(result))
|
||||
}
|
||||
(Some(packet), None) => Ok(Some(vec![packet])),
|
||||
(None, Some(result)) => Ok(Some(result)),
|
||||
(None, None) => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// 检查是否可以恢复丢失的包
|
||||
fn try_decode(
|
||||
group: &mut FecGroup,
|
||||
key: (Ipv4Addr, u64),
|
||||
net_packet: &NetPacket<TransmissionBytes>,
|
||||
) -> Result<Option<Vec<NetPacket<TransmissionBytes>>>> {
|
||||
if group.data_shards == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
if group.received_original_count == group.data_shards {
|
||||
return Ok(None);
|
||||
}
|
||||
let received_count = group.received_shards.iter().filter(|s| s.is_some()).count();
|
||||
|
||||
if received_count < group.data_shards {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Self::decode_with_rs(group, key, net_packet)
|
||||
}
|
||||
|
||||
/// Reed-Solomon解码恢复丢失的包
|
||||
fn decode_with_rs(
|
||||
group: &mut FecGroup,
|
||||
key: (Ipv4Addr, u64),
|
||||
net_packet: &NetPacket<TransmissionBytes>,
|
||||
) -> Result<Option<Vec<NetPacket<TransmissionBytes>>>> {
|
||||
let (src_ip, group_id) = key;
|
||||
|
||||
if group.received_shards.len() != group.data_shards + group.parity_shards {
|
||||
bail!(
|
||||
"received_shards.len()({}) != data_shards({})+parity_shards({}) src_ip={src_ip},group_id={group_id}",
|
||||
group.received_shards.len(),
|
||||
group.data_shards,
|
||||
group.parity_shards,
|
||||
)
|
||||
}
|
||||
|
||||
let mut delivered_packets = vec![false; group.data_shards];
|
||||
for (index, x) in group.received_shards[..group.data_shards]
|
||||
.iter()
|
||||
.enumerate()
|
||||
{
|
||||
if x.is_some() {
|
||||
delivered_packets[index] = true;
|
||||
}
|
||||
}
|
||||
|
||||
let rs = ReedSolomon::new(group.data_shards, group.parity_shards)?;
|
||||
rs.reconstruct(&mut group.received_shards)?;
|
||||
|
||||
let mut result = Vec::new();
|
||||
for (i, shard) in group
|
||||
.received_shards
|
||||
.iter()
|
||||
.enumerate()
|
||||
.take(group.data_shards)
|
||||
{
|
||||
if let Some(shard) = shard {
|
||||
if delivered_packets[i] {
|
||||
continue;
|
||||
}
|
||||
|
||||
let net_packet = Self::rebuild_net_packet(net_packet, shard)?;
|
||||
result.push(net_packet);
|
||||
}
|
||||
}
|
||||
|
||||
group.done();
|
||||
|
||||
Ok(Some(result))
|
||||
}
|
||||
|
||||
/// 从恢复的数据重建NetPacket
|
||||
fn rebuild_net_packet(
|
||||
current_packet: &NetPacket<TransmissionBytes>,
|
||||
recovered_data: &[u8],
|
||||
) -> Result<NetPacket<TransmissionBytes>> {
|
||||
if recovered_data.len() < 4 {
|
||||
bail!("recovered_data too short");
|
||||
}
|
||||
|
||||
let type_byte = recovered_data[0];
|
||||
let flags_byte = recovered_data[1];
|
||||
let payload_len = u16::from_be_bytes([recovered_data[2], recovered_data[3]]) as usize;
|
||||
|
||||
if payload_len + 4 > recovered_data.len() {
|
||||
bail!("invalid payload_len in recovered_data");
|
||||
}
|
||||
|
||||
let payload = &recovered_data[4..4 + payload_len];
|
||||
|
||||
let buffer = TransmissionBytes::zeroed(HEAD_LENGTH + payload.len());
|
||||
let mut net_packet = NetPacket::new(buffer)?;
|
||||
|
||||
net_packet.head_mut()[0] = type_byte;
|
||||
net_packet.head_mut()[2] = flags_byte;
|
||||
net_packet.set_src_id(current_packet.src_id());
|
||||
net_packet.set_dest_id(current_packet.dest_id());
|
||||
net_packet.set_ttl(current_packet.ttl());
|
||||
net_packet.set_fec_flag(false);
|
||||
net_packet.set_payload(payload)?;
|
||||
|
||||
Ok(net_packet)
|
||||
}
|
||||
|
||||
fn cleanup_old_groups(groups: &mut HashMap<(Ipv4Addr, u64), FecGroup>) {
|
||||
let now = Instant::now();
|
||||
|
||||
groups.retain(|_, group| now.duration_since(group.last_update) < GROUP_TIMEOUT);
|
||||
|
||||
if groups.len() > MAX_GROUPS {
|
||||
let mut group_ids: Vec<_> = groups.iter().map(|(id, g)| (*id, g.last_update)).collect();
|
||||
group_ids.sort_by_key(|(_, last_update)| *last_update);
|
||||
|
||||
let to_remove = group_ids.len() - MAX_GROUPS;
|
||||
for (group_id, _) in group_ids.iter().take(to_remove) {
|
||||
groups.remove(group_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,247 @@
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::outbound::BasicOutbound;
|
||||
use anyhow::{Result, bail};
|
||||
use parking_lot::Mutex;
|
||||
use prost::Message;
|
||||
use reed_solomon_erasure::galois_8::ReedSolomon;
|
||||
use std::collections::HashMap;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
mod fec_proto {
|
||||
include!(concat!(env!("OUT_DIR"), "/protocol.fec.rs"));
|
||||
}
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
pub use fec_proto::FecPacket;
|
||||
|
||||
const BATCH_SIZE: usize = 10;
|
||||
const REDUNDANCY_RATE: f32 = 0.2;
|
||||
const BATCH_TIMEOUT_MS: u64 = 20;
|
||||
const MIN_PARITY: usize = 1;
|
||||
const BATCH_CHANNEL_SIZE: usize = 1024;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct FecEncoder {
|
||||
batch_states: Arc<Mutex<HashMap<Ipv4Addr, DestBatchState>>>,
|
||||
batch_tx: mpsc::Sender<(Ipv4Addr, Ipv4Addr, u64, Vec<TransmissionBytes>)>,
|
||||
}
|
||||
|
||||
struct DestBatchState {
|
||||
group_id: u64,
|
||||
current_batch: Vec<TransmissionBytes>,
|
||||
deadline: Instant,
|
||||
src_ip: Ipv4Addr,
|
||||
}
|
||||
|
||||
impl FecEncoder {
|
||||
pub fn new(task_group: &TaskGroup, basic_outbound: BasicOutbound) -> Self {
|
||||
let (batch_tx, batch_rx) = mpsc::channel(BATCH_CHANNEL_SIZE);
|
||||
let batch_states = Arc::new(Mutex::new(HashMap::new()));
|
||||
let encoder = Self {
|
||||
batch_states: batch_states.clone(),
|
||||
batch_tx,
|
||||
};
|
||||
|
||||
task_group.spawn(fec_encoder_worker(batch_rx, basic_outbound, batch_states));
|
||||
|
||||
encoder
|
||||
}
|
||||
|
||||
/// 将数据包加入FEC批次并返回包装后的包
|
||||
pub fn encode(
|
||||
&self,
|
||||
mut packet: NetPacket<TransmissionBytes>,
|
||||
) -> Result<NetPacket<TransmissionBytes>> {
|
||||
let src_ip = Ipv4Addr::from(packet.src_id());
|
||||
let dest = Ipv4Addr::from(packet.dest_id());
|
||||
if packet.payload().len() > u16::MAX as usize {
|
||||
bail!("Payload too big");
|
||||
}
|
||||
let original_payload = packet.payload().to_vec();
|
||||
let original_payload_len = original_payload.len();
|
||||
|
||||
let type_byte = packet.head()[0];
|
||||
let flags_byte = packet.head()[2];
|
||||
|
||||
// 组装FEC数据: [type_byte, flags_byte, payload_len(u16), payload...]
|
||||
let batch_len = 4 + original_payload_len;
|
||||
let mut batch_buffer = TransmissionBytes::zeroed(batch_len);
|
||||
batch_buffer[0] = type_byte;
|
||||
batch_buffer[1] = flags_byte;
|
||||
batch_buffer[2..4].copy_from_slice(&(original_payload_len as u16).to_be_bytes());
|
||||
batch_buffer[4..batch_len].copy_from_slice(&original_payload);
|
||||
|
||||
let (group_id, packet_index) = {
|
||||
let mut states = self.batch_states.lock();
|
||||
let state = states.entry(dest).or_insert_with(|| DestBatchState {
|
||||
group_id: 0,
|
||||
current_batch: Vec::with_capacity(BATCH_SIZE),
|
||||
deadline: Instant::now() + Duration::from_millis(BATCH_TIMEOUT_MS),
|
||||
src_ip,
|
||||
});
|
||||
|
||||
let group_id = state.group_id;
|
||||
let packet_index = state.current_batch.len();
|
||||
|
||||
state.current_batch.push(batch_buffer);
|
||||
|
||||
if state.current_batch.len() >= BATCH_SIZE {
|
||||
let batch = std::mem::take(&mut state.current_batch);
|
||||
state.group_id += 1;
|
||||
state.deadline = Instant::now() + Duration::from_millis(BATCH_TIMEOUT_MS);
|
||||
|
||||
if self
|
||||
.batch_tx
|
||||
.try_send((src_ip, dest, group_id, batch))
|
||||
.is_err()
|
||||
{
|
||||
log::warn!(
|
||||
"failed to send batch to worker (channel full), dest={}, group_id={}",
|
||||
dest,
|
||||
group_id
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
(group_id, packet_index)
|
||||
};
|
||||
|
||||
let fec_packet = FecPacket {
|
||||
group_id,
|
||||
packet_index: packet_index as u32,
|
||||
payload: original_payload,
|
||||
parity_data: None,
|
||||
};
|
||||
|
||||
let fec_payload = fec_packet.encode_to_vec();
|
||||
packet
|
||||
.source_buf_mut()
|
||||
.resize(HEAD_LENGTH + fec_payload.len(), 0);
|
||||
packet.set_payload(&fec_payload)?;
|
||||
packet.set_fec_flag(true);
|
||||
|
||||
Ok(packet)
|
||||
}
|
||||
}
|
||||
|
||||
/// 后台worker,处理满批次和超时批次
|
||||
async fn fec_encoder_worker(
|
||||
mut batch_rx: mpsc::Receiver<(Ipv4Addr, Ipv4Addr, u64, Vec<TransmissionBytes>)>,
|
||||
basic_outbound: BasicOutbound,
|
||||
batch_states: Arc<Mutex<HashMap<Ipv4Addr, DestBatchState>>>,
|
||||
) {
|
||||
let mut timer = tokio::time::interval(Duration::from_millis(5));
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
Some((src,dest, group_id, mut items)) = batch_rx.recv() => {
|
||||
if let Err(e) = encode_and_send_parity(src,dest, group_id, &mut items, &basic_outbound).await {
|
||||
log::warn!("encode_and_send_parity error for {} group {}: {:?}", dest, group_id, e);
|
||||
}
|
||||
}
|
||||
|
||||
_ = timer.tick() => {
|
||||
let now = Instant::now();
|
||||
|
||||
let timeout_batches = {
|
||||
let mut states = batch_states.lock();
|
||||
let mut batches = Vec::new();
|
||||
for (dest, state) in states.iter_mut() {
|
||||
if !state.current_batch.is_empty() && now >= state.deadline {
|
||||
let items = std::mem::take(&mut state.current_batch);
|
||||
let group_id = state.group_id;
|
||||
state.group_id += 1;
|
||||
state.deadline = Instant::now() + Duration::from_millis(BATCH_TIMEOUT_MS);
|
||||
batches.push((state.src_ip,*dest, group_id, items));
|
||||
}
|
||||
}
|
||||
batches
|
||||
};
|
||||
|
||||
for (src,dest, group_id, mut items) in timeout_batches {
|
||||
if let Err(e) = encode_and_send_parity(src, dest, group_id, &mut items, &basic_outbound).await {
|
||||
log::warn!("encode_and_send_parity timeout error for {} group {}: {:?}", dest, group_id, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Reed-Solomon编码并发送冗余包
|
||||
async fn encode_and_send_parity(
|
||||
src: Ipv4Addr,
|
||||
dest: Ipv4Addr,
|
||||
group_id: u64,
|
||||
items: &mut Vec<TransmissionBytes>,
|
||||
basic_outbound: &BasicOutbound,
|
||||
) -> Result<()> {
|
||||
if items.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let data_shards = items.len();
|
||||
let parity_shards = (data_shards as f32 * REDUNDANCY_RATE).ceil() as usize;
|
||||
let parity_shards = parity_shards.max(MIN_PARITY);
|
||||
|
||||
let max_len = items.iter().map(|buf| buf.len()).max().unwrap_or(0);
|
||||
|
||||
if max_len == 0 {
|
||||
log::warn!("max_len is 0, dest={}, group_id={}", dest, group_id);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
for buf in items.iter_mut() {
|
||||
if buf.len() < max_len {
|
||||
let padding = max_len - buf.len();
|
||||
buf.extend_end(padding);
|
||||
}
|
||||
}
|
||||
|
||||
for _ in 0..parity_shards {
|
||||
items.push(TransmissionBytes::zeroed(max_len));
|
||||
}
|
||||
|
||||
let rs = ReedSolomon::new(data_shards, parity_shards)?;
|
||||
let mut shard_refs: Vec<&mut [u8]> = items.iter_mut().map(|buf| buf.as_mut()).collect();
|
||||
rs.encode(&mut shard_refs)?;
|
||||
|
||||
for (i, parity_buf) in items[data_shards..].iter().enumerate() {
|
||||
let packet_index = (data_shards + i) as u32;
|
||||
let fec_packet = FecPacket {
|
||||
group_id,
|
||||
packet_index,
|
||||
payload: parity_buf.as_ref().to_vec(),
|
||||
parity_data: Some(fec_proto::ParityData {
|
||||
data_shards: data_shards as u32,
|
||||
parity_shards: parity_shards as u32,
|
||||
}),
|
||||
};
|
||||
|
||||
let fec_payload = fec_packet.encode_to_vec();
|
||||
|
||||
let buffer = TransmissionBytes::zeroed(HEAD_LENGTH + fec_payload.len());
|
||||
let mut net_packet = NetPacket::new(buffer)?;
|
||||
net_packet.set_msg_type(MsgType::Turn);
|
||||
net_packet.set_src_id(src.into());
|
||||
net_packet.set_dest_id(dest.into());
|
||||
net_packet.set_ttl(5);
|
||||
net_packet.set_payload(&fec_payload)?;
|
||||
net_packet.set_fec_flag(true);
|
||||
|
||||
if let Err(e) = basic_outbound.send_encrypted_packet(dest, net_packet).await {
|
||||
log::warn!(
|
||||
"failed to send parity packet {}: {:?}, dest={}, group_id={}",
|
||||
packet_index,
|
||||
e,
|
||||
dest,
|
||||
group_id
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
mod decoder;
|
||||
mod encoder;
|
||||
|
||||
pub(crate) use decoder::FecDecoder;
|
||||
pub(crate) use encoder::FecEncoder;
|
||||
@@ -0,0 +1,15 @@
|
||||
pub(crate) mod compression;
|
||||
pub mod context;
|
||||
pub mod core;
|
||||
pub mod crypto;
|
||||
pub(crate) mod fec;
|
||||
pub mod nat;
|
||||
pub mod protocol;
|
||||
pub mod tls;
|
||||
pub(crate) mod tun;
|
||||
pub mod tunnel_core;
|
||||
pub mod utils;
|
||||
|
||||
pub mod api;
|
||||
pub(crate) mod enhanced_tunnel;
|
||||
pub mod port_mapping;
|
||||
@@ -0,0 +1,139 @@
|
||||
use crate::context::SharedNetworkAddr;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::Context;
|
||||
use pnet_packet::Packet;
|
||||
use pnet_packet::icmp::echo_reply::{Identifier, SequenceNumber};
|
||||
use pnet_packet::icmp::{IcmpPacket, IcmpTypes};
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
use std::collections::HashMap;
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
|
||||
use tcp_ip::IpStack;
|
||||
use tcp_ip::icmp::IcmpSocket;
|
||||
use tokio::net::UdpSocket;
|
||||
|
||||
pub async fn start_icmp_nat(
|
||||
task_group: &TaskGroup,
|
||||
ip_stack: &IpStack,
|
||||
no_tun: bool,
|
||||
network: SharedNetworkAddr,
|
||||
) -> anyhow::Result<()> {
|
||||
let net_icmp_socket = socket2::Socket::new(
|
||||
socket2::Domain::IPV4,
|
||||
socket2::Type::RAW,
|
||||
Some(socket2::Protocol::ICMPV4),
|
||||
)
|
||||
.context("new Socket RAW ICMPV4 failed")?;
|
||||
let addr: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
|
||||
net_icmp_socket
|
||||
.bind(&socket2::SockAddr::from(addr))
|
||||
.context("bind Socket ICMPV4 failed")?;
|
||||
net_icmp_socket.set_nonblocking(true)?;
|
||||
|
||||
let std_socket: std::net::UdpSocket = net_icmp_socket.into();
|
||||
|
||||
let tokio_icmp_socket = UdpSocket::from_std(std_socket)?;
|
||||
|
||||
let inner_icmp_socket = IcmpSocket::bind_all(ip_stack.clone()).await?;
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = task(tokio_icmp_socket, inner_icmp_socket, no_tun, network).await {
|
||||
log::error!("icmp task failed: {:?}", e);
|
||||
}
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
async fn task(
|
||||
tokio_icmp_socket: UdpSocket,
|
||||
inner_icmp_socket: IcmpSocket,
|
||||
no_tun: bool,
|
||||
network: SharedNetworkAddr,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut buf1 = vec![0u8; 65536];
|
||||
let mut buf2 = vec![0u8; 65536];
|
||||
let mut map = HashMap::new();
|
||||
loop {
|
||||
tokio::select! {
|
||||
rs = tokio_icmp_socket.recv(&mut buf1) => {
|
||||
let len = rs?;
|
||||
tokio_icmp_socket_recv(&buf1[..len],&inner_icmp_socket,&map,no_tun,&network).await?;
|
||||
}
|
||||
rs = inner_icmp_socket.recv_from_to(&mut buf2) => {
|
||||
let (len,src,dst) = rs?;
|
||||
inner_icmp_socket_recv(&buf2[..len],src,dst,&tokio_icmp_socket,&mut map,no_tun,&network).await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
async fn tokio_icmp_socket_recv(
|
||||
buf: &[u8],
|
||||
inner_icmp_socket: &IcmpSocket,
|
||||
map: &HashMap<(Ipv4Addr, Identifier, SequenceNumber), Ipv4Addr>,
|
||||
no_tun: bool,
|
||||
network: &SharedNetworkAddr,
|
||||
) -> anyhow::Result<()> {
|
||||
let Some(ipv4) = Ipv4Packet::new(buf) else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(icmp) = IcmpPacket::new(ipv4.payload()) else {
|
||||
return Ok(());
|
||||
};
|
||||
if icmp.get_icmp_type() != IcmpTypes::EchoReply
|
||||
&& icmp.get_icmp_type() != IcmpTypes::EchoRequest
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
let payload = icmp.payload();
|
||||
if payload.len() < 4 {
|
||||
return Ok(());
|
||||
}
|
||||
let mut src = ipv4.get_source();
|
||||
let identifier = Identifier::new(u16::from_be_bytes([payload[0], payload[1]]));
|
||||
let sequence_number = SequenceNumber::new(u16::from_be_bytes([payload[2], payload[3]]));
|
||||
let Some(dst) = map.get(&(src, identifier, sequence_number)) else {
|
||||
return Ok(());
|
||||
};
|
||||
if no_tun && src == Ipv4Addr::LOCALHOST {
|
||||
src = network.ip().context("not ip")?;
|
||||
}
|
||||
|
||||
inner_icmp_socket
|
||||
.send_from_to(ipv4.payload(), src.into(), (*dst).into())
|
||||
.await
|
||||
.context("sending ICMPv4 failed")?;
|
||||
Ok(())
|
||||
}
|
||||
async fn inner_icmp_socket_recv(
|
||||
buf: &[u8],
|
||||
src: IpAddr,
|
||||
dst: IpAddr,
|
||||
tokio_icmp_socket: &UdpSocket,
|
||||
map: &mut HashMap<(Ipv4Addr, Identifier, SequenceNumber), Ipv4Addr>,
|
||||
no_tun: bool,
|
||||
network: &SharedNetworkAddr,
|
||||
) -> anyhow::Result<()> {
|
||||
let (IpAddr::V4(src), IpAddr::V4(mut dst)) = (src, dst) else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(icmp) = IcmpPacket::new(buf) else {
|
||||
return Ok(());
|
||||
};
|
||||
if icmp.get_icmp_type() != IcmpTypes::EchoReply
|
||||
&& icmp.get_icmp_type() != IcmpTypes::EchoRequest
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
let payload = icmp.payload();
|
||||
if payload.len() < 4 {
|
||||
return Ok(());
|
||||
}
|
||||
if no_tun && dst == network.ip().context("not ip")? {
|
||||
dst = Ipv4Addr::LOCALHOST;
|
||||
}
|
||||
|
||||
let identifier = Identifier::new(u16::from_be_bytes([payload[0], payload[1]]));
|
||||
let sequence_number = SequenceNumber::new(u16::from_be_bytes([payload[2], payload[3]]));
|
||||
map.insert((dst, identifier, sequence_number), src);
|
||||
tokio_icmp_socket
|
||||
.send_to(buf, SocketAddr::new(dst.into(), 0))
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,225 @@
|
||||
use crate::context::{NetworkAddr, SharedNetworkAddr};
|
||||
use crate::nat::AllowSubnetExternalRoute;
|
||||
use crate::protocol::ip_packet_protocol::HEAD_LENGTH;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::outbound::HybridOutbound;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::Context;
|
||||
use bytes::BytesMut;
|
||||
use pnet_packet::ip::IpNextHeaderProtocol;
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
use std::net::{Ipv4Addr, SocketAddr};
|
||||
use std::str::FromStr;
|
||||
use std::sync::Arc;
|
||||
use tcp_ip::{IpStackConfig, IpStackRecv, IpStackSend};
|
||||
use tokio::io::{AsyncRead, AsyncWrite};
|
||||
|
||||
#[cfg(not(target_os = "android"))]
|
||||
mod icmp_nat;
|
||||
mod tcp_nat;
|
||||
mod udp_nat;
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct InternalNatInbound {
|
||||
no_tun: bool,
|
||||
ip_stack_send: Arc<IpStackSend>,
|
||||
allow_subnet: AllowSubnetExternalRoute,
|
||||
network: SharedNetworkAddr,
|
||||
}
|
||||
impl InternalNatInbound {
|
||||
pub async fn create(
|
||||
task_group: &TaskGroup,
|
||||
mtu: u16,
|
||||
hybrid_outbound: HybridOutbound,
|
||||
allow_subnet: AllowSubnetExternalRoute,
|
||||
network: SharedNetworkAddr,
|
||||
no_tun: bool,
|
||||
) -> anyhow::Result<Self> {
|
||||
let ip_stack_config = IpStackConfig {
|
||||
mtu,
|
||||
..Default::default()
|
||||
};
|
||||
let (ip_stack, ip_stack_send, ip_stack_recv) = tcp_ip::ip_stack(ip_stack_config)?;
|
||||
#[cfg(not(target_os = "android"))]
|
||||
icmp_nat::start_icmp_nat(task_group, &ip_stack, no_tun, network.clone()).await?;
|
||||
tcp_nat::start_tcp_nat(task_group, &ip_stack, no_tun, network.clone()).await?;
|
||||
udp_nat::start_udp_nat(task_group, &ip_stack).await?;
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = ip_stack_recv_task(ip_stack_recv, hybrid_outbound).await {
|
||||
log::error!("ip stack recv task error: {e:?}");
|
||||
}
|
||||
});
|
||||
Ok(Self {
|
||||
no_tun,
|
||||
ip_stack_send: Arc::new(ip_stack_send),
|
||||
allow_subnet,
|
||||
network,
|
||||
})
|
||||
}
|
||||
pub async fn send(&self, data: &[u8], net: &NetworkAddr) -> anyhow::Result<()> {
|
||||
if data[0] >> 4 != 4 {
|
||||
return Ok(());
|
||||
}
|
||||
let Some(ipv4) = Ipv4Packet::new(data) else {
|
||||
return Ok(());
|
||||
};
|
||||
let dest = ipv4.get_destination();
|
||||
if net.network().contains(&dest)
|
||||
|| dest == net.broadcast
|
||||
|| dest.is_broadcast()
|
||||
|| dest.is_multicast()
|
||||
|| self.allow_subnet.allow(&dest)
|
||||
{
|
||||
self.ip_stack_send.send_ip_packet(data).await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
pub async fn send_ipv4_payload(
|
||||
&self,
|
||||
protocol: IpNextHeaderProtocol,
|
||||
src_ip: Ipv4Addr,
|
||||
dest_ip: Ipv4Addr,
|
||||
payload: BytesMut,
|
||||
) -> anyhow::Result<()> {
|
||||
self.ip_stack_send
|
||||
.send_ipv4_payload(protocol, src_ip, dest_ip, payload)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
async fn ip_stack_recv_task(
|
||||
mut ip_stack_recv: IpStackRecv,
|
||||
hybrid_outbound: HybridOutbound,
|
||||
) -> anyhow::Result<()> {
|
||||
loop {
|
||||
let mut bytes = TransmissionBytes::new_offset_zeroed(HEAD_LENGTH);
|
||||
let len = ip_stack_recv.recv(&mut bytes).await?;
|
||||
bytes.set_len(len)?;
|
||||
if let Err(e) = hybrid_outbound.ipv4_outbound_common(bytes).await {
|
||||
log::warn!("ip_stack_recv_task,{e:?}");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl InternalNatInbound {
|
||||
fn network_contains(&self, ip: &Ipv4Addr) -> bool {
|
||||
self.network
|
||||
.network()
|
||||
.map(|net| net.contains(ip))
|
||||
.unwrap_or(false)
|
||||
}
|
||||
pub fn use_nat(&self, dst: &Ipv4Addr) -> bool {
|
||||
if self.no_tun {
|
||||
return self.allow_nat(dst);
|
||||
}
|
||||
if self.network_contains(dst) {
|
||||
return false;
|
||||
}
|
||||
self.allow_subnet.allow(dst)
|
||||
}
|
||||
pub fn no_tun(&self) -> bool {
|
||||
self.no_tun
|
||||
}
|
||||
pub fn allow_nat(&self, dst: &Ipv4Addr) -> bool {
|
||||
self.allow_subnet.allow(dst) || self.network_contains(dst)
|
||||
}
|
||||
pub async fn tcp_nat<R, W>(
|
||||
&self,
|
||||
recv_stream: R,
|
||||
send_stream: W,
|
||||
mut dest_ip: Ipv4Addr,
|
||||
dest_port: u16,
|
||||
) -> anyhow::Result<()>
|
||||
where
|
||||
R: AsyncRead + Unpin,
|
||||
W: AsyncWrite + Unpin,
|
||||
{
|
||||
if self.no_tun {
|
||||
let net = self.network.get().context("no network")?;
|
||||
if dest_ip == net.ip {
|
||||
dest_ip = Ipv4Addr::LOCALHOST;
|
||||
} else if net.network().contains(&dest_ip) {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
let dst = SocketAddr::new(dest_ip.into(), dest_port);
|
||||
tcp_nat::stream_nat(recv_stream, send_stream, dst).await
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct PortMappingManager {
|
||||
no_tun: bool,
|
||||
allow_port_mapping: bool,
|
||||
network: SharedNetworkAddr,
|
||||
}
|
||||
|
||||
impl PortMappingManager {
|
||||
pub fn new(no_tun: bool, allow_port_mapping: bool, network: SharedNetworkAddr) -> Self {
|
||||
Self {
|
||||
no_tun,
|
||||
allow_port_mapping,
|
||||
network,
|
||||
}
|
||||
}
|
||||
pub async fn tcp_mapping<R, W>(
|
||||
&self,
|
||||
recv_stream: R,
|
||||
send_stream: W,
|
||||
dest: String,
|
||||
dest_port: u16,
|
||||
) -> anyhow::Result<()>
|
||||
where
|
||||
R: AsyncRead + Unpin,
|
||||
W: AsyncWrite + Unpin,
|
||||
{
|
||||
if !self.allow_port_mapping {
|
||||
log::debug!("port mapping not enabled");
|
||||
return Ok(());
|
||||
}
|
||||
if self.no_tun
|
||||
&& let Ok(dest_ip) = Ipv4Addr::from_str(&dest)
|
||||
{
|
||||
let net = self.network.get().context("no network")?;
|
||||
|
||||
if dest_ip == net.ip {
|
||||
let dst = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), dest_port);
|
||||
return tcp_nat::stream_nat(recv_stream, send_stream, dst).await;
|
||||
} else if net.network().contains(&dest_ip) {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
let dst = format!("{}:{}", dest, dest_port);
|
||||
tcp_nat::stream_nat(recv_stream, send_stream, dst).await
|
||||
}
|
||||
pub async fn udp_mapping<R, W>(
|
||||
&self,
|
||||
recv_stream: R,
|
||||
send_stream: W,
|
||||
dest: String,
|
||||
dest_port: u16,
|
||||
) -> anyhow::Result<()>
|
||||
where
|
||||
R: AsyncRead + Unpin,
|
||||
W: AsyncWrite + Unpin,
|
||||
{
|
||||
if !self.allow_port_mapping {
|
||||
log::debug!("port mapping not enabled");
|
||||
return Ok(());
|
||||
}
|
||||
if self.no_tun
|
||||
&& let Ok(dest_ip) = Ipv4Addr::from_str(&dest)
|
||||
{
|
||||
let net = self.network.get().context("no network")?;
|
||||
|
||||
if dest_ip == net.ip {
|
||||
let dst = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), dest_port);
|
||||
return udp_nat::stream_nat(recv_stream, send_stream, dst).await;
|
||||
} else if net.network().contains(&dest_ip) {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
let dst = format!("{}:{}", dest, dest_port);
|
||||
udp_nat::stream_nat(recv_stream, send_stream, dst).await
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
use crate::context::SharedNetworkAddr;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::Context;
|
||||
use std::fmt::Debug;
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
use tcp_ip::IpStack;
|
||||
use tcp_ip::tcp::TcpListener;
|
||||
use tokio::io::{AsyncRead, AsyncWrite};
|
||||
use tokio::net::{TcpStream, ToSocketAddrs};
|
||||
|
||||
pub async fn start_tcp_nat(
|
||||
task_group: &TaskGroup,
|
||||
ip_stack: &IpStack,
|
||||
no_tun: bool,
|
||||
network: SharedNetworkAddr,
|
||||
) -> anyhow::Result<()> {
|
||||
let tcp_listener = TcpListener::bind_all(ip_stack.clone()).await?;
|
||||
let group = task_group.clone();
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = listen_task(&group, tcp_listener, no_tun, network).await {
|
||||
log::error!("listen task error: {:?}", e);
|
||||
}
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn listen_task(
|
||||
task_group: &TaskGroup,
|
||||
mut tcp_listener: TcpListener,
|
||||
no_tun: bool,
|
||||
network: SharedNetworkAddr,
|
||||
) -> anyhow::Result<()> {
|
||||
loop {
|
||||
let (stream, _addr) = tcp_listener.accept().await?;
|
||||
let mut local_addr = stream.local_addr()?;
|
||||
let peer_addr = stream.peer_addr()?;
|
||||
if no_tun {
|
||||
let IpAddr::V4(ip) = local_addr.ip() else {
|
||||
continue;
|
||||
};
|
||||
if ip == network.ip().context("not ip")? {
|
||||
// 无tun的情况下写入本机的则写到localhost
|
||||
local_addr.set_ip(IpAddr::V4(Ipv4Addr::LOCALHOST));
|
||||
}
|
||||
}
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = stream_task(stream, local_addr).await {
|
||||
log::error!("stream task Error: {:?},{peer_addr}->{local_addr}", e);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async fn stream_task(
|
||||
mut inner_stream: tcp_ip::tcp::TcpStream,
|
||||
addr: SocketAddr,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut tokio_stream = TcpStream::connect(addr).await?;
|
||||
tokio::io::copy_bidirectional(&mut inner_stream, &mut tokio_stream).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) async fn stream_nat<R, W, A: ToSocketAddrs + Debug>(
|
||||
mut recv_stream: R,
|
||||
mut send_stream: W,
|
||||
addr: A,
|
||||
) -> anyhow::Result<()>
|
||||
where
|
||||
R: AsyncRead + Unpin,
|
||||
W: AsyncWrite + Unpin,
|
||||
{
|
||||
let mut tokio_stream = TcpStream::connect(&addr)
|
||||
.await
|
||||
.with_context(|| format!("error connecting to {:?}", addr))?;
|
||||
let (mut tcp_r, mut tcp_w) = tokio_stream.split();
|
||||
tokio::select! {
|
||||
_ = tokio::io::copy(&mut recv_stream, &mut tcp_w) => {},
|
||||
_ = tokio::io::copy(&mut tcp_r, &mut send_stream) => {},
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::Context;
|
||||
use bytes::Bytes;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use std::collections::HashMap;
|
||||
use std::fmt::Debug;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tcp_ip::IpStack;
|
||||
use tokio::io::{AsyncRead, AsyncWrite};
|
||||
use tokio::net::ToSocketAddrs;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_util::codec::{FramedRead, FramedWrite, LengthDelimitedCodec};
|
||||
|
||||
struct NatEntry {
|
||||
socket: Arc<tokio::net::UdpSocket>,
|
||||
last_active: Instant,
|
||||
}
|
||||
|
||||
type NatTable = Arc<Mutex<HashMap<(SocketAddr, SocketAddr), NatEntry>>>;
|
||||
|
||||
const NAT_IDLE_TIMEOUT: Duration = Duration::from_secs(60 * 5);
|
||||
const NAT_GC_INTERVAL: Duration = Duration::from_secs(60);
|
||||
|
||||
pub async fn start_udp_nat(task_group: &TaskGroup, ip_stack: &IpStack) -> anyhow::Result<()> {
|
||||
let inner_socket = tcp_ip::udp::UdpSocket::bind_all(ip_stack.clone()).await?;
|
||||
let inner_socket = Arc::new(inner_socket);
|
||||
let nat_table: NatTable = Arc::new(Mutex::new(HashMap::new()));
|
||||
let mut buf = vec![0u8; 65536];
|
||||
let group = task_group.clone();
|
||||
let nat_table_clone = nat_table.clone();
|
||||
task_group.spawn(async move {
|
||||
loop {
|
||||
let (len, src, dst) = match inner_socket.recv_from_to(&mut buf).await {
|
||||
Ok(rs) => rs,
|
||||
Err(e) => {
|
||||
log::warn!("{e:?}");
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(e) =
|
||||
handle_outbound(&group, &inner_socket, &nat_table, src, dst, &buf[..len]).await
|
||||
{
|
||||
log::warn!("udp nat outbound error: {e:?}");
|
||||
}
|
||||
}
|
||||
});
|
||||
spawn_nat_gc(task_group, nat_table_clone);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn handle_outbound(
|
||||
task_group: &TaskGroup,
|
||||
inner: &Arc<tcp_ip::udp::UdpSocket>,
|
||||
nat: &NatTable,
|
||||
src: SocketAddr,
|
||||
dst: SocketAddr,
|
||||
packet: &[u8],
|
||||
) -> anyhow::Result<()> {
|
||||
let key = (src, dst);
|
||||
|
||||
let socket = {
|
||||
let mut table = nat.lock().await;
|
||||
if let Some(entry) = table.get_mut(&key) {
|
||||
entry.last_active = Instant::now();
|
||||
entry.socket.clone()
|
||||
} else {
|
||||
// 创建真实 UDP socket
|
||||
let sock = tokio::net::UdpSocket::bind("0.0.0.0:0").await?;
|
||||
sock.connect(dst).await?;
|
||||
let sock = Arc::new(sock);
|
||||
table.insert(
|
||||
key,
|
||||
NatEntry {
|
||||
socket: sock.clone(),
|
||||
last_active: Instant::now(),
|
||||
},
|
||||
);
|
||||
|
||||
// 启动反向转发
|
||||
spawn_inbound(
|
||||
task_group,
|
||||
inner.clone(),
|
||||
nat.clone(),
|
||||
src,
|
||||
dst,
|
||||
sock.clone(),
|
||||
);
|
||||
|
||||
sock
|
||||
}
|
||||
};
|
||||
|
||||
socket.send(packet).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn spawn_inbound(
|
||||
task_group: &TaskGroup,
|
||||
inner: Arc<tcp_ip::udp::UdpSocket>,
|
||||
nat: NatTable,
|
||||
src: SocketAddr,
|
||||
dst: SocketAddr,
|
||||
socket: Arc<tokio::net::UdpSocket>,
|
||||
) {
|
||||
task_group.spawn(async move {
|
||||
let mut buf = vec![0u8; 65536];
|
||||
|
||||
loop {
|
||||
let len = match socket.recv(&mut buf).await {
|
||||
Ok(n) => n,
|
||||
Err(_) => break,
|
||||
};
|
||||
|
||||
// 反向写回 inner socket
|
||||
if inner.send_from_to(&buf[..len], dst, src).await.is_err() {
|
||||
break;
|
||||
}
|
||||
|
||||
// 更新活跃时间
|
||||
if let Some(entry) = nat.lock().await.get_mut(&(src, dst)) {
|
||||
entry.last_active = Instant::now();
|
||||
}
|
||||
}
|
||||
|
||||
// 回收 NAT
|
||||
nat.lock().await.remove(&(src, dst));
|
||||
});
|
||||
}
|
||||
|
||||
fn spawn_nat_gc(task_group: &TaskGroup, nat: NatTable) {
|
||||
task_group.spawn(async move {
|
||||
let mut interval = tokio::time::interval(NAT_GC_INTERVAL);
|
||||
|
||||
loop {
|
||||
interval.tick().await;
|
||||
|
||||
let now = Instant::now();
|
||||
let mut table = nat.lock().await;
|
||||
|
||||
table.retain(|(src, dst), entry| {
|
||||
let alive = now.duration_since(entry.last_active) < NAT_IDLE_TIMEOUT;
|
||||
if !alive {
|
||||
log::debug!("udp nat expired: {} -> {}", src, dst);
|
||||
}
|
||||
alive
|
||||
});
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
pub(crate) async fn stream_nat<R, W, A: ToSocketAddrs + Debug>(
|
||||
recv_stream: R,
|
||||
send_stream: W,
|
||||
addr: A,
|
||||
) -> anyhow::Result<()>
|
||||
where
|
||||
R: AsyncRead + Unpin,
|
||||
W: AsyncWrite + Unpin,
|
||||
{
|
||||
let udp_socket = tokio::net::UdpSocket::bind("0.0.0.0:0").await?;
|
||||
udp_socket
|
||||
.connect(&addr)
|
||||
.await
|
||||
.with_context(|| format!("error connecting to {:?}", addr))?;
|
||||
let mut framed_read = FramedRead::new(recv_stream, LengthDelimitedCodec::new());
|
||||
let mut framed_write = FramedWrite::new(send_stream, LengthDelimitedCodec::new());
|
||||
let mut buf = vec![0u8; 65536];
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
Some(buf) = framed_read.next() => {
|
||||
let buf = buf?;
|
||||
udp_socket.send(&buf).await?;
|
||||
},
|
||||
rs = udp_socket.recv(&mut buf) =>{
|
||||
let len = rs?;
|
||||
framed_write.send(Bytes::copy_from_slice(&buf[..len])).await?;
|
||||
},
|
||||
else => {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
use ipnet::Ipv4Net;
|
||||
use parking_lot::Mutex;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
use std::fmt;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::str::FromStr;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub(crate) mod internal_nat;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct NetInput {
|
||||
pub net: Ipv4Net,
|
||||
pub target_ip: Ipv4Addr,
|
||||
}
|
||||
impl FromStr for NetInput {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
let parts: Vec<&str> = s.split(',').map(|x| x.trim()).collect();
|
||||
if parts.len() != 2 {
|
||||
return Err("格式错误,应为 net,target_ip 例如: 192.168.0.0/24,10.26.0.2".into());
|
||||
}
|
||||
|
||||
let net = Ipv4Net::from_str(parts[0]).map_err(|e| format!("网络段格式错误: {}", e))?;
|
||||
|
||||
let target_ip =
|
||||
Ipv4Addr::from_str(parts[1]).map_err(|e| format!("目标 IP 格式错误: {}", e))?;
|
||||
|
||||
Ok(NetInput { net, target_ip })
|
||||
}
|
||||
}
|
||||
impl fmt::Display for NetInput {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "{},{}", self.net, self.target_ip)
|
||||
}
|
||||
}
|
||||
impl Serialize for NetInput {
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
serializer.serialize_str(&self.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for NetInput {
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let s = String::deserialize(deserializer)?;
|
||||
s.parse().map_err(serde::de::Error::custom)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct SubnetExternalRoute {
|
||||
route_table: Arc<Mutex<Vec<NetInput>>>,
|
||||
}
|
||||
impl SubnetExternalRoute {
|
||||
pub fn new(mut route_table: Vec<NetInput>) -> Self {
|
||||
route_table.sort_by_key(|r| std::cmp::Reverse(r.net.prefix_len()));
|
||||
SubnetExternalRoute {
|
||||
route_table: Arc::new(Mutex::new(route_table)),
|
||||
}
|
||||
}
|
||||
pub fn set_route_table(&self, mut route_table: Vec<NetInput>) {
|
||||
route_table.sort_by_key(|r| std::cmp::Reverse(r.net.prefix_len()));
|
||||
*self.route_table.lock() = route_table;
|
||||
}
|
||||
pub fn route(&self, ip: &Ipv4Addr) -> Option<Ipv4Addr> {
|
||||
let route_table = self.route_table.lock();
|
||||
if route_table.is_empty() {
|
||||
return None;
|
||||
}
|
||||
for net in route_table.iter() {
|
||||
if net.net.contains(ip) {
|
||||
return Some(net.target_ip);
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
pub fn all_route(&self) -> Vec<NetInput> {
|
||||
self.route_table.lock().clone()
|
||||
}
|
||||
pub fn reset_route(&self, route_table: Vec<NetInput>) {
|
||||
*self.route_table.lock() = route_table;
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct AllowSubnetExternalRoute {
|
||||
route_table: Arc<Vec<Ipv4Net>>,
|
||||
}
|
||||
impl AllowSubnetExternalRoute {
|
||||
pub fn new(mut route_table: Vec<Ipv4Net>) -> Self {
|
||||
route_table.sort_by_key(|r| r.prefix_len());
|
||||
Self {
|
||||
route_table: Arc::new(route_table),
|
||||
}
|
||||
}
|
||||
pub fn allow(&self, ip: &Ipv4Addr) -> bool {
|
||||
if self.route_table.is_empty() {
|
||||
return false;
|
||||
}
|
||||
for net in self.route_table.iter() {
|
||||
if net.contains(ip) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
use crate::enhanced_tunnel::quic_over::quic_client::QuicTunnelClient;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use pnet_packet::ip::{IpNextHeaderProtocol, IpNextHeaderProtocols};
|
||||
use std::fmt;
|
||||
use std::net::{Ipv4Addr, SocketAddr};
|
||||
use std::str::FromStr;
|
||||
|
||||
pub(crate) mod tcp_port_mapping;
|
||||
pub(crate) mod udp_port_mapping;
|
||||
|
||||
pub(crate) async fn port_mapping_start(
|
||||
task_group: &TaskGroup,
|
||||
list: Vec<PortMapping>,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
tcp_port_mapping::start(task_group, &list, quic_tunnel_client.clone()).await?;
|
||||
udp_port_mapping::start(task_group, &list, quic_tunnel_client).await?;
|
||||
Ok(())
|
||||
}
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PortMapping {
|
||||
pub protocol: IpNextHeaderProtocol,
|
||||
pub src_addr: SocketAddr,
|
||||
pub virtual_target_ip: Ipv4Addr,
|
||||
pub dst_host: String,
|
||||
pub dst_port: u16,
|
||||
}
|
||||
|
||||
impl fmt::Display for PortMapping {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(
|
||||
f,
|
||||
"{}://{}-{}-{}:{}",
|
||||
protocol_to_str(self.protocol),
|
||||
self.src_addr,
|
||||
self.virtual_target_ip,
|
||||
self.dst_host,
|
||||
self.dst_port
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl FromStr for PortMapping {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
let (proto_str, rest) = s.split_once("://").ok_or("missing '://'")?;
|
||||
|
||||
let protocol =
|
||||
str_to_protocol(proto_str).ok_or_else(|| format!("unknown protocol: {}", proto_str))?;
|
||||
|
||||
let mut parts = rest.splitn(3, '-');
|
||||
|
||||
let src_addr = parts
|
||||
.next()
|
||||
.ok_or("missing src_addr")?
|
||||
.parse::<SocketAddr>()
|
||||
.map_err(|e| format!("invalid src_addr: {}", e))?;
|
||||
|
||||
let virtual_target_ip = parts
|
||||
.next()
|
||||
.ok_or("missing virtual_target_ip")?
|
||||
.parse::<Ipv4Addr>()
|
||||
.map_err(|e| format!("invalid virtual_target_ip: {}", e))?;
|
||||
|
||||
let dst = parts.next().ok_or("missing destination")?;
|
||||
|
||||
let (dst_host, dst_port) = dst.rsplit_once(':').ok_or("missing dst port")?;
|
||||
|
||||
let dst_port = dst_port
|
||||
.parse::<u16>()
|
||||
.map_err(|e| format!("invalid dst_port: {}", e))?;
|
||||
if dst_port == 0 {
|
||||
return Err("invalid dst port: 0".to_string());
|
||||
}
|
||||
Ok(Self {
|
||||
protocol,
|
||||
src_addr,
|
||||
virtual_target_ip,
|
||||
dst_host: dst_host.to_string(),
|
||||
dst_port,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn protocol_to_str(p: IpNextHeaderProtocol) -> &'static str {
|
||||
match p {
|
||||
IpNextHeaderProtocols::Tcp => "tcp",
|
||||
IpNextHeaderProtocols::Udp => "udp",
|
||||
_ => "unknown",
|
||||
}
|
||||
}
|
||||
|
||||
fn str_to_protocol(s: &str) -> Option<IpNextHeaderProtocol> {
|
||||
match s.to_ascii_lowercase().as_str() {
|
||||
"tcp" => Some(IpNextHeaderProtocols::Tcp),
|
||||
"udp" => Some(IpNextHeaderProtocols::Udp),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
use crate::enhanced_tunnel::quic_over::quic_client::{QuicTunnelClient, send_handshake};
|
||||
use crate::port_mapping::PortMapping;
|
||||
use crate::protocol::client_message::{
|
||||
PortProxyHandshake, QuicProxyHandshake, quic_proxy_handshake,
|
||||
};
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::Context;
|
||||
use pnet_packet::ip::IpNextHeaderProtocols;
|
||||
use std::net::{Ipv4Addr, SocketAddr};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
pub async fn start(
|
||||
task_group: &TaskGroup,
|
||||
list: &Vec<PortMapping>,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
for x in list {
|
||||
if x.protocol != IpNextHeaderProtocols::Tcp {
|
||||
continue;
|
||||
}
|
||||
log::info!("Starting TCP port mapping on {}", x);
|
||||
let listener = TcpListener::bind(x.src_addr)
|
||||
.await
|
||||
.with_context(|| format!("Tcp port mapping Failed to bind to {}", x.src_addr))?;
|
||||
let group = task_group.clone();
|
||||
let tunnel_client = quic_tunnel_client.clone();
|
||||
let mapping = x.clone();
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = listen(&group, listener, &mapping, tunnel_client).await {
|
||||
log::error!("listen {:?},mapping:{mapping}", e);
|
||||
}
|
||||
});
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn listen(
|
||||
task_group: &TaskGroup,
|
||||
listener: TcpListener,
|
||||
mapping: &PortMapping,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
loop {
|
||||
let (stream, addr) = listener.accept().await?;
|
||||
let tunnel_client = quic_tunnel_client.clone();
|
||||
let target_ip = mapping.virtual_target_ip;
|
||||
let dst_host = mapping.dst_host.clone();
|
||||
let dst_port = mapping.dst_port;
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) =
|
||||
stream_copy(stream, addr, target_ip, dst_host, dst_port, tunnel_client).await
|
||||
{
|
||||
log::error!("TCP TCP Stream Error: {:?}", e);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async fn stream_copy(
|
||||
mut tcp_stream: TcpStream,
|
||||
src: SocketAddr,
|
||||
target_ip: Ipv4Addr,
|
||||
dst_host: String,
|
||||
dst_port: u16,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
let (mut send_stream, mut recv_stream) = quic_tunnel_client.open_bi(target_ip).await?;
|
||||
let handshake = QuicProxyHandshake {
|
||||
handshake: Some(quic_proxy_handshake::Handshake::TcpPortMapping(
|
||||
PortProxyHandshake {
|
||||
src_ip: src.ip().to_string(),
|
||||
src_port: src.port().into(),
|
||||
dst_host,
|
||||
dst_port: dst_port as _,
|
||||
},
|
||||
)),
|
||||
};
|
||||
send_handshake(&mut send_stream, handshake).await?;
|
||||
let (mut tcp_r, mut tcp_w) = tcp_stream.split();
|
||||
tokio::select! {
|
||||
_ = tokio::io::copy(&mut recv_stream, &mut tcp_w) => {},
|
||||
_ = tokio::io::copy(&mut tcp_r, &mut send_stream) => {},
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
use crate::enhanced_tunnel::quic_over::quic_client::{QuicTunnelClient, send_handshake};
|
||||
use crate::port_mapping::PortMapping;
|
||||
use crate::protocol::client_message::{
|
||||
PortProxyHandshake, QuicProxyHandshake, quic_proxy_handshake,
|
||||
};
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::Context;
|
||||
use bytes::Bytes;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use parking_lot::Mutex;
|
||||
use pnet_packet::ip::IpNextHeaderProtocols;
|
||||
use std::collections::HashMap;
|
||||
use std::net::{Ipv4Addr, SocketAddr};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tokio::net::UdpSocket;
|
||||
use tokio::sync::mpsc::Sender;
|
||||
use tokio::sync::mpsc::error::TrySendError;
|
||||
use tokio_util::codec::{FramedRead, FramedWrite, LengthDelimitedCodec};
|
||||
|
||||
pub async fn start(
|
||||
task_group: &TaskGroup,
|
||||
list: &Vec<PortMapping>,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
for x in list {
|
||||
if x.protocol != IpNextHeaderProtocols::Udp {
|
||||
continue;
|
||||
}
|
||||
log::info!("Starting UDP port mapping on {}", x);
|
||||
let udp = UdpSocket::bind(x.src_addr)
|
||||
.await
|
||||
.with_context(|| format!("Udp port mapping Failed to bind to {}", x.src_addr))?;
|
||||
let udp_socket = Arc::new(udp);
|
||||
let group = task_group.clone();
|
||||
let tunnel_client = quic_tunnel_client.clone();
|
||||
let mapping = x.clone();
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = recv(&group, udp_socket, &mapping, tunnel_client).await {
|
||||
log::error!("recv {:?},mapping:{mapping}", e);
|
||||
}
|
||||
});
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn recv(
|
||||
task_group: &TaskGroup,
|
||||
udp_socket: Arc<UdpSocket>,
|
||||
mapping: &PortMapping,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut buf = vec![0u8; 65536];
|
||||
let dest_map = Arc::new(Mutex::new(HashMap::<SocketAddr, Sender<Bytes>>::new()));
|
||||
loop {
|
||||
let (len, src) = udp_socket.recv_from(&mut buf).await?;
|
||||
let bytes = Bytes::copy_from_slice(&buf[..len]);
|
||||
|
||||
let tx = {
|
||||
let mut map = dest_map.lock();
|
||||
if let Some(tx) = map.get(&src) {
|
||||
tx.clone()
|
||||
} else {
|
||||
let (tx, rx) = tokio::sync::mpsc::channel::<Bytes>(128);
|
||||
let udp_socket = udp_socket.clone();
|
||||
let virtual_target_ip = mapping.virtual_target_ip;
|
||||
let dst_host = mapping.dst_host.clone();
|
||||
let dst_port = mapping.dst_port;
|
||||
let tunnel_client = quic_tunnel_client.clone();
|
||||
task_group.spawn(async move {
|
||||
if let Err(e) = udp_mapping_handle(
|
||||
udp_socket,
|
||||
src,
|
||||
virtual_target_ip,
|
||||
dst_host,
|
||||
dst_port,
|
||||
rx,
|
||||
tunnel_client,
|
||||
)
|
||||
.await
|
||||
{
|
||||
log::error!("udp_mapping_handle {e:?},src:{src}");
|
||||
}
|
||||
});
|
||||
|
||||
map.insert(src, tx.clone());
|
||||
tx
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(err) = tx.try_send(bytes) {
|
||||
match err {
|
||||
TrySendError::Full(_) => {}
|
||||
TrySendError::Closed(_) => {
|
||||
let mut map = dest_map.lock();
|
||||
map.remove(&src);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn udp_mapping_handle(
|
||||
udp_socket: Arc<UdpSocket>,
|
||||
src: SocketAddr,
|
||||
target_ip: Ipv4Addr,
|
||||
dst_host: String,
|
||||
dst_port: u16,
|
||||
mut rx: tokio::sync::mpsc::Receiver<Bytes>,
|
||||
quic_tunnel_client: QuicTunnelClient,
|
||||
) -> anyhow::Result<()> {
|
||||
let (mut send_stream, recv_stream) = quic_tunnel_client.open_bi(target_ip).await?;
|
||||
|
||||
let handshake = QuicProxyHandshake {
|
||||
handshake: Some(quic_proxy_handshake::Handshake::UdpPortMapping(
|
||||
PortProxyHandshake {
|
||||
src_ip: src.ip().to_string(),
|
||||
src_port: src.port().into(),
|
||||
dst_host,
|
||||
dst_port: dst_port as _,
|
||||
},
|
||||
)),
|
||||
};
|
||||
send_handshake(&mut send_stream, handshake).await?;
|
||||
let mut framed_write = FramedWrite::new(send_stream, LengthDelimitedCodec::new());
|
||||
let mut framed_read = FramedRead::new(recv_stream, LengthDelimitedCodec::new());
|
||||
loop {
|
||||
tokio::select! {
|
||||
Some(buf) = framed_read.next()=>{
|
||||
let buf = buf?;
|
||||
udp_socket.send_to(&buf, src).await?;
|
||||
},
|
||||
Some(buf) = rx.recv()=>{
|
||||
framed_write.send(buf).await?
|
||||
},
|
||||
_ = tokio::time::sleep(Duration::from_secs(60)) =>{
|
||||
break;
|
||||
},
|
||||
else => {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
mod proto {
|
||||
include!(concat!(env!("OUT_DIR"), "/protocol.client.rs"));
|
||||
}
|
||||
|
||||
use anyhow::bail;
|
||||
use bytes::BytesMut;
|
||||
use prost::Message;
|
||||
use std::net::{Ipv4Addr, Ipv6Addr};
|
||||
|
||||
use crate::protocol::ProtoToBytesMut;
|
||||
pub use proto::*;
|
||||
|
||||
pub fn encode_nat_info(nat_info: &rust_p2p_core::nat::NatInfo) -> proto::NatInfo {
|
||||
let nat_type = match nat_info.nat_type {
|
||||
rust_p2p_core::nat::NatType::Cone => proto::NatType::Cone,
|
||||
rust_p2p_core::nat::NatType::Symmetric => proto::NatType::Symmetric,
|
||||
};
|
||||
|
||||
proto::NatInfo {
|
||||
nat_type: nat_type.into(),
|
||||
public_ips: nat_info.public_ips.iter().map(|v| (*v).into()).collect(),
|
||||
public_udp_ports: nat_info
|
||||
.public_udp_ports
|
||||
.iter()
|
||||
.map(|v| (*v).into())
|
||||
.collect(),
|
||||
public_port_range: nat_info.public_port_range.into(),
|
||||
local_ipv4s: nat_info.local_ipv4s.iter().map(|v| (*v).into()).collect(),
|
||||
ipv6: nat_info.ipv6.map(|v| v.octets().to_vec()),
|
||||
local_udp_ports: nat_info
|
||||
.local_udp_ports
|
||||
.iter()
|
||||
.map(|v| (*v).into())
|
||||
.collect(),
|
||||
local_tcp_port: nat_info.local_tcp_port.into(),
|
||||
public_tcp_port: nat_info.public_tcp_port.into(),
|
||||
}
|
||||
}
|
||||
pub fn decode_nat_info(msg: proto::NatInfo) -> anyhow::Result<rust_p2p_core::nat::NatInfo> {
|
||||
let nat_type = match msg.nat_type() {
|
||||
proto::NatType::Cone => rust_p2p_core::nat::NatType::Cone,
|
||||
proto::NatType::Symmetric => rust_p2p_core::nat::NatType::Symmetric,
|
||||
};
|
||||
let ipv6: Option<[u8; 16]> = msg.ipv6.and_then(|v| v.as_slice().try_into().ok());
|
||||
|
||||
// Validate all ports fit in u16
|
||||
let validate_port = |p: u32| -> anyhow::Result<u16> {
|
||||
u16::try_from(p).map_err(|_| anyhow::anyhow!("invalid port number: {}", p))
|
||||
};
|
||||
|
||||
let public_udp_ports: Result<Vec<_>, _> = msg
|
||||
.public_udp_ports
|
||||
.into_iter()
|
||||
.map(validate_port)
|
||||
.collect();
|
||||
let local_udp_ports: Result<Vec<_>, _> =
|
||||
msg.local_udp_ports.into_iter().map(validate_port).collect();
|
||||
|
||||
Ok(rust_p2p_core::nat::NatInfo {
|
||||
nat_type,
|
||||
public_ips: msg.public_ips.into_iter().map(|v| v.into()).collect(),
|
||||
public_udp_ports: public_udp_ports?,
|
||||
mapping_tcp_addr: vec![],
|
||||
mapping_udp_addr: vec![],
|
||||
public_port_range: validate_port(msg.public_port_range)?,
|
||||
local_ipv4: msg
|
||||
.local_ipv4s
|
||||
.first()
|
||||
.map(|v| (*v).into())
|
||||
.unwrap_or(Ipv4Addr::UNSPECIFIED),
|
||||
local_ipv4s: msg.local_ipv4s.into_iter().map(|v| v.into()).collect(),
|
||||
ipv6: ipv6.map(Ipv6Addr::from),
|
||||
local_udp_ports: local_udp_ports?,
|
||||
local_tcp_port: validate_port(msg.local_tcp_port)?,
|
||||
public_tcp_port: validate_port(msg.public_tcp_port)?,
|
||||
})
|
||||
}
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct PunchInfo {
|
||||
pub nat_info: rust_p2p_core::nat::NatInfo,
|
||||
}
|
||||
|
||||
impl PunchInfo {
|
||||
pub fn from_slice(buf: &[u8]) -> anyhow::Result<Self> {
|
||||
let msg = proto::PunchInfo::decode(buf)?;
|
||||
let Some(nat_info) = msg.nat_info else {
|
||||
bail!("Punched info decode failed.");
|
||||
};
|
||||
let nat_info = decode_nat_info(nat_info)?;
|
||||
Ok(Self { nat_info })
|
||||
}
|
||||
pub fn encode(&self) -> BytesMut {
|
||||
let message = proto::PunchInfo {
|
||||
nat_info: Some(encode_nat_info(&self.nat_info)),
|
||||
};
|
||||
message.encode_bytes_mut()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,275 @@
|
||||
use crate::protocol::ProtoToBytesMut;
|
||||
pub(crate) use crate::protocol::control_message::proto::SelectiveBroadcast;
|
||||
use crate::protocol::control_message::proto::request_message::RequestPayload;
|
||||
use crate::protocol::control_message::proto::response_message::ResponsePayload;
|
||||
use anyhow::bail;
|
||||
use bytes::BytesMut;
|
||||
use prost::Message;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
mod proto {
|
||||
include!(concat!(env!("OUT_DIR"), "/protocol.control_message.rs"));
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Eq, PartialEq, Default)]
|
||||
pub enum RegistrationMode {
|
||||
#[default]
|
||||
Normal = 0,
|
||||
PreRegister = 1,
|
||||
}
|
||||
|
||||
impl From<RegistrationMode> for proto::RegistrationMode {
|
||||
fn from(mode: RegistrationMode) -> Self {
|
||||
match mode {
|
||||
RegistrationMode::Normal => proto::RegistrationMode::Normal,
|
||||
RegistrationMode::PreRegister => proto::RegistrationMode::PreRegister,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<proto::RegistrationMode> for RegistrationMode {
|
||||
fn from(mode: proto::RegistrationMode) -> Self {
|
||||
match mode {
|
||||
proto::RegistrationMode::Normal => RegistrationMode::Normal,
|
||||
proto::RegistrationMode::PreRegister => RegistrationMode::PreRegister,
|
||||
}
|
||||
}
|
||||
}
|
||||
pub(crate) struct RegRequestMsg {
|
||||
pub network_code: String,
|
||||
pub device_id: String,
|
||||
pub ip: Option<Ipv4Addr>,
|
||||
pub name: String,
|
||||
pub version: String,
|
||||
pub key_sign: Option<String>,
|
||||
pub ip_variable: bool,
|
||||
pub server_id: u32,
|
||||
pub registration_mode: RegistrationMode,
|
||||
}
|
||||
impl RegRequestMsg {
|
||||
// pub fn check(&self) -> anyhow::Result<()> {
|
||||
// if self.network_code.is_empty() {
|
||||
// return Err(anyhow!("network_code cannot be empty"));
|
||||
// }
|
||||
// if self.network_code.len() > MAX_NETWORK_CODE_LEN {
|
||||
// return Err(anyhow!(
|
||||
// "network_code length exceeds {} characters (current: {})",
|
||||
// MAX_NETWORK_CODE_LEN,
|
||||
// self.network_code.len()
|
||||
// ));
|
||||
// }
|
||||
// if self.device_id.is_empty() {
|
||||
// return Err(anyhow!("device_id cannot be empty"));
|
||||
// }
|
||||
// if self.device_id.len() > MAX_DEVICE_ID_LEN {
|
||||
// return Err(anyhow!(
|
||||
// "device_id length exceeds {} characters (current: {})",
|
||||
// MAX_DEVICE_ID_LEN,
|
||||
// self.device_id.len()
|
||||
// ));
|
||||
// }
|
||||
//
|
||||
// if self.name.len() > MAX_NAME_LEN {
|
||||
// return Err(anyhow!(
|
||||
// "name length exceeds {} characters (current: {})",
|
||||
// MAX_NAME_LEN,
|
||||
// self.name.len()
|
||||
// ));
|
||||
// }
|
||||
//
|
||||
// if self.version.len() > MAX_VERSION_LEN {
|
||||
// return Err(anyhow!(
|
||||
// "version length exceeds {} characters (current: {})",
|
||||
// MAX_VERSION_LEN,
|
||||
// self.version.len()
|
||||
// ));
|
||||
// }
|
||||
//
|
||||
// Ok(())
|
||||
// }
|
||||
// pub fn from(msg: proto::RegRequestMsg) -> anyhow::Result<Self> {
|
||||
// Ok(Self {
|
||||
// network_code: msg.network_code,
|
||||
// device_id: msg.device_id,
|
||||
// ip: msg.ip.map(|ip| ip.into()),
|
||||
// name: msg.name,
|
||||
// version: msg.version,
|
||||
// key_sign: msg.key_sign,
|
||||
// ip_variable: msg.ip_variable,
|
||||
// server_id: msg.server_id,
|
||||
// })
|
||||
// }
|
||||
pub fn to(self) -> proto::RegRequestMsg {
|
||||
proto::RegRequestMsg {
|
||||
network_code: self.network_code,
|
||||
device_id: self.device_id,
|
||||
ip: self.ip.map(|ip| ip.into()),
|
||||
name: self.name,
|
||||
version: self.version,
|
||||
key_sign: self.key_sign,
|
||||
ip_variable: self.ip_variable,
|
||||
server_id: self.server_id,
|
||||
registration_mode: proto::RegistrationMode::from(self.registration_mode).into(),
|
||||
}
|
||||
}
|
||||
}
|
||||
#[derive(Debug, Clone, Eq, PartialEq)]
|
||||
pub struct RegResponseMsg {
|
||||
pub ip: Ipv4Addr,
|
||||
pub prefix_len: u8,
|
||||
pub gateway: Ipv4Addr,
|
||||
pub server_version: String,
|
||||
}
|
||||
impl RegResponseMsg {
|
||||
pub fn from(msg: proto::RegResponseMsg) -> anyhow::Result<Self> {
|
||||
Ok(Self {
|
||||
ip: msg.ip.into(),
|
||||
prefix_len: (msg.prefix_len & 0xFF) as u8,
|
||||
gateway: msg.gateway.into(),
|
||||
server_version: msg.server_version,
|
||||
})
|
||||
}
|
||||
pub fn to(self) -> proto::RegResponseMsg {
|
||||
proto::RegResponseMsg {
|
||||
ip: self.ip.into(),
|
||||
prefix_len: self.prefix_len as _,
|
||||
gateway: self.gateway.into(),
|
||||
server_version: self.server_version,
|
||||
}
|
||||
}
|
||||
}
|
||||
#[derive(Debug, Clone, Eq, PartialEq)]
|
||||
pub struct ErrorResponseMsg {
|
||||
pub code: u32,
|
||||
pub message: String,
|
||||
}
|
||||
impl ErrorResponseMsg {
|
||||
pub fn from(msg: proto::ErrorResponseMsg) -> anyhow::Result<Self> {
|
||||
Ok(Self {
|
||||
code: msg.code,
|
||||
message: msg.message,
|
||||
})
|
||||
}
|
||||
pub fn to(self) -> proto::ErrorResponseMsg {
|
||||
proto::ErrorResponseMsg {
|
||||
code: self.code,
|
||||
message: self.message,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Eq, PartialEq)]
|
||||
pub struct ConfirmRegResponseMsg {
|
||||
pub success: bool,
|
||||
}
|
||||
impl ConfirmRegResponseMsg {
|
||||
pub fn from(msg: proto::ConfirmRegResponseMsg) -> anyhow::Result<Self> {
|
||||
Ok(Self {
|
||||
success: msg.success,
|
||||
})
|
||||
}
|
||||
pub fn to(self) -> proto::ConfirmRegResponseMsg {
|
||||
proto::ConfirmRegResponseMsg {
|
||||
success: self.success,
|
||||
}
|
||||
}
|
||||
}
|
||||
pub(crate) enum RequestMessage {
|
||||
Reg(RegRequestMsg),
|
||||
ConfirmReg,
|
||||
}
|
||||
impl RequestMessage {
|
||||
pub fn encode(self) -> BytesMut {
|
||||
let request_payload = match self {
|
||||
RequestMessage::Reg(reg) => RequestPayload::Reg(reg.to()),
|
||||
RequestMessage::ConfirmReg => RequestPayload::ConfirmReg(proto::ConfirmRegMsg {}),
|
||||
};
|
||||
proto::RequestMessage {
|
||||
request_payload: Some(request_payload),
|
||||
}
|
||||
.encode_bytes_mut()
|
||||
}
|
||||
}
|
||||
#[derive(Debug, Clone, Eq, PartialEq)]
|
||||
pub enum ResponseMessage {
|
||||
Reg(RegResponseMsg),
|
||||
Error(ErrorResponseMsg),
|
||||
ConfirmReg(ConfirmRegResponseMsg),
|
||||
}
|
||||
impl ResponseMessage {
|
||||
pub fn from_slice(buf: &[u8]) -> anyhow::Result<Self> {
|
||||
let msg = proto::ResponseMessage::decode(buf)?;
|
||||
let Some(payload) = msg.response_payload else {
|
||||
bail!("unsupported")
|
||||
};
|
||||
match payload {
|
||||
ResponsePayload::Reg(reg) => Ok(ResponseMessage::Reg(RegResponseMsg::from(reg)?)),
|
||||
ResponsePayload::Error(e) => Ok(ResponseMessage::Error(ErrorResponseMsg::from(e)?)),
|
||||
ResponsePayload::ConfirmReg(c) => {
|
||||
Ok(ResponseMessage::ConfirmReg(ConfirmRegResponseMsg::from(c)?))
|
||||
}
|
||||
}
|
||||
}
|
||||
pub fn encode(self) -> BytesMut {
|
||||
let response_payload = match self {
|
||||
ResponseMessage::Reg(reg) => ResponsePayload::Reg(reg.to()),
|
||||
ResponseMessage::Error(e) => ResponsePayload::Error(e.to()),
|
||||
ResponseMessage::ConfirmReg(c) => ResponsePayload::ConfirmReg(c.to()),
|
||||
};
|
||||
proto::ResponseMessage {
|
||||
response_payload: Some(response_payload),
|
||||
}
|
||||
.encode_bytes_mut()
|
||||
}
|
||||
}
|
||||
|
||||
impl SelectiveBroadcast {
|
||||
pub fn new(ips: &[Ipv4Addr], data: Vec<u8>) -> Self {
|
||||
SelectiveBroadcast {
|
||||
ips: ips.iter().map(|v| (*v).into()).collect(),
|
||||
data,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ClientSimpleInfo {
|
||||
pub ip: Ipv4Addr,
|
||||
pub online: bool,
|
||||
}
|
||||
impl ClientSimpleInfo {
|
||||
pub fn from(msg: proto::ClientSimpleInfo) -> anyhow::Result<Self> {
|
||||
Ok(Self {
|
||||
ip: msg.ip.into(),
|
||||
online: msg.online,
|
||||
})
|
||||
}
|
||||
pub fn to(self) -> proto::ClientSimpleInfo {
|
||||
proto::ClientSimpleInfo {
|
||||
ip: self.ip.into(),
|
||||
online: self.online,
|
||||
}
|
||||
}
|
||||
}
|
||||
#[derive(Debug)]
|
||||
pub struct ClientSimpleInfoList {
|
||||
pub data_version: u64,
|
||||
pub list: Vec<ClientSimpleInfo>,
|
||||
pub is_all: bool,
|
||||
pub time: i64,
|
||||
}
|
||||
impl ClientSimpleInfoList {
|
||||
pub fn from_slice(buf: &[u8]) -> anyhow::Result<Self> {
|
||||
let msg = proto::ClientSimpleInfoList::decode(buf)?;
|
||||
let mut list = Vec::with_capacity(msg.list.len());
|
||||
for x in msg.list {
|
||||
list.push(ClientSimpleInfo::from(x)?);
|
||||
}
|
||||
Ok(Self {
|
||||
data_version: msg.data_version,
|
||||
list,
|
||||
is_all: msg.is_all,
|
||||
time: msg.time,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,302 @@
|
||||
/*
|
||||
0 15 31
|
||||
0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| 1 | msg_type(7) |max ttl(4) |curr ttl(4)| C | G | R | reserve(13) |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| seq(32) |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| src ID(32) |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| dest ID(32) |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
| payload(n) |
|
||||
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|
||||
*/
|
||||
#![allow(dead_code)]
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use std::io;
|
||||
use zerocopy::byteorder::{NetworkEndian, U32};
|
||||
use zerocopy::{FromBytes, Immutable, IntoBytes, KnownLayout, Ref, Unaligned};
|
||||
|
||||
#[derive(Debug, FromBytes, IntoBytes, Unaligned, KnownLayout, Immutable)]
|
||||
#[repr(C)]
|
||||
pub struct NetHeader {
|
||||
/// Byte 0: bit7 = 1, bit0..6 = msg_type
|
||||
pub type_byte: u8,
|
||||
/// Byte 1: high 4 = max ttl, low 4 = curr ttl
|
||||
pub ttl_byte: u8,
|
||||
/// Byte 2: C(0x80) | G(0x40) | reserve
|
||||
pub flags_byte: u8,
|
||||
/// Byte 3: reserve
|
||||
pub _reserved: u8,
|
||||
|
||||
pub seq: U32<NetworkEndian>,
|
||||
pub src_id: U32<NetworkEndian>,
|
||||
pub dest_id: U32<NetworkEndian>,
|
||||
}
|
||||
const COMPRESSED: u8 = 0x80;
|
||||
const GATEWAY: u8 = 0x40;
|
||||
const FEC: u8 = 0x20;
|
||||
impl NetHeader {
|
||||
#[inline]
|
||||
pub fn msg_type(&self) -> u8 {
|
||||
self.type_byte & 0x7F
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn set_msg_type(&mut self, msg_type: u8) {
|
||||
self.type_byte = (msg_type & 0x7F) | 0x80;
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn max_ttl(&self) -> u8 {
|
||||
self.ttl_byte >> 4
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn curr_ttl(&self) -> u8 {
|
||||
self.ttl_byte & 0x0F
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn set_ttl(&mut self, max: u8, curr: u8) {
|
||||
self.ttl_byte = (max << 4) | (curr & 0x0F);
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn decr_ttl(&mut self) {
|
||||
let curr = self.curr_ttl();
|
||||
if curr == 0 {
|
||||
return;
|
||||
}
|
||||
self.ttl_byte = (self.ttl_byte & 0xF0) | (curr - 1);
|
||||
}
|
||||
|
||||
fn set_flag(&mut self, mask: u8, val: bool) {
|
||||
if val {
|
||||
self.flags_byte |= mask;
|
||||
} else {
|
||||
self.flags_byte &= !mask;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Copy, Clone, Eq, PartialEq, Debug)]
|
||||
pub enum MsgType {
|
||||
Turn = 1,
|
||||
Broadcast = 2,
|
||||
ExcludeBroadcast = 3,
|
||||
TargetBroadcast = 4,
|
||||
|
||||
Ping = 5,
|
||||
Pong = 6,
|
||||
PingTurn = 7,
|
||||
PongTurn = 8,
|
||||
|
||||
PunchStart1 = 9,
|
||||
PunchStart2 = 10,
|
||||
PunchReq = 11,
|
||||
PunchRes = 12,
|
||||
|
||||
PushClientIps = 13,
|
||||
|
||||
RpcReq = 14,
|
||||
RpcRes = 15,
|
||||
|
||||
Quic = 17,
|
||||
}
|
||||
impl From<MsgType> for u8 {
|
||||
fn from(val: MsgType) -> Self {
|
||||
val as u8
|
||||
}
|
||||
}
|
||||
|
||||
impl TryFrom<u8> for MsgType {
|
||||
type Error = io::Error;
|
||||
|
||||
fn try_from(value: u8) -> Result<Self, Self::Error> {
|
||||
let val = match value {
|
||||
1 => MsgType::Turn,
|
||||
2 => MsgType::Broadcast,
|
||||
3 => MsgType::ExcludeBroadcast,
|
||||
4 => MsgType::TargetBroadcast,
|
||||
|
||||
5 => MsgType::Ping,
|
||||
6 => MsgType::Pong,
|
||||
7 => MsgType::PingTurn,
|
||||
8 => MsgType::PongTurn,
|
||||
|
||||
9 => MsgType::PunchStart1,
|
||||
10 => MsgType::PunchStart2,
|
||||
11 => MsgType::PunchReq,
|
||||
12 => MsgType::PunchRes,
|
||||
|
||||
13 => MsgType::PushClientIps,
|
||||
|
||||
14 => MsgType::RpcReq,
|
||||
15 => MsgType::RpcRes,
|
||||
|
||||
17 => MsgType::Quic,
|
||||
_ => {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!("invalid msg type:{value}"),
|
||||
));
|
||||
}
|
||||
};
|
||||
Ok(val)
|
||||
}
|
||||
}
|
||||
|
||||
pub const HEAD_LENGTH: usize = std::mem::size_of::<NetHeader>();
|
||||
|
||||
pub struct NetPacket<B> {
|
||||
buffer: B,
|
||||
}
|
||||
impl<B: AsRef<[u8]>> NetPacket<B> {
|
||||
pub fn new(buffer: B) -> io::Result<NetPacket<B>> {
|
||||
if buffer.as_ref().len() < HEAD_LENGTH {
|
||||
return Err(io::ErrorKind::InvalidInput.into());
|
||||
}
|
||||
Ok(NetPacket { buffer })
|
||||
}
|
||||
fn header(&self) -> Ref<&[u8], NetHeader> {
|
||||
// Safe: NetHeader is Unaligned and length is validated in new()
|
||||
let (header, _) = Ref::<&[u8], NetHeader>::from_prefix(self.buffer.as_ref()).unwrap();
|
||||
header
|
||||
}
|
||||
pub fn buffer(&self) -> &[u8] {
|
||||
self.buffer.as_ref()
|
||||
}
|
||||
pub fn into_buffer(self) -> B {
|
||||
self.buffer
|
||||
}
|
||||
pub fn source_buf(&self) -> &B {
|
||||
&self.buffer
|
||||
}
|
||||
pub fn msg_type(&self) -> io::Result<MsgType> {
|
||||
self.header().msg_type().try_into()
|
||||
}
|
||||
pub fn max_ttl(&self) -> u8 {
|
||||
self.header().max_ttl()
|
||||
}
|
||||
pub fn ttl(&self) -> u8 {
|
||||
self.header().curr_ttl()
|
||||
}
|
||||
|
||||
pub fn seq(&self) -> u32 {
|
||||
self.header().seq.get()
|
||||
}
|
||||
|
||||
pub fn src_id(&self) -> u32 {
|
||||
self.header().src_id.get()
|
||||
}
|
||||
|
||||
pub fn dest_id(&self) -> u32 {
|
||||
self.header().dest_id.get()
|
||||
}
|
||||
pub fn is_compressed(&self) -> bool {
|
||||
(self.header().flags_byte & COMPRESSED) != 0
|
||||
}
|
||||
pub fn is_gateway(&self) -> bool {
|
||||
(self.header().flags_byte & GATEWAY) != 0
|
||||
}
|
||||
pub fn is_fec(&self) -> bool {
|
||||
(self.header().flags_byte & FEC) != 0
|
||||
}
|
||||
pub fn head(&self) -> &[u8] {
|
||||
&self.buffer.as_ref()[..HEAD_LENGTH]
|
||||
}
|
||||
pub fn payload(&self) -> &[u8] {
|
||||
&self.buffer.as_ref()[HEAD_LENGTH..]
|
||||
}
|
||||
}
|
||||
|
||||
impl<B: AsRef<[u8]> + AsMut<[u8]>> NetPacket<B> {
|
||||
fn header_mut(&mut self) -> Ref<&mut [u8], NetHeader> {
|
||||
// Safe: NetHeader is Unaligned and length is validated in new()
|
||||
let (header, _) = Ref::<&mut [u8], NetHeader>::from_prefix(self.buffer.as_mut()).unwrap();
|
||||
header
|
||||
}
|
||||
|
||||
pub fn set_msg_type(&mut self, msg_type: MsgType) {
|
||||
self.header_mut().set_msg_type(msg_type.into());
|
||||
}
|
||||
|
||||
pub fn decr_ttl(&mut self){
|
||||
self.header_mut().decr_ttl()
|
||||
}
|
||||
|
||||
pub fn set_ttl(&mut self, ttl: u8) {
|
||||
self.header_mut().set_ttl(ttl, ttl);
|
||||
}
|
||||
|
||||
pub fn set_seq(&mut self, seq: u32) {
|
||||
self.header_mut().seq.set(seq);
|
||||
}
|
||||
|
||||
pub fn set_src_id(&mut self, id: u32) {
|
||||
self.header_mut().src_id.set(id);
|
||||
}
|
||||
|
||||
pub fn set_dest_id(&mut self, id: u32) {
|
||||
self.header_mut().dest_id.set(id);
|
||||
}
|
||||
|
||||
pub fn set_compressed_flag(&mut self, compressed: bool) {
|
||||
self.header_mut().set_flag(COMPRESSED, compressed);
|
||||
}
|
||||
pub fn set_gateway_flag(&mut self, gateway: bool) {
|
||||
self.header_mut().set_flag(GATEWAY, gateway);
|
||||
}
|
||||
pub fn set_fec_flag(&mut self, fec: bool) {
|
||||
self.header_mut().set_flag(FEC, fec);
|
||||
}
|
||||
|
||||
pub fn set_payload(&mut self, data: &[u8]) -> io::Result<()> {
|
||||
let buf = self.buffer.as_mut();
|
||||
if buf.len() < HEAD_LENGTH + data.len() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"Invalid message length",
|
||||
));
|
||||
}
|
||||
buf[HEAD_LENGTH..HEAD_LENGTH + data.len()].copy_from_slice(data);
|
||||
Ok(())
|
||||
}
|
||||
pub fn head_mut(&mut self) -> &mut [u8] {
|
||||
&mut self.buffer.as_mut()[..HEAD_LENGTH]
|
||||
}
|
||||
pub fn payload_mut(&mut self) -> &mut [u8] {
|
||||
&mut self.buffer.as_mut()[HEAD_LENGTH..]
|
||||
}
|
||||
pub fn source_buf_mut(&mut self) -> &mut B {
|
||||
&mut self.buffer
|
||||
}
|
||||
}
|
||||
|
||||
impl Clone for NetPacket<Bytes> {
|
||||
fn clone(&self) -> Self {
|
||||
NetPacket {
|
||||
buffer: self.buffer.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl NetPacket<BytesMut> {
|
||||
pub fn into_bytes(self) -> NetPacket<Bytes> {
|
||||
NetPacket {
|
||||
buffer: self.buffer.freeze(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl NetPacket<TransmissionBytes> {
|
||||
pub fn into_bytes(self) -> NetPacket<Bytes> {
|
||||
NetPacket {
|
||||
buffer: self.buffer.into_bytes().freeze(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
use bytes::BytesMut;
|
||||
use prost::Message;
|
||||
|
||||
pub(crate) mod client_message;
|
||||
pub mod control_message;
|
||||
pub(crate) mod ip_packet_protocol;
|
||||
pub(crate) mod rpc_message;
|
||||
pub(crate) mod transmission;
|
||||
|
||||
pub trait ProtoToBytesMut: Message {
|
||||
fn encode_bytes_mut(&self) -> BytesMut
|
||||
where
|
||||
Self: Sized,
|
||||
{
|
||||
let mut bytes_mut = BytesMut::with_capacity(self.encoded_len());
|
||||
self.encode_raw(&mut bytes_mut);
|
||||
bytes_mut
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Message> ProtoToBytesMut for T {}
|
||||
@@ -0,0 +1,4 @@
|
||||
mod proto {
|
||||
include!(concat!(env!("OUT_DIR"), "/protocol.rpc.rs"));
|
||||
}
|
||||
pub use proto::*;
|
||||
@@ -0,0 +1,266 @@
|
||||
use bytes::{Buf, Bytes, BytesMut};
|
||||
use std::borrow::{Borrow, BorrowMut};
|
||||
use std::io;
|
||||
use std::ops::{Deref, DerefMut};
|
||||
|
||||
const DEFAULT_BUF_SIZE: usize = 2048;
|
||||
#[derive(Clone)]
|
||||
pub struct TransmissionBytes {
|
||||
buf: BytesMut,
|
||||
start: usize,
|
||||
end: usize,
|
||||
}
|
||||
impl From<BytesMut> for TransmissionBytes {
|
||||
fn from(buf: BytesMut) -> TransmissionBytes {
|
||||
let end = buf.len();
|
||||
Self { buf, start: 0, end }
|
||||
}
|
||||
}
|
||||
impl From<Bytes> for TransmissionBytes {
|
||||
fn from(buf: Bytes) -> TransmissionBytes {
|
||||
let end = buf.len();
|
||||
Self {
|
||||
buf: BytesMut::from(buf),
|
||||
start: 0,
|
||||
end,
|
||||
}
|
||||
}
|
||||
}
|
||||
impl From<&[u8]> for TransmissionBytes {
|
||||
fn from(buf: &[u8]) -> TransmissionBytes {
|
||||
let end = buf.len();
|
||||
Self {
|
||||
buf: BytesMut::from(buf),
|
||||
start: 0,
|
||||
end,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl TransmissionBytes {
|
||||
pub fn new_offset(start: usize) -> Self {
|
||||
TransmissionBytes {
|
||||
buf: BytesMut::zeroed(DEFAULT_BUF_SIZE),
|
||||
start,
|
||||
end: start,
|
||||
}
|
||||
}
|
||||
pub fn new_offset_zeroed(start: usize) -> Self {
|
||||
TransmissionBytes {
|
||||
buf: BytesMut::zeroed(DEFAULT_BUF_SIZE),
|
||||
start,
|
||||
end: DEFAULT_BUF_SIZE,
|
||||
}
|
||||
}
|
||||
pub fn zeroed(cap: usize) -> Self {
|
||||
TransmissionBytes {
|
||||
buf: BytesMut::zeroed(cap),
|
||||
start: 0,
|
||||
end: cap,
|
||||
}
|
||||
}
|
||||
pub fn zeroed_size(size: usize, reserve: usize) -> Self {
|
||||
TransmissionBytes {
|
||||
buf: BytesMut::zeroed(size + reserve),
|
||||
start: 0,
|
||||
end: size,
|
||||
}
|
||||
}
|
||||
#[allow(dead_code)]
|
||||
pub fn with_capacity(head_room: usize, capacity: usize) -> Self {
|
||||
TransmissionBytes {
|
||||
buf: BytesMut::zeroed(capacity),
|
||||
start: head_room,
|
||||
end: head_room,
|
||||
}
|
||||
}
|
||||
pub fn len(&self) -> usize {
|
||||
self.end - self.start
|
||||
}
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.len() == 0
|
||||
}
|
||||
#[allow(dead_code)]
|
||||
pub fn capacity(&self) -> usize {
|
||||
self.buf.capacity()
|
||||
}
|
||||
/// 头部可用空间(可向前扩展的字节数)
|
||||
#[inline]
|
||||
pub fn head_room(&self) -> usize {
|
||||
self.start
|
||||
}
|
||||
|
||||
/// 尾部可用空间(可向后扩展的字节数)
|
||||
#[inline]
|
||||
#[allow(dead_code)]
|
||||
pub fn tail_room(&self) -> usize {
|
||||
self.buf.capacity() - self.end
|
||||
}
|
||||
#[inline]
|
||||
fn as_slice(&self) -> &[u8] {
|
||||
&self.buf[self.start..self.end]
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn as_slice_mut(&mut self) -> &mut [u8] {
|
||||
&mut self.buf[self.start..self.end]
|
||||
}
|
||||
pub fn put(&mut self, data: &[u8]) -> io::Result<()> {
|
||||
let need = data.len();
|
||||
let free = self.buf.capacity() - self.end;
|
||||
|
||||
if need > free {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!("data too large:need={need},free={free}"),
|
||||
));
|
||||
}
|
||||
|
||||
self.buf[self.end..self.end + need].copy_from_slice(data);
|
||||
self.end += need;
|
||||
Ok(())
|
||||
}
|
||||
pub fn retreat_head(&mut self, len: usize) -> io::Result<()> {
|
||||
if len > self.head_room() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!(
|
||||
"retreat_head beyond start: len={len}, head_room={}",
|
||||
self.head_room()
|
||||
),
|
||||
));
|
||||
}
|
||||
self.start -= len;
|
||||
Ok(())
|
||||
}
|
||||
pub fn advance_head(&mut self, len: usize) -> io::Result<()> {
|
||||
if len > self.len() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!(
|
||||
"advance_head beyond end: len={len}, data_len={}",
|
||||
self.len()
|
||||
),
|
||||
));
|
||||
}
|
||||
self.start += len;
|
||||
Ok(())
|
||||
}
|
||||
pub fn set_len(&mut self, new_len: usize) -> io::Result<()> {
|
||||
let new_end = self.start + new_len;
|
||||
if new_end > self.buf.capacity() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!(
|
||||
"set_len exceeds capacity: new_len={new_len}, max={}",
|
||||
self.buf.capacity() - self.start
|
||||
),
|
||||
));
|
||||
}
|
||||
self.end = new_end;
|
||||
Ok(())
|
||||
}
|
||||
pub fn resize(&mut self, new_len: usize, value: u8) {
|
||||
let new_end = self.start + new_len;
|
||||
self.buf.resize(new_end, value);
|
||||
self.end = new_end;
|
||||
}
|
||||
pub fn extend_end(&mut self, n: usize) {
|
||||
if self.end + n > self.buf.len() {
|
||||
self.buf.resize(self.end + n, 0);
|
||||
}
|
||||
self.end += n;
|
||||
}
|
||||
pub fn shrink_end(&mut self, n: usize) {
|
||||
if n >= self.end - self.start {
|
||||
self.end = self.start;
|
||||
} else {
|
||||
self.end -= n;
|
||||
}
|
||||
}
|
||||
#[allow(dead_code)]
|
||||
pub fn clear(&mut self) {
|
||||
self.start = 0;
|
||||
self.end = 0;
|
||||
}
|
||||
pub fn into_bytes(mut self) -> BytesMut {
|
||||
self.buf.truncate(self.end);
|
||||
if self.start > 0 {
|
||||
self.buf.advance(self.start);
|
||||
}
|
||||
self.buf
|
||||
}
|
||||
}
|
||||
|
||||
impl AsRef<[u8]> for TransmissionBytes {
|
||||
#[inline]
|
||||
fn as_ref(&self) -> &[u8] {
|
||||
self.as_slice()
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for TransmissionBytes {
|
||||
type Target = [u8];
|
||||
|
||||
#[inline]
|
||||
fn deref(&self) -> &[u8] {
|
||||
self.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl AsMut<[u8]> for TransmissionBytes {
|
||||
#[inline]
|
||||
fn as_mut(&mut self) -> &mut [u8] {
|
||||
self.as_slice_mut()
|
||||
}
|
||||
}
|
||||
|
||||
impl DerefMut for TransmissionBytes {
|
||||
#[inline]
|
||||
fn deref_mut(&mut self) -> &mut [u8] {
|
||||
self.as_mut()
|
||||
}
|
||||
}
|
||||
|
||||
impl Borrow<[u8]> for TransmissionBytes {
|
||||
fn borrow(&self) -> &[u8] {
|
||||
self.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl BorrowMut<[u8]> for TransmissionBytes {
|
||||
fn borrow_mut(&mut self) -> &mut [u8] {
|
||||
self.as_mut()
|
||||
}
|
||||
}
|
||||
|
||||
pub trait ShrinkEnd {
|
||||
fn shrink_end(&mut self, n: usize);
|
||||
}
|
||||
pub trait ExtendEnd {
|
||||
fn extend_end(&mut self, n: usize);
|
||||
}
|
||||
|
||||
impl ShrinkEnd for TransmissionBytes {
|
||||
fn shrink_end(&mut self, n: usize) {
|
||||
self.shrink_end(n);
|
||||
}
|
||||
}
|
||||
|
||||
impl ExtendEnd for TransmissionBytes {
|
||||
fn extend_end(&mut self, n: usize) {
|
||||
self.extend_end(n);
|
||||
}
|
||||
}
|
||||
|
||||
impl ShrinkEnd for &mut TransmissionBytes {
|
||||
fn shrink_end(&mut self, n: usize) {
|
||||
TransmissionBytes::shrink_end(self, n);
|
||||
}
|
||||
}
|
||||
|
||||
impl ExtendEnd for &mut TransmissionBytes {
|
||||
fn extend_end(&mut self, n: usize) {
|
||||
TransmissionBytes::extend_end(self, n);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
use anyhow::{Context, Result};
|
||||
use rcgen::{CertificateParams, DnType, KeyPair, PKCS_ED25519, SerialNumber};
|
||||
use rustls::pki_types;
|
||||
use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer};
|
||||
use sha2::{Digest, Sha256};
|
||||
use time::{Duration, OffsetDateTime};
|
||||
|
||||
const ED25519_PKCS8_V1_PREFIX: [u8; 16] = [
|
||||
0x30, 0x2e, // Sequence (len 46)
|
||||
0x02, 0x01, 0x00, // Version 0
|
||||
0x30, 0x05, // Sequence (len 5)
|
||||
0x06, 0x03, 0x2b, 0x65, 0x70, // OID: 1.3.101.112 (Ed25519)
|
||||
0x04, 0x22, // Octet String (len 34) - 包装私钥
|
||||
0x04, 0x20, // Octet String (len 32) - 内部 CurvePrivateKey
|
||||
];
|
||||
|
||||
pub fn generate_deterministic_cert(
|
||||
password: &str,
|
||||
) -> Result<(CertificateDer<'static>, PrivateKeyDer<'static>)> {
|
||||
// 基于密码生成 32 字节的确定性种子
|
||||
let seed = derive_seed_from_password(password);
|
||||
|
||||
let mut pkcs8_bytes = Vec::with_capacity(ED25519_PKCS8_V1_PREFIX.len() + seed.len());
|
||||
pkcs8_bytes.extend_from_slice(&ED25519_PKCS8_V1_PREFIX);
|
||||
pkcs8_bytes.extend_from_slice(&seed);
|
||||
let private_key_der = pki_types::PrivateKeyDer::try_from(pkcs8_bytes.clone())
|
||||
.map_err(|e| anyhow::anyhow!("Failed to convert private key: {}", e))?;
|
||||
let key_pair = KeyPair::from_der_and_sign_algo(&private_key_der, &PKCS_ED25519)
|
||||
.context("Failed to load determinstic Ed25519 key")?;
|
||||
|
||||
let mut params = CertificateParams::new(vec!["deterministic-node".to_string()])?;
|
||||
params
|
||||
.distinguished_name
|
||||
.push(DnType::CommonName, "Deterministic Self-Signed Cert");
|
||||
|
||||
let not_before = OffsetDateTime::UNIX_EPOCH;
|
||||
let not_after = not_before + Duration::days(365 * 1000);
|
||||
params.not_before = not_before;
|
||||
params.not_after = not_after;
|
||||
|
||||
let serial_number_bytes = derive_serial_number(password);
|
||||
params.serial_number = Some(SerialNumber::from_slice(&serial_number_bytes));
|
||||
|
||||
// Ed25519 签名是确定性的 (RFC 8032),不需要随机数,因此每次运行结果字节完全一致
|
||||
let cert = params
|
||||
.self_signed(&key_pair)
|
||||
.context("Failed to sign certificate")?;
|
||||
|
||||
let cert_der = cert.der().clone();
|
||||
|
||||
let private_key_der = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(pkcs8_bytes));
|
||||
|
||||
Ok((cert_der, private_key_der))
|
||||
}
|
||||
|
||||
fn derive_serial_number(password: &str) -> [u8; 20] {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(b"vnt-serial-v1:");
|
||||
hasher.update(password.as_bytes());
|
||||
let result = hasher.finalize();
|
||||
|
||||
let mut serial = [0u8; 20];
|
||||
serial.copy_from_slice(&result[..20]);
|
||||
serial
|
||||
}
|
||||
|
||||
fn derive_seed_from_password(password: &str) -> [u8; 32] {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(b"vnt-ed25519-seed-v1:");
|
||||
hasher.update(password.as_bytes());
|
||||
|
||||
let mut result = hasher.finalize();
|
||||
for _ in 0..10 {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(result);
|
||||
result = hasher.finalize();
|
||||
}
|
||||
|
||||
result.into()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_deterministic_cert_generation() {
|
||||
let password = "test_password_123";
|
||||
|
||||
// 生成两次证书
|
||||
let (cert1, key1) = generate_deterministic_cert(password).unwrap();
|
||||
let (cert2, key2) = generate_deterministic_cert(password).unwrap();
|
||||
|
||||
// 验证私钥相同
|
||||
assert_eq!(
|
||||
key1.secret_der(),
|
||||
key2.secret_der(),
|
||||
"Private keys should be identical"
|
||||
);
|
||||
|
||||
// 验证证书相同
|
||||
assert_eq!(cert1, cert2, "Certificates should be identical");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_different_passwords_generate_different_certs() {
|
||||
let (cert1, key1) = generate_deterministic_cert("password1").unwrap();
|
||||
let (cert2, key2) = generate_deterministic_cert("password2").unwrap();
|
||||
|
||||
// 不同密码应该生成不同的证书和密钥
|
||||
assert_ne!(cert1, cert2);
|
||||
assert_ne!(key1.secret_der(), key2.secret_der());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
pub(crate) mod cert;
|
||||
pub mod verifier;
|
||||
@@ -0,0 +1,228 @@
|
||||
use anyhow::Context;
|
||||
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
|
||||
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
|
||||
use rustls::{CertificateError, ClientConfig, Error, SignatureScheme};
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::fmt;
|
||||
use std::str::FromStr;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct FingerprintVerifier {
|
||||
pub expected_fingerprint: [u8; 32],
|
||||
}
|
||||
impl FingerprintVerifier {
|
||||
pub fn new(expected_fingerprint: [u8; 32]) -> Self {
|
||||
Self {
|
||||
expected_fingerprint,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ServerCertVerifier for FingerprintVerifier {
|
||||
fn verify_server_cert(
|
||||
&self,
|
||||
end_entity: &CertificateDer<'_>,
|
||||
_intermediates: &[CertificateDer<'_>],
|
||||
_server_name: &ServerName,
|
||||
_ocsp_response: &[u8],
|
||||
_now: UnixTime,
|
||||
) -> Result<ServerCertVerified, Error> {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(end_entity.as_ref());
|
||||
let calculated_hash: [u8; 32] = hasher.finalize().into();
|
||||
|
||||
if calculated_hash == self.expected_fingerprint {
|
||||
Ok(ServerCertVerified::assertion())
|
||||
} else {
|
||||
log::error!(
|
||||
"Certificate fingerprint mismatch. Expected: {:X?}, Got: {:X?}",
|
||||
self.expected_fingerprint,
|
||||
calculated_hash
|
||||
);
|
||||
Err(Error::InvalidCertificate(CertificateError::BadSignature))
|
||||
}
|
||||
}
|
||||
|
||||
fn verify_tls12_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, rustls::Error> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls13_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, rustls::Error> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
|
||||
vec![
|
||||
// RSA schemes
|
||||
rustls::SignatureScheme::RSA_PKCS1_SHA256,
|
||||
rustls::SignatureScheme::RSA_PKCS1_SHA384,
|
||||
rustls::SignatureScheme::RSA_PKCS1_SHA512,
|
||||
rustls::SignatureScheme::RSA_PSS_SHA256,
|
||||
rustls::SignatureScheme::RSA_PSS_SHA384,
|
||||
rustls::SignatureScheme::RSA_PSS_SHA512,
|
||||
// ECDSA schemes
|
||||
rustls::SignatureScheme::ECDSA_NISTP256_SHA256,
|
||||
rustls::SignatureScheme::ECDSA_NISTP384_SHA384,
|
||||
rustls::SignatureScheme::ECDSA_NISTP521_SHA512,
|
||||
// EdDSA schemes
|
||||
rustls::SignatureScheme::ED25519,
|
||||
rustls::SignatureScheme::ED448,
|
||||
]
|
||||
}
|
||||
}
|
||||
#[derive(Debug)]
|
||||
pub struct InsecureVerifier;
|
||||
|
||||
impl ServerCertVerifier for InsecureVerifier {
|
||||
fn verify_server_cert(
|
||||
&self,
|
||||
_end_entity: &CertificateDer<'_>,
|
||||
_intermediates: &[CertificateDer<'_>],
|
||||
_server_name: &ServerName,
|
||||
_ocsp_response: &[u8],
|
||||
_now: UnixTime,
|
||||
) -> Result<rustls::client::danger::ServerCertVerified, Error> {
|
||||
Ok(rustls::client::danger::ServerCertVerified::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls12_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, Error> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls13_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &rustls::DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, Error> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
|
||||
vec![
|
||||
SignatureScheme::RSA_PKCS1_SHA256,
|
||||
SignatureScheme::RSA_PKCS1_SHA384,
|
||||
SignatureScheme::RSA_PKCS1_SHA512,
|
||||
SignatureScheme::RSA_PSS_SHA256,
|
||||
SignatureScheme::RSA_PSS_SHA384,
|
||||
SignatureScheme::RSA_PSS_SHA512,
|
||||
SignatureScheme::ECDSA_NISTP256_SHA256,
|
||||
SignatureScheme::ECDSA_NISTP384_SHA384,
|
||||
SignatureScheme::ECDSA_NISTP521_SHA512,
|
||||
SignatureScheme::ED25519,
|
||||
SignatureScheme::ED448,
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
pub fn load_root_cert() -> anyhow::Result<rustls::RootCertStore> {
|
||||
let mut root_cert_store = rustls::RootCertStore::empty();
|
||||
let certs = rustls_native_certs::load_native_certs().certs;
|
||||
for cert in certs {
|
||||
root_cert_store
|
||||
.add(cert)
|
||||
.context("Failed to add native cert to store")?;
|
||||
}
|
||||
Ok(root_cert_store)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub enum CertValidationMode {
|
||||
#[default]
|
||||
InsecureSkipVerification,
|
||||
VerifyFingerprint([u8; 32]),
|
||||
Standard,
|
||||
}
|
||||
impl FromStr for CertValidationMode {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(value: &str) -> Result<Self, Self::Err> {
|
||||
let val = value.trim().to_lowercase();
|
||||
if val == "skip" {
|
||||
return Ok(CertValidationMode::InsecureSkipVerification);
|
||||
}
|
||||
if val == "standard" {
|
||||
return Ok(CertValidationMode::Standard);
|
||||
}
|
||||
if let Some(hex_str) = val.strip_prefix("finger:") {
|
||||
let decoded =
|
||||
hex::decode(hex_str).map_err(|e| format!("Invalid hex in fingerprint: {}", e))?;
|
||||
|
||||
if decoded.len() != 32 {
|
||||
return Err(format!(
|
||||
"Fingerprint must be 32 bytes (64 hex chars), got {} bytes",
|
||||
decoded.len()
|
||||
));
|
||||
}
|
||||
|
||||
let mut arr = [0u8; 32];
|
||||
arr.copy_from_slice(&decoded);
|
||||
return Ok(CertValidationMode::VerifyFingerprint(arr));
|
||||
}
|
||||
Err(format!("Unknown certificate validation mode: {}", value))
|
||||
}
|
||||
}
|
||||
impl fmt::Display for CertValidationMode {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
CertValidationMode::InsecureSkipVerification => {
|
||||
write!(f, "skip")
|
||||
}
|
||||
CertValidationMode::Standard => {
|
||||
write!(f, "standard")
|
||||
}
|
||||
CertValidationMode::VerifyFingerprint(fingerprint) => {
|
||||
let hex_str = hex::encode(fingerprint);
|
||||
|
||||
write!(f, "finger:{}", hex_str)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl CertValidationMode {
|
||||
pub fn build_verifier(&self) -> anyhow::Result<std::sync::Arc<dyn ServerCertVerifier>> {
|
||||
match self {
|
||||
CertValidationMode::InsecureSkipVerification => {
|
||||
Ok(std::sync::Arc::new(InsecureVerifier))
|
||||
}
|
||||
CertValidationMode::VerifyFingerprint(fingerprint) => {
|
||||
Ok(std::sync::Arc::new(FingerprintVerifier {
|
||||
expected_fingerprint: *fingerprint,
|
||||
}))
|
||||
}
|
||||
CertValidationMode::Standard => {
|
||||
let root_store = load_root_cert()?;
|
||||
let verifier =
|
||||
rustls::client::WebPkiServerVerifier::builder(std::sync::Arc::new(root_store))
|
||||
.build()?;
|
||||
|
||||
Ok(verifier)
|
||||
}
|
||||
}
|
||||
}
|
||||
pub fn create_tls_client_config(&self) -> anyhow::Result<ClientConfig> {
|
||||
let verifier = self.build_verifier()?;
|
||||
|
||||
let config = ClientConfig::builder()
|
||||
.dangerous()
|
||||
.with_custom_certificate_verifier(verifier)
|
||||
.with_no_client_auth();
|
||||
Ok(config)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
use crate::context::NetworkAddr;
|
||||
use crate::nat::internal_nat::InternalNatInbound;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tun::TunDataInbound;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub enum EnhancedTunInbound {
|
||||
Tun(TunDataInbound),
|
||||
Nat(InternalNatInbound),
|
||||
}
|
||||
impl EnhancedTunInbound {
|
||||
pub async fn inbound(&self, data: TransmissionBytes, net: &NetworkAddr) -> anyhow::Result<()> {
|
||||
match self {
|
||||
EnhancedTunInbound::Tun(tun) => tun.send(data, net).await,
|
||||
EnhancedTunInbound::Nat(nat) => nat.send(&data, net).await,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
use crate::enhanced_tunnel::outbound::EnhancedOutbound;
|
||||
use crate::protocol::ip_packet_protocol::HEAD_LENGTH;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::utils::task_control::{SubTask, TaskGroup};
|
||||
use anyhow::{Context, bail};
|
||||
use bytes::BytesMut;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use std::io;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::mpsc::{Receiver, Sender};
|
||||
use tun_rs::async_framed::{Decoder, DeviceFramedRead, DeviceFramedWrite, Encoder};
|
||||
use tun_rs::{AsyncDevice, DeviceBuilder};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct DeviceIOManager {
|
||||
task_group: TaskGroup,
|
||||
device: DeviceMutex,
|
||||
}
|
||||
type DeviceMutex = Arc<tokio::sync::Mutex<(Option<DeviceTask>, Option<(Ipv4Addr, u8)>)>>;
|
||||
pub struct DeviceTask {
|
||||
device: Arc<AsyncDevice>,
|
||||
task_recv: SubTask,
|
||||
task_send: SubTask,
|
||||
}
|
||||
#[derive(Debug, Default)]
|
||||
pub struct DeviceConfig {
|
||||
pub tun_name: Option<String>,
|
||||
#[cfg(unix)]
|
||||
pub tun_fd: Option<i32>,
|
||||
pub mtu: Option<u16>,
|
||||
}
|
||||
|
||||
impl DeviceConfig {
|
||||
pub fn set_tun_name(mut self, tun_name: String) -> Self {
|
||||
self.tun_name = Some(tun_name);
|
||||
self
|
||||
}
|
||||
#[cfg(unix)]
|
||||
pub fn set_tun_fd(mut self, tun_fd: i32) -> Self {
|
||||
self.tun_fd = Some(tun_fd);
|
||||
self
|
||||
}
|
||||
pub fn set_mtu(mut self, mtu: u16) -> Self {
|
||||
self.mtu = Some(mtu);
|
||||
self
|
||||
}
|
||||
}
|
||||
#[derive(Clone)]
|
||||
pub struct TunInbound {
|
||||
pub(crate) sender: Sender<TransmissionBytes>,
|
||||
}
|
||||
|
||||
pub struct TunReceiver {
|
||||
receiver: Receiver<TransmissionBytes>,
|
||||
}
|
||||
pub fn tun_channel() -> (TunInbound, TunReceiver) {
|
||||
let (sender, receiver) = tokio::sync::mpsc::channel(1024);
|
||||
(TunInbound { sender }, TunReceiver { receiver })
|
||||
}
|
||||
|
||||
impl DeviceIOManager {
|
||||
pub fn new(task_group: TaskGroup) -> DeviceIOManager {
|
||||
Self {
|
||||
task_group,
|
||||
device: Arc::new(Default::default()),
|
||||
}
|
||||
}
|
||||
pub async fn stop_task(&self) {
|
||||
let mut guard = self.device.lock().await;
|
||||
if let Some(dev) = guard.0.take() {
|
||||
dev.task_recv.stop().await;
|
||||
dev.task_send.stop().await;
|
||||
}
|
||||
}
|
||||
pub async fn start_task(
|
||||
&self,
|
||||
device_config: DeviceConfig,
|
||||
receiver: TunReceiver,
|
||||
enhanced_outbound: EnhancedOutbound,
|
||||
) -> anyhow::Result<()> {
|
||||
self.stop_task().await;
|
||||
let task = create(
|
||||
&self.task_group,
|
||||
device_config,
|
||||
receiver.receiver,
|
||||
enhanced_outbound,
|
||||
)?;
|
||||
self.device.lock().await.0.replace(task);
|
||||
Ok(())
|
||||
}
|
||||
#[cfg(not(target_os = "android"))]
|
||||
pub async fn tun_if_index(&self) -> anyhow::Result<u32> {
|
||||
let guard = self.device.lock().await;
|
||||
if let Some(v) = &guard.0 {
|
||||
Ok(v.device.if_index()?)
|
||||
} else {
|
||||
bail!("device doesn't exist")
|
||||
}
|
||||
}
|
||||
pub async fn set_network(&self, ip: Ipv4Addr, prefix_len: u8) -> anyhow::Result<()> {
|
||||
let mut guard = self.device.lock().await;
|
||||
let Some(dev) = guard.0.as_ref() else {
|
||||
bail!("未启动tun")
|
||||
};
|
||||
if let Some(v) = guard.1.as_ref()
|
||||
&& v.0 == ip
|
||||
&& v.1 == prefix_len
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
dev.device
|
||||
.set_network_address(ip, prefix_len, None)
|
||||
.context("设置IP失败")?;
|
||||
guard.1 = Some((ip, prefix_len));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn create_tun(config: DeviceConfig) -> anyhow::Result<AsyncDevice> {
|
||||
#[cfg(unix)]
|
||||
if let Some(fd) = config.tun_fd {
|
||||
// SAFETY: Caller must ensure fd is a valid, open file descriptor for a TUN device.
|
||||
// Using an invalid fd may cause undefined behavior.
|
||||
unsafe { return Ok(AsyncDevice::from_fd(fd)?) }
|
||||
}
|
||||
let mut builder = DeviceBuilder::new();
|
||||
if let Some(tun_name) = config.tun_name {
|
||||
builder = builder.name(tun_name);
|
||||
}
|
||||
if let Some(mtu) = config.mtu {
|
||||
builder = builder.mtu(mtu);
|
||||
}
|
||||
#[cfg(windows)]
|
||||
{
|
||||
builder = builder.metric(1);
|
||||
}
|
||||
#[cfg(target_os = "linux")]
|
||||
{
|
||||
builder = builder.offload(true);
|
||||
}
|
||||
let dev = builder.build_async().context("创建tun失败")?;
|
||||
#[cfg(target_os = "linux")]
|
||||
{
|
||||
_ = dev.set_tx_queue_len(1000);
|
||||
}
|
||||
Ok(dev)
|
||||
}
|
||||
fn create(
|
||||
task_group: &TaskGroup,
|
||||
config: DeviceConfig,
|
||||
receiver: Receiver<TransmissionBytes>,
|
||||
enhanced_outbound: EnhancedOutbound,
|
||||
) -> anyhow::Result<DeviceTask> {
|
||||
let device = Arc::new(create_tun(config)?);
|
||||
|
||||
let device_framed_read = DeviceFramedRead::new(device.clone(), BytesCodec::new());
|
||||
let device_framed_write = DeviceFramedWrite::new(device.clone(), BytesCodec::new());
|
||||
|
||||
let task_recv = task_group.spawn(async move {
|
||||
if let Err(e) = in_tun_loop(receiver, device_framed_write).await {
|
||||
log::error!("in_tun_loop error: {e:?}")
|
||||
}
|
||||
});
|
||||
let task_send = task_group.spawn(async move {
|
||||
if let Err(e) = out_tun_loop(device_framed_read, enhanced_outbound).await {
|
||||
log::error!("out_tun_loop error: {e:?}");
|
||||
}
|
||||
});
|
||||
|
||||
Ok(DeviceTask {
|
||||
device,
|
||||
task_recv,
|
||||
task_send,
|
||||
})
|
||||
}
|
||||
|
||||
async fn in_tun_loop(
|
||||
mut receiver: Receiver<TransmissionBytes>,
|
||||
mut device_framed_write: DeviceFramedWrite<BytesCodec, Arc<AsyncDevice>>,
|
||||
) -> anyhow::Result<()> {
|
||||
while let Some(data) = receiver.recv().await {
|
||||
match device_framed_write.send(data).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
log::error!("send to tun error: {:?}", e);
|
||||
return Err(anyhow::anyhow!(e));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn out_tun_loop(
|
||||
mut device_framed_read: DeviceFramedRead<BytesCodec, Arc<AsyncDevice>>,
|
||||
enhanced_outbound: EnhancedOutbound,
|
||||
) -> anyhow::Result<()> {
|
||||
while let Some(rs) = device_framed_read.next().await {
|
||||
let bytes_mut = rs?;
|
||||
enhanced_outbound.ipv4_outbound(bytes_mut).await;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash, Default)]
|
||||
pub struct BytesCodec(());
|
||||
impl BytesCodec {
|
||||
pub fn new() -> BytesCodec {
|
||||
BytesCodec(())
|
||||
}
|
||||
}
|
||||
impl Decoder for BytesCodec {
|
||||
type Item = TransmissionBytes;
|
||||
type Error = io::Error;
|
||||
|
||||
fn decode(&mut self, buf: &mut BytesMut) -> Result<Option<TransmissionBytes>, io::Error> {
|
||||
if !buf.is_empty() {
|
||||
let mut bytes = TransmissionBytes::new_offset(HEAD_LENGTH);
|
||||
bytes.put(buf)?;
|
||||
buf.clear();
|
||||
Ok(Some(bytes))
|
||||
} else {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Encoder<TransmissionBytes> for BytesCodec {
|
||||
type Error = io::Error;
|
||||
|
||||
fn encode(&mut self, data: TransmissionBytes, buf: &mut BytesMut) -> Result<(), io::Error> {
|
||||
buf.reserve(data.len());
|
||||
buf.extend_from_slice(&data);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
mod general;
|
||||
pub use general::*;
|
||||
mod sender;
|
||||
pub use sender::*;
|
||||
|
||||
pub mod enhanced_tun;
|
||||
@@ -0,0 +1,40 @@
|
||||
use crate::context::NetworkAddr;
|
||||
use crate::nat::AllowSubnetExternalRoute;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tun::TunInbound;
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct TunDataInbound {
|
||||
allow_subnet: AllowSubnetExternalRoute,
|
||||
tun_inbound: TunInbound,
|
||||
}
|
||||
impl TunDataInbound {
|
||||
pub fn new(tun_inbound: TunInbound, allow_subnet: AllowSubnetExternalRoute) -> Self {
|
||||
Self {
|
||||
allow_subnet,
|
||||
tun_inbound,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl TunDataInbound {
|
||||
pub async fn send(&self, data: TransmissionBytes, net: &NetworkAddr) -> anyhow::Result<()> {
|
||||
if data[0] >> 4 != 4 {
|
||||
return Ok(());
|
||||
}
|
||||
let Some(ipv4) = Ipv4Packet::new(data.as_ref()) else {
|
||||
return Ok(());
|
||||
};
|
||||
let dest = ipv4.get_destination();
|
||||
if net.network().contains(&dest)
|
||||
|| dest == net.broadcast
|
||||
|| dest.is_broadcast()
|
||||
|| dest.is_multicast()
|
||||
|| self.allow_subnet.allow(&dest)
|
||||
{
|
||||
self.tun_inbound.sender.send(data).await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
pub(crate) mod p2p;
|
||||
pub mod server;
|
||||
|
||||
pub(crate) mod outbound;
|
||||
@@ -0,0 +1,286 @@
|
||||
use crate::compression::PacketCompression;
|
||||
use crate::context::{NetworkAddr, ServerInfoCollection, SharedNetworkAddr, TrafficStats};
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::fec::FecEncoder;
|
||||
use crate::nat::SubnetExternalRoute;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::p2p::outbound::P2pOutbound;
|
||||
use crate::tunnel_core::server::outbound::ServerOutbound;
|
||||
use anyhow::bail;
|
||||
use bytes::Bytes;
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct BasicOutbound {
|
||||
server_outbound: ServerOutbound,
|
||||
p2p_outbound: Option<P2pOutbound>,
|
||||
packet_crypto: PacketCrypto,
|
||||
}
|
||||
|
||||
impl BasicOutbound {
|
||||
pub fn new(
|
||||
server_outbound: ServerOutbound,
|
||||
p2p_outbound: Option<P2pOutbound>,
|
||||
packet_crypto: PacketCrypto,
|
||||
) -> Self {
|
||||
Self {
|
||||
server_outbound,
|
||||
p2p_outbound,
|
||||
packet_crypto,
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取加密保留空间大小
|
||||
pub fn encrypt_reserve(&self) -> usize {
|
||||
self.packet_crypto.encrypt_reserve()
|
||||
}
|
||||
|
||||
/// 加密数据包
|
||||
pub fn encrypt_in_place(
|
||||
&self,
|
||||
packet: &mut NetPacket<TransmissionBytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
Ok(self.packet_crypto.encrypt_in_place(packet)?)
|
||||
}
|
||||
|
||||
/// 发送原始数据包到指定目标(通过P2P或服务器)
|
||||
pub async fn send_raw(
|
||||
&self,
|
||||
dest: Ipv4Addr,
|
||||
packet: NetPacket<TransmissionBytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
let packet = packet.into_bytes();
|
||||
if let Some(p2p) = self.p2p_outbound.as_ref()
|
||||
&& let Some(route) = p2p.get_route_by_id(&dest)
|
||||
{
|
||||
p2p.send_raw_to(packet, &route.route_key()).await?;
|
||||
} else {
|
||||
self.server_outbound.send_raw(dest, packet).await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 发送到默认服务器
|
||||
pub async fn send_default_raw(
|
||||
&self,
|
||||
packet: NetPacket<TransmissionBytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
let bytes = packet.into_buffer().into_bytes().freeze();
|
||||
self.server_outbound
|
||||
.send_default_raw(NetPacket::new(bytes)?)
|
||||
.await
|
||||
}
|
||||
|
||||
/// 广播发送
|
||||
pub async fn send_raw_broadcast(
|
||||
&self,
|
||||
exclude_ips: Option<Vec<Ipv4Addr>>,
|
||||
packet: NetPacket<Bytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
self.server_outbound
|
||||
.send_raw_broadcast(exclude_ips, packet)
|
||||
.await
|
||||
}
|
||||
|
||||
/// 检查是否存在到目标的路由
|
||||
pub fn exists_route(&self, dest: &Ipv4Addr) -> bool {
|
||||
if let Some(p2p) = self.p2p_outbound.as_ref()
|
||||
&& p2p.exists_route_by_id(dest)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
self.server_outbound.exists_route(dest)
|
||||
}
|
||||
|
||||
/// P2P广播(内部转换类型)
|
||||
pub fn p2p_broadcast_transmission(
|
||||
&self,
|
||||
list: &[Ipv4Addr],
|
||||
max_count: usize,
|
||||
packet: &NetPacket<Bytes>,
|
||||
) -> Option<Vec<Ipv4Addr>> {
|
||||
if let Some(p2p) = self.p2p_outbound.as_ref() {
|
||||
let vec = p2p.p2p_broadcast(list, max_count, packet);
|
||||
if vec.is_empty() { None } else { Some(vec) }
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// 发送加密后的数据包
|
||||
pub async fn send_encrypted_packet(
|
||||
&self,
|
||||
dest: Ipv4Addr,
|
||||
mut packet: NetPacket<TransmissionBytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
// 加密
|
||||
self.packet_crypto.encrypt_in_place(&mut packet)?;
|
||||
|
||||
// 发送
|
||||
if let Some(p2p) = self.p2p_outbound.as_ref()
|
||||
&& let Some(route) = p2p.get_route_by_id(&dest)
|
||||
{
|
||||
let bytes = packet.into_buffer().into_bytes().freeze();
|
||||
p2p.send_raw_to(NetPacket::new(bytes)?, &route.route_key())
|
||||
.await?;
|
||||
} else {
|
||||
let bytes = packet.into_buffer().into_bytes().freeze();
|
||||
self.server_outbound
|
||||
.send_raw(dest, NetPacket::new(bytes)?)
|
||||
.await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct HybridOutbound {
|
||||
network: SharedNetworkAddr,
|
||||
server_info: ServerInfoCollection,
|
||||
traffic_stats: TrafficStats,
|
||||
basic_outbound: BasicOutbound,
|
||||
packet_compression: PacketCompression,
|
||||
external_route: SubnetExternalRoute,
|
||||
fec_encoder: Option<FecEncoder>,
|
||||
}
|
||||
impl HybridOutbound {
|
||||
pub fn new(
|
||||
network: SharedNetworkAddr,
|
||||
server_info: ServerInfoCollection,
|
||||
traffic_stats: TrafficStats,
|
||||
basic_outbound: BasicOutbound,
|
||||
packet_compression: PacketCompression,
|
||||
external_route: SubnetExternalRoute,
|
||||
fec_encoder: Option<FecEncoder>,
|
||||
) -> Self {
|
||||
Self {
|
||||
network,
|
||||
server_info,
|
||||
traffic_stats,
|
||||
basic_outbound,
|
||||
packet_compression,
|
||||
external_route,
|
||||
fec_encoder,
|
||||
}
|
||||
}
|
||||
pub async fn outbound_raw(
|
||||
&self,
|
||||
dest: Ipv4Addr,
|
||||
mut packet: NetPacket<TransmissionBytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
if packet.src_id() == 0 {
|
||||
if let Some(ip) = self.network.ip() {
|
||||
packet.set_src_id(ip.into());
|
||||
} else {
|
||||
bail!("Not src ip")
|
||||
}
|
||||
}
|
||||
|
||||
let len = packet.buffer().len() as u64;
|
||||
|
||||
if let Some(fec_encoder) = &self.fec_encoder {
|
||||
packet = fec_encoder.encode(packet)?;
|
||||
}
|
||||
|
||||
self.basic_outbound.send_raw(dest, packet).await?;
|
||||
self.traffic_stats.record_tx(dest, len);
|
||||
Ok(())
|
||||
}
|
||||
pub async fn ipv4_outbound_common(&self, data: TransmissionBytes) -> anyhow::Result<()> {
|
||||
let Some(net) = self.network.get() else {
|
||||
bail!("Not src ip")
|
||||
};
|
||||
self.ipv4_outbound(net, data).await
|
||||
}
|
||||
pub async fn ipv4_outbound(
|
||||
&self,
|
||||
net: NetworkAddr,
|
||||
mut data: TransmissionBytes,
|
||||
) -> anyhow::Result<()> {
|
||||
let Some(ipv4) = Ipv4Packet::new(data.as_ref()) else {
|
||||
return Ok(());
|
||||
};
|
||||
let mut dest = ipv4.get_destination();
|
||||
let len = data.len() as u64;
|
||||
data.retreat_head(HEAD_LENGTH)?;
|
||||
let mut packet = NetPacket::new(data)?;
|
||||
packet.set_msg_type(MsgType::Turn);
|
||||
packet.set_src_id(net.ip.into());
|
||||
packet.set_ttl(5);
|
||||
// 路由
|
||||
if !net.network().contains(&dest) {
|
||||
if let Some(v) = self.external_route.route(&dest) {
|
||||
dest = v;
|
||||
} else {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
packet.set_dest_id(dest.into());
|
||||
|
||||
packet = self
|
||||
.packet_compression
|
||||
.compress(packet, self.basic_outbound.encrypt_reserve())?;
|
||||
|
||||
if let Some(fec_encoder) = &self.fec_encoder {
|
||||
packet = fec_encoder.encode(packet)?;
|
||||
}
|
||||
|
||||
// 发送
|
||||
self.basic_outbound
|
||||
.send_encrypted_packet(dest, packet)
|
||||
.await?;
|
||||
self.traffic_stats.record_tx(dest, len);
|
||||
Ok(())
|
||||
}
|
||||
pub async fn ipv4_gateway_outbound(
|
||||
&self,
|
||||
net: NetworkAddr,
|
||||
mut data: TransmissionBytes,
|
||||
) -> anyhow::Result<()> {
|
||||
data.retreat_head(HEAD_LENGTH)?;
|
||||
let mut packet = NetPacket::new(data)?;
|
||||
packet.set_msg_type(MsgType::Turn);
|
||||
packet.set_src_id(net.ip.into());
|
||||
packet.set_dest_id(net.gateway.into());
|
||||
packet.set_ttl(5);
|
||||
packet.set_gateway_flag(true);
|
||||
self.basic_outbound.send_default_raw(packet).await?;
|
||||
Ok(())
|
||||
}
|
||||
pub async fn ipv4_broadcast_outbound(
|
||||
&self,
|
||||
net: NetworkAddr,
|
||||
mut data: TransmissionBytes,
|
||||
) -> anyhow::Result<()> {
|
||||
data.retreat_head(HEAD_LENGTH)?;
|
||||
let mut packet = NetPacket::new(data)?;
|
||||
packet.set_msg_type(MsgType::Broadcast);
|
||||
packet.set_src_id(net.ip.into());
|
||||
packet.set_dest_id(Ipv4Addr::BROADCAST.into());
|
||||
packet.set_ttl(5);
|
||||
let mut packet = self
|
||||
.packet_compression
|
||||
.compress(packet, self.basic_outbound.encrypt_reserve())?;
|
||||
self.basic_outbound.encrypt_in_place(&mut packet)?;
|
||||
let packet_bytes = packet.into_bytes();
|
||||
let list = self.server_info.client_online_ips();
|
||||
let exclude_ips = self
|
||||
.basic_outbound
|
||||
.p2p_broadcast_transmission(&list, 16, &packet_bytes);
|
||||
if let Some(exclude_ips) = &exclude_ips
|
||||
&& exclude_ips.len() == list.len()
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
self.basic_outbound
|
||||
.send_raw_broadcast(exclude_ips, packet_bytes)
|
||||
.await
|
||||
}
|
||||
#[allow(dead_code)]
|
||||
pub fn has_route(&self, dest: &Ipv4Addr) -> bool {
|
||||
self.basic_outbound.exists_route(dest)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,268 @@
|
||||
use crate::compression::PacketCompression;
|
||||
use crate::context::{NetworkRoute, PacketLossStats};
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::enhanced_tunnel::inbound::EnhancedInbound;
|
||||
use crate::fec::FecDecoder;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::p2p::outbound::P2pOutbound;
|
||||
use crate::tunnel_core::p2p::route_table::{Route, RouteTable};
|
||||
use anyhow::bail;
|
||||
use rust_p2p_core::route::RouteKey;
|
||||
use rust_p2p_core::tunnel::Tunnel;
|
||||
use std::net::{IpAddr, Ipv4Addr};
|
||||
|
||||
struct PacketContext {
|
||||
msg_type: MsgType,
|
||||
src_ip: Ipv4Addr,
|
||||
dest_ip: Ipv4Addr,
|
||||
max_ttl: u8,
|
||||
ttl: u8,
|
||||
}
|
||||
|
||||
pub(crate) struct P2pInboundConfig {
|
||||
pub network_route: NetworkRoute,
|
||||
pub route_table: RouteTable,
|
||||
pub packet_loss_stats: PacketLossStats,
|
||||
pub packet_crypto: PacketCrypto,
|
||||
pub packet_compression: PacketCompression,
|
||||
pub enhanced_inbound: EnhancedInbound,
|
||||
pub fec_decoder: FecDecoder,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct P2pInboundHandler {
|
||||
network_route: NetworkRoute,
|
||||
route_table: RouteTable,
|
||||
packet_loss_stats: PacketLossStats,
|
||||
packet_crypto: PacketCrypto,
|
||||
packet_compression: PacketCompression,
|
||||
enhanced_inbound: EnhancedInbound,
|
||||
fec_decoder: FecDecoder,
|
||||
}
|
||||
|
||||
impl P2pInboundHandler {
|
||||
pub fn new(config: P2pInboundConfig) -> Self {
|
||||
Self {
|
||||
network_route: config.network_route,
|
||||
route_table: config.route_table,
|
||||
packet_loss_stats: config.packet_loss_stats,
|
||||
packet_crypto: config.packet_crypto,
|
||||
packet_compression: config.packet_compression,
|
||||
enhanced_inbound: config.enhanced_inbound,
|
||||
fec_decoder: config.fec_decoder,
|
||||
}
|
||||
}
|
||||
fn network_contains(&self, ip: &Ipv4Addr) -> bool {
|
||||
self.network_route.network_contains(ip)
|
||||
}
|
||||
pub async fn next_handle(
|
||||
&self,
|
||||
buf: TransmissionBytes,
|
||||
route_key: RouteKey,
|
||||
p2p_socket_manager: &P2pOutbound,
|
||||
tunnel: &mut Tunnel,
|
||||
) {
|
||||
if let Err(e) = self
|
||||
.next_handle_impl(buf, route_key, p2p_socket_manager, tunnel)
|
||||
.await
|
||||
{
|
||||
log::warn!(
|
||||
"Error while handling P2pInboundHandler: {:?},route={route_key:?}",
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
async fn next_handle_impl(
|
||||
&self,
|
||||
buf: TransmissionBytes,
|
||||
route_key: RouteKey,
|
||||
p2p_socket_manager: &P2pOutbound,
|
||||
tunnel: &mut Tunnel,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut net_packet = NetPacket::new(buf)?;
|
||||
let msg_type = net_packet.msg_type()?;
|
||||
let src_ip = Ipv4Addr::from(net_packet.src_id());
|
||||
let dest_ip = Ipv4Addr::from(net_packet.dest_id());
|
||||
if src_ip == dest_ip {
|
||||
return Ok(());
|
||||
}
|
||||
net_packet.decr_ttl();
|
||||
|
||||
let max_ttl = net_packet.max_ttl();
|
||||
let ttl = net_packet.ttl();
|
||||
if max_ttl <= ttl {
|
||||
return Ok(());
|
||||
}
|
||||
let Some(net) = self.network_route.network.get() else {
|
||||
bail!("未找到自身IP")
|
||||
};
|
||||
if net.ip != dest_ip
|
||||
&& !dest_ip.is_broadcast()
|
||||
&& !dest_ip.is_unspecified()
|
||||
&& dest_ip != net.broadcast
|
||||
{
|
||||
// 帮忙转发数据包
|
||||
if ttl >= 1 {
|
||||
if let Some(route) = p2p_socket_manager.get_route_by_id(&dest_ip) {
|
||||
p2p_socket_manager
|
||||
.send_raw_to(net_packet.into_bytes(), &route.route_key())
|
||||
.await?;
|
||||
} else {
|
||||
log::debug!("未找到到 {} 的路由,无法转发", dest_ip);
|
||||
}
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if msg_type == MsgType::Quic {
|
||||
if net_packet.is_fec() {
|
||||
let packets = self.fec_decoder.receive(net_packet)?;
|
||||
if let Some(packets) = packets {
|
||||
for pkt in packets {
|
||||
self.enhanced_inbound
|
||||
.inbound(&net, msg_type, src_ip, pkt)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
self.enhanced_inbound
|
||||
.inbound(&net, msg_type, src_ip, net_packet)
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// 解密
|
||||
self.packet_crypto.decrypt_in_place(&mut net_packet)?;
|
||||
|
||||
let ctx = PacketContext {
|
||||
msg_type,
|
||||
src_ip,
|
||||
dest_ip,
|
||||
max_ttl,
|
||||
ttl,
|
||||
};
|
||||
|
||||
// FEC 解码(始终尝试解码,如果有 FEC 标志)
|
||||
if net_packet.is_fec() {
|
||||
let packets = self.fec_decoder.receive(net_packet)?;
|
||||
if let Some(packets) = packets {
|
||||
for pkt in packets {
|
||||
let pkt = self.packet_compression.decompress(pkt)?;
|
||||
self.process_decompressed_packet(&net, route_key, tunnel, pkt, &ctx)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// 解压缩
|
||||
let net_packet = self.packet_compression.decompress(net_packet)?;
|
||||
self.process_decompressed_packet(&net, route_key, tunnel, net_packet, &ctx)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn process_decompressed_packet(
|
||||
&self,
|
||||
net: &crate::context::NetworkAddr,
|
||||
route_key: RouteKey,
|
||||
tunnel: &mut Tunnel,
|
||||
net_packet: NetPacket<TransmissionBytes>,
|
||||
ctx: &PacketContext,
|
||||
) -> anyhow::Result<()> {
|
||||
match ctx.msg_type {
|
||||
MsgType::Turn | MsgType::Broadcast | MsgType::ExcludeBroadcast => {
|
||||
self.enhanced_inbound
|
||||
.inbound(net, ctx.msg_type, ctx.src_ip, net_packet)
|
||||
.await?;
|
||||
}
|
||||
MsgType::Ping => {
|
||||
let metric = ctx.max_ttl - ctx.ttl;
|
||||
self.route_table
|
||||
.add_route_if_absent(ctx.src_ip, Route::from_default_rt(route_key, metric));
|
||||
let mut packet = NetPacket::new(TransmissionBytes::zeroed_size(
|
||||
HEAD_LENGTH + 8,
|
||||
self.packet_crypto.encrypt_reserve(),
|
||||
))?;
|
||||
packet.set_msg_type(MsgType::Pong);
|
||||
packet.set_ttl(1);
|
||||
packet.set_src_id(ctx.dest_ip.into());
|
||||
packet.set_dest_id(ctx.src_ip.into());
|
||||
packet.set_payload(net_packet.payload())?;
|
||||
self.packet_crypto.encrypt_in_place(&mut packet)?;
|
||||
tunnel
|
||||
.send_to(packet.into_bytes().into_buffer(), route_key.addr())
|
||||
.await?;
|
||||
}
|
||||
MsgType::Pong => {
|
||||
if net_packet.payload().len() >= 8 {
|
||||
let metric = ctx.max_ttl - ctx.ttl;
|
||||
let time = i64::from_be_bytes(net_packet.payload()[..8].try_into()?);
|
||||
let now = crate::utils::time::now_ts_ms();
|
||||
if now >= time {
|
||||
self.route_table.add_route(
|
||||
ctx.src_ip,
|
||||
Route::from(route_key, metric, (now - time) as _),
|
||||
);
|
||||
self.packet_loss_stats.record_received(ctx.src_ip);
|
||||
}
|
||||
}
|
||||
}
|
||||
MsgType::PunchStart1 => {}
|
||||
MsgType::PunchStart2 => {}
|
||||
MsgType::PunchReq => {
|
||||
if let IpAddr::V4(ip) = route_key.addr().ip()
|
||||
&& self.network_contains(&ip)
|
||||
{
|
||||
log::info!("===========loop PunchReq {route_key:?} {:?}", ctx.src_ip);
|
||||
return Ok(());
|
||||
}
|
||||
log::info!(
|
||||
"PunchReq 打洞成功 {}->{},route={route_key:?}",
|
||||
ctx.src_ip,
|
||||
ctx.dest_ip
|
||||
);
|
||||
self.route_table.add_owner_route(ctx.src_ip, route_key);
|
||||
let mut packet = NetPacket::new(TransmissionBytes::zeroed_size(
|
||||
HEAD_LENGTH + 8,
|
||||
self.packet_crypto.encrypt_reserve(),
|
||||
))?;
|
||||
packet.set_msg_type(MsgType::PunchRes);
|
||||
packet.set_ttl(1);
|
||||
packet.set_src_id(ctx.dest_ip.into());
|
||||
packet.set_dest_id(ctx.src_ip.into());
|
||||
packet.set_payload(&crate::utils::time::now_ts_ms().to_be_bytes())?;
|
||||
|
||||
self.packet_crypto.encrypt_in_place(&mut packet)?;
|
||||
tunnel
|
||||
.send_to(packet.into_bytes().into_buffer(), route_key.addr())
|
||||
.await?;
|
||||
}
|
||||
MsgType::PunchRes => {
|
||||
if let IpAddr::V4(ip) = route_key.addr().ip()
|
||||
&& self.network_contains(&ip)
|
||||
{
|
||||
log::info!("===========loop PunchRes {route_key:?} {:?}", ctx.src_ip);
|
||||
return Ok(());
|
||||
}
|
||||
log::info!(
|
||||
"PunchRes 打洞成功 {}->{},route={route_key:?}",
|
||||
ctx.src_ip,
|
||||
ctx.dest_ip
|
||||
);
|
||||
self.route_table.add_owner_route(ctx.src_ip, route_key);
|
||||
}
|
||||
MsgType::PingTurn => {}
|
||||
MsgType::PongTurn => {}
|
||||
_ => {}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn tcp_disconnect(&self, route_key: RouteKey) {
|
||||
if let Some(ip) = self.route_table.get_id_by_route_key(&route_key) {
|
||||
self.route_table.remove_route(&ip, &route_key);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
pub(crate) mod inbound;
|
||||
pub(crate) mod outbound;
|
||||
pub(crate) mod transport;
|
||||
|
||||
pub(crate) mod route_table;
|
||||
@@ -0,0 +1,133 @@
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::protocol::ip_packet_protocol::NetPacket;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::p2p::route_table::{Route, RouteTable};
|
||||
use bytes::Bytes;
|
||||
use rust_p2p_core::route::RouteKey;
|
||||
use rust_p2p_core::tunnel::SocketManager;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct P2pOutbound {
|
||||
manager: SocketManager,
|
||||
route_table: RouteTable,
|
||||
packet_crypto: PacketCrypto,
|
||||
}
|
||||
impl P2pOutbound {
|
||||
pub fn new(
|
||||
manager: SocketManager,
|
||||
route_table: RouteTable,
|
||||
packet_crypto: PacketCrypto,
|
||||
) -> Self {
|
||||
Self {
|
||||
manager,
|
||||
route_table,
|
||||
packet_crypto,
|
||||
}
|
||||
}
|
||||
pub fn encrypt_reserve(&self) -> usize {
|
||||
self.packet_crypto.encrypt_reserve()
|
||||
}
|
||||
// pub async fn send_raw(&self, buf: NetPacket<Bytes>) -> anyhow::Result<()> {
|
||||
// let dest_id = Ipv4Addr::from(buf.dest_id());
|
||||
// let route = self.route_table.get_route_by_id(&dest_id)?;
|
||||
// self.manager
|
||||
// .send_to(buf.into_buffer(), &route.route_key())
|
||||
// .await?;
|
||||
// Ok(())
|
||||
// }
|
||||
// pub async fn send(&self, mut buf: NetPacket<TransmissionBytes>) -> anyhow::Result<()> {
|
||||
// let dest_id = Ipv4Addr::from(buf.dest_id());
|
||||
// let route = self.route_table.get_route_by_id(&dest_id)?;
|
||||
// self.packet_crypto.encrypt_in_place(&mut buf)?;
|
||||
// self.manager
|
||||
// .send_to(buf.into_buffer().into_bytes().freeze(), &route.route_key())
|
||||
// .await?;
|
||||
// Ok(())
|
||||
// }
|
||||
pub async fn send_raw_to(
|
||||
&self,
|
||||
buf: NetPacket<Bytes>,
|
||||
route_key: &RouteKey,
|
||||
) -> anyhow::Result<()> {
|
||||
self.manager.send_to(buf.into_buffer(), route_key).await?;
|
||||
Ok(())
|
||||
}
|
||||
pub async fn send_to(
|
||||
&self,
|
||||
mut buf: NetPacket<TransmissionBytes>,
|
||||
route_key: &RouteKey,
|
||||
) -> anyhow::Result<()> {
|
||||
self.packet_crypto.encrypt_in_place(&mut buf)?;
|
||||
self.manager
|
||||
.send_to(buf.into_buffer().into_bytes().freeze(), route_key)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
pub fn get_route_by_id(&self, id: &Ipv4Addr) -> Option<Route> {
|
||||
self.route_table.get_route_by_id(id).ok()
|
||||
}
|
||||
pub fn get_p2p_route_by_id(&self, id: &Ipv4Addr) -> Option<Route> {
|
||||
self.route_table.get_route_by_id(id).ok().filter(|v| v.is_direct())
|
||||
}
|
||||
pub fn exists_route_by_id(&self, id: &Ipv4Addr) -> bool {
|
||||
self.route_table.exists(id)
|
||||
}
|
||||
|
||||
// pub async fn send_to_id(
|
||||
// &self,
|
||||
// buf: NetPacket<TransmissionBytes>,
|
||||
// id: &Ipv4Addr,
|
||||
// ) -> anyhow::Result<bool> {
|
||||
// let Ok(route) = self.route_table.get_route_by_id(id) else {
|
||||
// return Ok(false);
|
||||
// };
|
||||
// self.send_to(buf, &route.route_key()).await?;
|
||||
// Ok(true)
|
||||
// }
|
||||
// pub fn try_send_to_id(
|
||||
// &self,
|
||||
// buf: NetPacket<TransmissionBytes>,
|
||||
// id: &Ipv4Addr,
|
||||
// ) -> anyhow::Result<bool> {
|
||||
// let Ok(route) = self.route_table.get_route_by_id(id) else {
|
||||
// return Ok(false);
|
||||
// };
|
||||
// self.try_send_to(buf, &route.route_key())?;
|
||||
// Ok(true)
|
||||
// }
|
||||
// pub fn try_send_to(
|
||||
// &self,
|
||||
// buf: NetPacket<TransmissionBytes>,
|
||||
// route_key: &RouteKey,
|
||||
// ) -> anyhow::Result<()> {
|
||||
// self.manager
|
||||
// .try_send_to(buf.into_buffer().into_bytes(), route_key)?;
|
||||
// Ok(())
|
||||
// }
|
||||
pub fn p2p_broadcast(
|
||||
&self,
|
||||
ips: &[Ipv4Addr],
|
||||
max: usize,
|
||||
buf: &NetPacket<Bytes>,
|
||||
) -> Vec<Ipv4Addr> {
|
||||
let mut list = Vec::with_capacity(ips.len().min(max));
|
||||
|
||||
for id in ips {
|
||||
let Some(route) = self.get_p2p_route_by_id(id) else {
|
||||
continue;
|
||||
};
|
||||
if self
|
||||
.manager
|
||||
.try_send_to(buf.source_buf().clone(), &route.route_key())
|
||||
.is_ok()
|
||||
{
|
||||
list.push(*id);
|
||||
if list.len() >= max {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
list
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,265 @@
|
||||
use parking_lot::{Mutex, RwLock};
|
||||
use rust_p2p_core::route::{RouteKey, DEFAULT_RTT};
|
||||
use std::collections::HashMap;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
use std::time::Instant;
|
||||
|
||||
|
||||
#[derive(Copy, Clone, Debug)]
|
||||
pub struct Route {
|
||||
route_key: RouteKey,
|
||||
metric: u8,
|
||||
rtt: u32,
|
||||
}
|
||||
impl Route {
|
||||
pub fn from(route_key: RouteKey, metric: u8, rtt: u32) -> Self {
|
||||
Self {
|
||||
route_key,
|
||||
metric,
|
||||
rtt,
|
||||
}
|
||||
}
|
||||
pub fn from_default_rt(route_key: RouteKey, metric: u8) -> Self {
|
||||
Self {
|
||||
route_key,
|
||||
metric,
|
||||
rtt: DEFAULT_RTT,
|
||||
}
|
||||
}
|
||||
pub fn route_key(&self) -> RouteKey {
|
||||
self.route_key
|
||||
}
|
||||
|
||||
pub fn is_direct(&self) -> bool {
|
||||
self.metric == 1
|
||||
}
|
||||
pub fn rtt(&self) -> u32 {
|
||||
self.rtt
|
||||
}
|
||||
pub fn metric(&self) -> u8 {
|
||||
self.metric
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RouteTable {
|
||||
inner: Arc<RouteTableInner>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct RouteTableInner {
|
||||
route_table: RwLock<HashMap<Ipv4Addr, Vec<Route>>>,
|
||||
route_key_time: Mutex<HashMap<(Ipv4Addr, RouteKey), Instant>>,
|
||||
route_key_owner: Mutex<HashMap<RouteKey, Ipv4Addr>>,
|
||||
}
|
||||
|
||||
impl Default for RouteTable {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl RouteTable {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
inner: Arc::new(RouteTableInner::default()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取指定 ID 的最优路由
|
||||
pub fn get_route_by_id(&self, id: &Ipv4Addr) -> anyhow::Result<Route> {
|
||||
self.inner
|
||||
.get_by_id(id)
|
||||
.ok_or_else(|| anyhow::anyhow!("route not found for {}", id))
|
||||
}
|
||||
|
||||
/// 检查是否存在到指定 ID 的路由
|
||||
pub fn exists(&self, id: &Ipv4Addr) -> bool {
|
||||
self.inner.get_by_id(id).is_some()
|
||||
}
|
||||
|
||||
/// 判断是否需要打洞(没有路由或只有中继路由)
|
||||
pub fn need_punch(&self, id: &Ipv4Addr) -> bool {
|
||||
let guard = self.inner.route_table.read();
|
||||
let Some(list) = guard.get(id) else {
|
||||
return true;
|
||||
};
|
||||
// 如果没有直连路由(metric=1),则需要打洞
|
||||
!list.iter().any(|r| r.is_direct())
|
||||
}
|
||||
|
||||
/// 获取直连路由数量(用于判断是否直连)
|
||||
pub fn p2p_num(&self, id: &Ipv4Addr) -> usize {
|
||||
let guard = self.inner.route_table.read();
|
||||
let Some(list) = guard.get(id) else {
|
||||
return 0;
|
||||
};
|
||||
list.iter().filter(|r| r.is_direct()).count()
|
||||
}
|
||||
|
||||
/// 添加 owner 路由(打洞请求响应时调用)
|
||||
pub fn add_owner_route(&self, id: Ipv4Addr, key: RouteKey) {
|
||||
self.inner.add_owner_route(id, key);
|
||||
}
|
||||
|
||||
/// 添加路由(心跳时调用,用于更新路由时间和添加跨节点转发路由)
|
||||
pub fn add_route(&self, id: Ipv4Addr, route: Route) {
|
||||
self.inner.add_route(id, route);
|
||||
}
|
||||
|
||||
/// 如果路由不存在则添加(用于 Ping 消息)
|
||||
pub fn add_route_if_absent(&self, id: Ipv4Addr, route: Route) {
|
||||
let guard = self.inner.route_table.read();
|
||||
if guard.contains_key(&id) {
|
||||
return;
|
||||
}
|
||||
drop(guard);
|
||||
self.inner.add_route(id, route);
|
||||
}
|
||||
|
||||
/// 获取所有路由表
|
||||
pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec<Route>)> {
|
||||
let guard = self.inner.route_table.read();
|
||||
guard.iter().map(|(k, v)| (*k, v.clone())).collect()
|
||||
}
|
||||
|
||||
/// 根据 RouteKey 查找对应的 IP
|
||||
pub fn get_id_by_route_key(&self, route_key: &RouteKey) -> Option<Ipv4Addr> {
|
||||
let owner_map = self.inner.route_key_owner.lock();
|
||||
owner_map.get(route_key).copied()
|
||||
}
|
||||
|
||||
/// 移除指定 IP 和 RouteKey 的路由
|
||||
pub fn remove_route(&self, id: &Ipv4Addr, route_key: &RouteKey) {
|
||||
let mut table = self.inner.route_table.write();
|
||||
let mut owner_map = self.inner.route_key_owner.lock();
|
||||
let mut time_map = self.inner.route_key_time.lock();
|
||||
|
||||
if let Some(list) = table.get_mut(id) {
|
||||
list.retain(|r| r.route_key() != *route_key);
|
||||
if list.is_empty() {
|
||||
table.remove(id);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(owner_id) = owner_map.get(route_key) {
|
||||
if owner_id == id {
|
||||
owner_map.remove(route_key);
|
||||
}
|
||||
}
|
||||
|
||||
time_map.remove(&(*id, *route_key));
|
||||
}
|
||||
|
||||
/// 移除过期的路由
|
||||
pub fn remove_oldest_route(&self, expired_time: Instant) {
|
||||
self.inner.remove_oldest_route(expired_time);
|
||||
}
|
||||
}
|
||||
|
||||
impl RouteTableInner {
|
||||
fn get_by_id(&self, id: &Ipv4Addr) -> Option<Route> {
|
||||
let guard = self.route_table.read();
|
||||
let list = guard.get(id)?;
|
||||
list.first().cloned()
|
||||
}
|
||||
|
||||
fn add_owner_route(&self, id: Ipv4Addr, key: RouteKey) {
|
||||
let route = Route::from_default_rt(key, 1);
|
||||
let mut guard = self.route_table.write();
|
||||
|
||||
self.route_key_owner.lock().insert(key, id);
|
||||
self.route_key_time.lock().insert((id, key), Instant::now());
|
||||
|
||||
let list = guard.entry(id).or_insert_with(|| Vec::with_capacity(6));
|
||||
if list.iter().any(|v| v.route_key() == key) {
|
||||
return;
|
||||
}
|
||||
list.push(route);
|
||||
}
|
||||
|
||||
fn add_route(&self, id: Ipv4Addr, route: Route) {
|
||||
let key = route.route_key();
|
||||
let mut guard = self.route_table.write();
|
||||
|
||||
// 检查是否是 owner 路由
|
||||
let mut route_key_owner = self.route_key_owner.lock();
|
||||
if route.is_direct() {
|
||||
route_key_owner.entry(key).or_insert(id);
|
||||
} else {
|
||||
if !guard.contains_key(&id) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
// 更新时间
|
||||
self.route_key_time.lock().insert((id, key), Instant::now());
|
||||
|
||||
let list = guard.entry(id).or_insert_with(|| Vec::with_capacity(6));
|
||||
|
||||
// 如果路由已存��,更新并重新排序
|
||||
if let Some(idx) = list.iter().position(|v| v.route_key() == key) {
|
||||
list[idx] = route;
|
||||
// 向前冒泡(如果 RTT 更小)
|
||||
let mut i = idx;
|
||||
while i > 0 && list[i].rtt() < list[i - 1].rtt() {
|
||||
list.swap(i, i - 1);
|
||||
i -= 1;
|
||||
}
|
||||
// 向后冒泡(如果 RTT 更大)
|
||||
while i + 1 < list.len() && list[i].rtt() > list[i + 1].rtt() {
|
||||
list.swap(i, i + 1);
|
||||
i += 1;
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
// 插入新路由,保持按 RTT 排序
|
||||
let mut pos = list.len();
|
||||
for (i, r) in list.iter().enumerate() {
|
||||
if route.rtt() < r.rtt() {
|
||||
pos = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
list.insert(pos, route);
|
||||
}
|
||||
|
||||
fn remove_oldest_route(&self, expired_time: Instant) {
|
||||
let mut expired_keys = Vec::new();
|
||||
{
|
||||
let mut time_map = self.route_key_time.lock();
|
||||
time_map.retain(|(id, route_key), t| {
|
||||
if *t <= expired_time {
|
||||
expired_keys.push((*id, *route_key));
|
||||
false
|
||||
} else {
|
||||
true
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
if expired_keys.is_empty() {
|
||||
return;
|
||||
}
|
||||
|
||||
let mut table = self.route_table.write();
|
||||
let mut owner_map = self.route_key_owner.lock();
|
||||
|
||||
for (id, route_key) in expired_keys {
|
||||
if let Some(list) = table.get_mut(&id) {
|
||||
list.retain(|r| r.route_key() != route_key);
|
||||
if list.is_empty() {
|
||||
table.remove(&id);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(owner_id) = owner_map.get(&route_key) {
|
||||
if *owner_id == id {
|
||||
owner_map.remove(&route_key);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
pub(crate) mod nat_test;
|
||||
pub(crate) mod punch;
|
||||
|
||||
pub(crate) mod task;
|
||||
@@ -0,0 +1,387 @@
|
||||
use crate::context::AppState;
|
||||
use rust_p2p_core::nat::{NatInfo, NatType};
|
||||
use rust_p2p_core::tunnel::SocketManager;
|
||||
use rust_p2p_core::tunnel::udp::Model;
|
||||
use std::collections::HashMap;
|
||||
use std::io;
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
|
||||
use std::time::Duration;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpStream;
|
||||
|
||||
pub async fn my_nat_info(app_context: AppState, socket_manager: SocketManager) {
|
||||
loop {
|
||||
my_nat_info_impl(&app_context, &socket_manager).await;
|
||||
tokio::time::sleep(Duration::from_secs(60 * 30)).await;
|
||||
}
|
||||
}
|
||||
async fn my_nat_info_impl(app_context: &AppState, socket_manager: &SocketManager) {
|
||||
let network = app_context.network.network();
|
||||
let mut local_ipv4s = Vec::new();
|
||||
let mut local_ipv6 = Vec::new();
|
||||
match getifaddrs::getifaddrs() {
|
||||
Ok(addrs) => {
|
||||
for x in addrs {
|
||||
let Some(ip) = x.address.ip_addr() else {
|
||||
continue;
|
||||
};
|
||||
if ip.is_loopback() {
|
||||
continue;
|
||||
}
|
||||
if ip.is_unspecified() {
|
||||
continue;
|
||||
}
|
||||
if ip.is_multicast() {
|
||||
continue;
|
||||
}
|
||||
|
||||
match ip {
|
||||
IpAddr::V4(addr) => {
|
||||
if addr.is_documentation() {
|
||||
continue;
|
||||
}
|
||||
if addr.is_broadcast() {
|
||||
continue;
|
||||
}
|
||||
if let Some(network) = &network
|
||||
&& network.contains(&addr)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
local_ipv4s.push(addr);
|
||||
}
|
||||
IpAddr::V6(addr) => {
|
||||
if addr.is_unique_local() {
|
||||
continue;
|
||||
}
|
||||
if addr.is_unicast_link_local() {
|
||||
continue;
|
||||
}
|
||||
local_ipv6.push(addr);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
log::error!("getifaddrs error: {e}");
|
||||
}
|
||||
}
|
||||
log::info!("local_ipv4s: {:?}", local_ipv4s);
|
||||
let local_ipv4 = rust_p2p_core::extend::addr::local_ipv4()
|
||||
.await
|
||||
.unwrap_or_else(|e| {
|
||||
log::warn!("local ipv4 failed {e:?}");
|
||||
local_ipv4s
|
||||
.first()
|
||||
.cloned()
|
||||
.unwrap_or(Ipv4Addr::UNSPECIFIED)
|
||||
});
|
||||
local_ipv4s = vec![local_ipv4];
|
||||
let mut ipv6 = rust_p2p_core::extend::addr::local_ipv6().await.ok();
|
||||
if let Some(addr) = ipv6 {
|
||||
if addr.is_loopback()
|
||||
|| addr.is_unique_local()
|
||||
|| addr.is_unicast_link_local()
|
||||
|| addr.is_unspecified()
|
||||
|| addr.is_multicast()
|
||||
{
|
||||
ipv6 = local_ipv6.first().cloned();
|
||||
}
|
||||
} else {
|
||||
ipv6 = local_ipv6.first().cloned();
|
||||
}
|
||||
let local_udp_ports = socket_manager
|
||||
.udp_socket_manager_as_ref()
|
||||
.unwrap()
|
||||
.local_ports()
|
||||
.unwrap();
|
||||
let local_tcp_port = socket_manager
|
||||
.tcp_socket_manager_as_ref()
|
||||
.unwrap()
|
||||
.local_addr()
|
||||
.port();
|
||||
log::info!(
|
||||
"local_ipv4={local_ipv4},ipv6={ipv6:?},local_udp_ports:{local_udp_ports:?},local_tcp_port:{local_tcp_port:?}"
|
||||
);
|
||||
let mut public_ports = local_udp_ports.clone();
|
||||
public_ports.fill(0);
|
||||
let mut nat_info = NatInfo {
|
||||
nat_type: NatType::Cone,
|
||||
public_ips: vec![],
|
||||
public_udp_ports: public_ports,
|
||||
mapping_tcp_addr: vec![],
|
||||
mapping_udp_addr: vec![],
|
||||
public_port_range: 0,
|
||||
local_ipv4s,
|
||||
local_ipv4,
|
||||
ipv6,
|
||||
local_udp_ports,
|
||||
local_tcp_port,
|
||||
public_tcp_port: 0,
|
||||
};
|
||||
let mut stun_server = app_context.udp_stun();
|
||||
if stun_server.is_empty() {
|
||||
stun_server = default_udp_stun();
|
||||
}
|
||||
let (nat_type, public_ips, port_range) = rust_p2p_core::stun::stun_test_nat(stun_server, None)
|
||||
.await
|
||||
.unwrap_or_else(|e| {
|
||||
log::warn!("stun_test_nat {e:?}");
|
||||
(NatType::Cone, vec![], 0)
|
||||
});
|
||||
log::info!("nat_type:{nat_type:?},public_ips:{public_ips:?},port_range={port_range}");
|
||||
nat_info.nat_type = nat_type;
|
||||
nat_info.public_ips = public_ips;
|
||||
nat_info.public_port_range = port_range;
|
||||
app_context.nat_info.replace_nat_info(nat_info);
|
||||
let model = match nat_type {
|
||||
NatType::Cone => Model::Low,
|
||||
NatType::Symmetric => Model::High,
|
||||
};
|
||||
if let Err(e) = socket_manager
|
||||
.udp_socket_manager_as_ref()
|
||||
.unwrap()
|
||||
.switch_model(model)
|
||||
{
|
||||
log::error!("switch_model error: {e:?}");
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn query_udp_public_addr_loop(app_context: AppState, socket_manager: SocketManager) {
|
||||
let mut udp_stun_servers = app_context.udp_stun();
|
||||
if udp_stun_servers.is_empty() {
|
||||
udp_stun_servers = default_udp_stun();
|
||||
}
|
||||
let udp_len = udp_stun_servers.len();
|
||||
let mut udp_count = 0;
|
||||
let stun_request = rust_p2p_core::stun::send_stun_request();
|
||||
loop {
|
||||
let stun = &udp_stun_servers[udp_count % udp_len];
|
||||
udp_count += 1;
|
||||
match tokio::net::lookup_host(stun.as_str()).await {
|
||||
Ok(mut addr) => {
|
||||
if let Some(addr) = addr.next()
|
||||
&& let Some(w) = socket_manager.udp_socket_manager_as_ref()
|
||||
&& let Err(e) = w.detect_pub_addrs(&stun_request, addr).await
|
||||
{
|
||||
log::info!("detect_pub_addrs {e:?} {addr:?}");
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
log::info!("query_public_addr lookup_host {e:?} {stun:?}",);
|
||||
}
|
||||
}
|
||||
let not_port = app_context
|
||||
.get_nat_info()
|
||||
.map(|v| v.public_udp_ports.contains(&0))
|
||||
.unwrap_or(true);
|
||||
if not_port {
|
||||
tokio::time::sleep(Duration::from_secs(2)).await;
|
||||
} else {
|
||||
tokio::time::sleep(Duration::from_secs(60)).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn query_tcp_public_addr_loop(
|
||||
app_context: AppState,
|
||||
socket_manager: SocketManager,
|
||||
) {
|
||||
use rand::Rng;
|
||||
use rand::seq::SliceRandom;
|
||||
|
||||
let tcp_stun_servers = {
|
||||
let servers = app_context.tcp_stun();
|
||||
if servers.is_empty() {
|
||||
default_tcp_stun()
|
||||
} else {
|
||||
servers
|
||||
}
|
||||
};
|
||||
|
||||
if tcp_stun_servers.is_empty() {
|
||||
return;
|
||||
}
|
||||
log::debug!("tcp_stun_servers = {tcp_stun_servers:?}");
|
||||
|
||||
let stun_request = rust_p2p_core::stun::send_stun_request();
|
||||
let target_conn_count = tcp_stun_servers.len().min(2);
|
||||
let mut active_connections: HashMap<SocketAddr, (TcpStream, SocketAddr)> = HashMap::new();
|
||||
|
||||
'outer: loop {
|
||||
while active_connections.len() < target_conn_count {
|
||||
let mut candidates: Vec<&String> = tcp_stun_servers.iter().collect();
|
||||
candidates.shuffle(&mut rand::rng());
|
||||
|
||||
let mut connected = false;
|
||||
for stun in candidates {
|
||||
let addr = match tokio::net::lookup_host(stun.as_str()).await {
|
||||
Ok(mut addrs) => addrs.next(),
|
||||
Err(e) => {
|
||||
log::debug!("lookup_host failed {stun} {e}");
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let Some(addr) = addr else {
|
||||
continue;
|
||||
};
|
||||
|
||||
if active_connections.contains_key(&addr) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let Some(w) = socket_manager.tcp_socket_manager_as_ref() else {
|
||||
continue;
|
||||
};
|
||||
|
||||
match tokio::time::timeout(Duration::from_secs(5), w.connect_reuse_port_raw(addr))
|
||||
.await
|
||||
{
|
||||
Ok(Ok(mut tcp_stream)) => {
|
||||
let write_result = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
tcp_stream.write_all(&stun_request),
|
||||
)
|
||||
.await;
|
||||
|
||||
if let Ok(Ok(_)) = write_result {
|
||||
match stun_tcp_read(&mut tcp_stream).await {
|
||||
Ok(pub_addr) => {
|
||||
log::debug!(
|
||||
"update_tcp_public_addr {stun} {addr} -> {pub_addr}"
|
||||
);
|
||||
|
||||
let existing_pub_addr =
|
||||
active_connections.values().next().map(|(_, p)| *p);
|
||||
|
||||
if let Some(existing) = existing_pub_addr
|
||||
&& existing != pub_addr
|
||||
{
|
||||
log::debug!(
|
||||
"pub_addr mismatch: {existing} != {pub_addr}, wait 60s"
|
||||
);
|
||||
active_connections.clear();
|
||||
app_context.nat_info.update_tcp_public_addr(
|
||||
SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into(),
|
||||
);
|
||||
tokio::time::sleep(Duration::from_secs(5 * 60)).await;
|
||||
continue 'outer;
|
||||
}
|
||||
|
||||
active_connections.insert(addr, (tcp_stream, pub_addr));
|
||||
connected = true;
|
||||
break;
|
||||
}
|
||||
Err(e) => {
|
||||
log::debug!("stun_tcp_read failed {stun} {addr} {e}");
|
||||
}
|
||||
}
|
||||
} else {
|
||||
log::debug!("write stun request failed {stun} {addr}");
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
log::debug!("connect_reuse_port_raw failed {stun} {addr} {e}");
|
||||
}
|
||||
Err(_) => {
|
||||
log::debug!("connect_reuse_port_raw timeout {stun} {addr}");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !connected {
|
||||
break;
|
||||
}
|
||||
}
|
||||
let existing_pub_addr = active_connections.values().next().map(|(_, p)| *p);
|
||||
if let Some(existing) = existing_pub_addr {
|
||||
app_context.nat_info.update_tcp_public_addr(existing);
|
||||
}
|
||||
|
||||
let sleep_secs = rand::rng().random_range(10u64..=15);
|
||||
tokio::time::sleep(Duration::from_secs(sleep_secs)).await;
|
||||
|
||||
let mut to_remove = Vec::new();
|
||||
let addrs: Vec<SocketAddr> = active_connections.keys().cloned().collect();
|
||||
|
||||
for addr in addrs {
|
||||
let (tcp_stream, _) = active_connections.get_mut(&addr).unwrap();
|
||||
let mut buf = [0u8; 1024];
|
||||
|
||||
match tcp_stream.try_read(&mut buf) {
|
||||
Ok(0) => {
|
||||
log::warn!("stun tcp close {addr} EOF");
|
||||
to_remove.push(addr);
|
||||
continue;
|
||||
}
|
||||
Err(e) if e.kind() != std::io::ErrorKind::WouldBlock => {
|
||||
log::warn!("stun tcp read error {addr} {e}");
|
||||
to_remove.push(addr);
|
||||
continue;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
match tokio::time::timeout(Duration::from_secs(3), tcp_stream.write_all(&stun_request))
|
||||
.await
|
||||
{
|
||||
Ok(Ok(_)) => {}
|
||||
Ok(Err(e)) => {
|
||||
log::warn!("stun tcp write error {addr} {e}");
|
||||
to_remove.push(addr);
|
||||
}
|
||||
Err(_) => {
|
||||
log::warn!("stun tcp write timeout {addr}");
|
||||
to_remove.push(addr);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for addr in to_remove {
|
||||
active_connections.remove(&addr);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn stun_tcp_read(tcp_stream: &mut TcpStream) -> io::Result<SocketAddr> {
|
||||
let mut head = [0; 20];
|
||||
match tokio::time::timeout(Duration::from_secs(5), tcp_stream.read_exact(&mut head)).await {
|
||||
Ok(rs) => rs?,
|
||||
Err(_) => Err(io::Error::from(io::ErrorKind::TimedOut))?,
|
||||
};
|
||||
let len = u16::from_be_bytes([head[2], head[3]]) as usize;
|
||||
let mut buf = vec![0; len + 20];
|
||||
buf[..20].copy_from_slice(&head);
|
||||
match tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
tcp_stream.read_exact(&mut buf[20..]),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(rs) => rs?,
|
||||
Err(_) => Err(io::Error::from(io::ErrorKind::TimedOut))?,
|
||||
};
|
||||
if let Some(addr) = rust_p2p_core::stun::recv_stun_response(&buf) {
|
||||
Ok(addr)
|
||||
} else {
|
||||
log::debug!("stun_tcp_read {buf:?}");
|
||||
Err(io::Error::from(io::ErrorKind::InvalidData))
|
||||
}
|
||||
}
|
||||
|
||||
fn default_udp_stun() -> Vec<String> {
|
||||
vec![
|
||||
"stun.miwifi.com:3478".to_string(),
|
||||
"stun.chat.bilibili.com:3478".to_string(),
|
||||
"stun.l.google.com:19302".to_string(),
|
||||
]
|
||||
}
|
||||
|
||||
fn default_tcp_stun() -> Vec<String> {
|
||||
vec![
|
||||
"stun.flashdance.cx:3478".to_string(),
|
||||
"stun.sipnet.net:3478".to_string(),
|
||||
"stun.nextcloud.com:443".to_string(),
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
use crate::context::nat::PunchBackoff;
|
||||
use crate::context::{ServerInfoCollection, SharedNetworkAddr};
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::protocol::client_message::PunchInfo;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::p2p::route_table::RouteTable;
|
||||
use crate::tunnel_core::server::outbound::ServerOutbound;
|
||||
use anyhow::bail;
|
||||
use log::error;
|
||||
use rand::seq::SliceRandom;
|
||||
use rust_p2p_core::punch::{PunchModel, Puncher};
|
||||
use std::net::Ipv4Addr;
|
||||
use std::time::Duration;
|
||||
|
||||
pub struct PunchTaskContext {
|
||||
pub network: SharedNetworkAddr,
|
||||
pub server_info: ServerInfoCollection,
|
||||
pub punch_backoff: PunchBackoff,
|
||||
pub punch_info_getter: PunchInfoGetter,
|
||||
}
|
||||
|
||||
pub type PunchInfoGetter = std::sync::Arc<dyn Fn() -> Option<PunchInfo> + Send + Sync>;
|
||||
|
||||
pub async fn punch_task(
|
||||
tunnel_to_server: ServerOutbound,
|
||||
route_table: RouteTable,
|
||||
ctx: PunchTaskContext,
|
||||
) -> anyhow::Result<()> {
|
||||
loop {
|
||||
tokio::time::sleep(Duration::from_secs(5)).await;
|
||||
let Some(src_ip) = ctx.network.ip() else {
|
||||
continue;
|
||||
};
|
||||
let Some(punch_info) = (ctx.punch_info_getter)() else {
|
||||
continue;
|
||||
};
|
||||
let mut list = ctx.server_info.client_online_ips();
|
||||
list.shuffle(&mut rand::rng());
|
||||
list.truncate(5);
|
||||
for dest_ip in list {
|
||||
if dest_ip <= src_ip {
|
||||
continue;
|
||||
}
|
||||
if ctx.server_info.is_any_server_connected(None) && route_table.need_punch(&dest_ip) {
|
||||
if !ctx.punch_backoff.should_punch(dest_ip) {
|
||||
continue;
|
||||
}
|
||||
log::info!("punching {dest_ip}");
|
||||
|
||||
let data = punch_info.encode();
|
||||
let mut net_packet = NetPacket::new(TransmissionBytes::zeroed_size(
|
||||
HEAD_LENGTH + data.len(),
|
||||
tunnel_to_server.encrypt_reserve(),
|
||||
))?;
|
||||
net_packet.set_msg_type(MsgType::PunchStart1);
|
||||
net_packet.set_ttl(2);
|
||||
net_packet.set_src_id(src_ip.into());
|
||||
net_packet.set_dest_id(dest_ip.into());
|
||||
net_packet.set_payload(data.as_ref())?;
|
||||
if let Err(e) = tunnel_to_server.send(dest_ip, net_packet).await {
|
||||
error!("punch send error {:?}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
#[derive(Clone)]
|
||||
pub struct NatPuncher {
|
||||
network: SharedNetworkAddr,
|
||||
punch_backoff: PunchBackoff,
|
||||
puncher: Option<Puncher>,
|
||||
packet_crypto: PacketCrypto,
|
||||
}
|
||||
|
||||
impl NatPuncher {
|
||||
pub fn new(
|
||||
network: SharedNetworkAddr,
|
||||
punch_backoff: PunchBackoff,
|
||||
puncher: Option<Puncher>,
|
||||
packet_crypto: PacketCrypto,
|
||||
) -> Self {
|
||||
Self {
|
||||
network,
|
||||
punch_backoff,
|
||||
puncher,
|
||||
packet_crypto,
|
||||
}
|
||||
}
|
||||
pub fn punch(&self, dest_ip: Ipv4Addr, punch_info: PunchInfo) -> anyhow::Result<bool> {
|
||||
if self.puncher.is_none() {
|
||||
return Ok(false);
|
||||
}
|
||||
if !self.punch_backoff.should_punch(dest_ip) {
|
||||
return Ok(false);
|
||||
}
|
||||
self.punch_uncheck_delay(dest_ip, punch_info, Some(Duration::from_millis(50)))?;
|
||||
Ok(true)
|
||||
}
|
||||
pub fn punch_uncheck(&self, dest_ip: Ipv4Addr, punch_info: PunchInfo) -> anyhow::Result<()> {
|
||||
self.punch_uncheck_delay(dest_ip, punch_info, None)
|
||||
}
|
||||
pub fn punch_uncheck_delay(
|
||||
&self,
|
||||
dest_ip: Ipv4Addr,
|
||||
punch_info: PunchInfo,
|
||||
time: Option<Duration>,
|
||||
) -> anyhow::Result<()> {
|
||||
let Some(puncher) = self.puncher.clone() else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(src_ip) = self.network.ip() else {
|
||||
bail!("not ip");
|
||||
};
|
||||
let packet_crypto = self.packet_crypto.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Some(time) = time {
|
||||
tokio::time::sleep(time).await;
|
||||
}
|
||||
if let Err(e) = punch_now(puncher, src_ip, dest_ip, punch_info, packet_crypto).await {
|
||||
log::warn!("punch send error {:?}", e);
|
||||
}
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
async fn punch_now(
|
||||
puncher: Puncher,
|
||||
src_ip: Ipv4Addr,
|
||||
dest_ip: Ipv4Addr,
|
||||
nat_info: PunchInfo,
|
||||
packet_crypto: PacketCrypto,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut packet = NetPacket::new(TransmissionBytes::zeroed_size(
|
||||
HEAD_LENGTH + 8,
|
||||
packet_crypto.encrypt_reserve(),
|
||||
))?;
|
||||
packet.set_msg_type(MsgType::PunchReq);
|
||||
packet.set_ttl(1);
|
||||
packet.set_src_id(src_ip.into());
|
||||
packet.set_dest_id(dest_ip.into());
|
||||
packet.set_payload(&crate::utils::time::now_ts_ms().to_be_bytes())?;
|
||||
packet_crypto.encrypt_in_place(&mut packet)?;
|
||||
let buf = packet.buffer();
|
||||
let punch_info = rust_p2p_core::punch::PunchInfo::new(PunchModel::all(), nat_info.nat_info);
|
||||
puncher.punch_now(Some(buf), buf, punch_info).await?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
use crate::context::nat::MyNatInfo;
|
||||
use crate::context::{AppState, PacketLossStats, SharedNetworkAddr};
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::p2p::inbound::P2pInboundHandler;
|
||||
use crate::tunnel_core::p2p::outbound::P2pOutbound;
|
||||
use crate::tunnel_core::p2p::route_table::RouteTable;
|
||||
use crate::tunnel_core::p2p::transport::nat_test::{
|
||||
my_nat_info, query_tcp_public_addr_loop, query_udp_public_addr_loop,
|
||||
};
|
||||
use crate::tunnel_core::p2p::transport::punch::{PunchTaskContext, punch_task};
|
||||
use crate::tunnel_core::server::outbound::ServerOutbound;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use rust_p2p_core::punch::Puncher;
|
||||
use rust_p2p_core::tunnel::{Tunnel, TunnelDispatcher, new_tunnel_component};
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
pub async fn init_tunnel(
|
||||
task_group: TaskGroup,
|
||||
app_state: AppState,
|
||||
tunnel_to_server: ServerOutbound,
|
||||
packet_crypto: PacketCrypto,
|
||||
) -> anyhow::Result<(Puncher, P2pOutbound, P2pTask)> {
|
||||
let udp_config = rust_p2p_core::tunnel::config::UdpTunnelConfig::default()
|
||||
.set_main_udp_count(2)
|
||||
.set_sub_udp_count(82);
|
||||
let tcp_config = rust_p2p_core::tunnel::config::TcpTunnelConfig::new(Box::new(
|
||||
rust_p2p_core::tunnel::tcp::LengthPrefixedInitCodec,
|
||||
))
|
||||
.set_tcp_multiplexing_limit(2);
|
||||
let config = rust_p2p_core::tunnel::config::TunnelConfig::empty()
|
||||
.set_udp_tunnel_config(udp_config)
|
||||
.set_tcp_tunnel_config(tcp_config);
|
||||
let (tunnel_dispatcher, puncher) = new_tunnel_component(config)?;
|
||||
let route_table = app_state.route_table.clone();
|
||||
let socket_manager = P2pOutbound::new(
|
||||
tunnel_dispatcher.socket_manager(),
|
||||
route_table.clone(),
|
||||
packet_crypto,
|
||||
);
|
||||
task_group.spawn(my_nat_info(
|
||||
app_state.clone(),
|
||||
tunnel_dispatcher.socket_manager(),
|
||||
));
|
||||
let manager = tunnel_dispatcher.socket_manager();
|
||||
task_group.spawn(query_udp_public_addr_loop(
|
||||
app_state.clone(),
|
||||
manager.clone(),
|
||||
));
|
||||
task_group.spawn(query_tcp_public_addr_loop(app_state.clone(), manager));
|
||||
|
||||
task_group.spawn(route_timeout_task(route_table.clone()));
|
||||
let app_state_for_punch = app_state.clone();
|
||||
let punch_ctx = PunchTaskContext {
|
||||
network: app_state.network.clone(),
|
||||
server_info: app_state.server_info_collection.clone(),
|
||||
punch_backoff: app_state.punch_backoff.clone(),
|
||||
punch_info_getter: Arc::new(move || app_state_for_punch.get_punch_info()),
|
||||
};
|
||||
task_group.spawn(punch_task(tunnel_to_server, route_table.clone(), punch_ctx));
|
||||
task_group.spawn(ping_all(
|
||||
app_state.network.clone(),
|
||||
app_state.packet_loss_stats.clone(),
|
||||
route_table.clone(),
|
||||
socket_manager.clone(),
|
||||
));
|
||||
task_group.spawn(relay_probe_task(
|
||||
app_state.network.clone(),
|
||||
app_state.server_info_collection.clone(),
|
||||
route_table.clone(),
|
||||
socket_manager.clone(),
|
||||
));
|
||||
let p2p_task = P2pTask {
|
||||
task_group,
|
||||
nat_info: app_state.nat_info.clone(),
|
||||
socket_manager: socket_manager.clone(),
|
||||
tunnel_dispatcher,
|
||||
};
|
||||
Ok((puncher, socket_manager, p2p_task))
|
||||
}
|
||||
pub struct P2pTask {
|
||||
task_group: TaskGroup,
|
||||
nat_info: MyNatInfo,
|
||||
socket_manager: P2pOutbound,
|
||||
tunnel_dispatcher: TunnelDispatcher,
|
||||
}
|
||||
impl P2pTask {
|
||||
pub fn start(self, p2p_inbound_handler: P2pInboundHandler) {
|
||||
self.task_group.spawn(tunnel_dispatch_task(
|
||||
self.nat_info,
|
||||
self.task_group.clone(),
|
||||
self.tunnel_dispatcher,
|
||||
p2p_inbound_handler,
|
||||
self.socket_manager,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn ping_all(
|
||||
network: SharedNetworkAddr,
|
||||
packet_loss_stats: PacketLossStats,
|
||||
route_table: RouteTable,
|
||||
socket_manager: P2pOutbound,
|
||||
) {
|
||||
loop {
|
||||
tokio::time::sleep(Duration::from_secs(5)).await;
|
||||
let Some(src) = network.ip() else {
|
||||
continue;
|
||||
};
|
||||
let vec = route_table.route_table();
|
||||
|
||||
for (id, list) in vec {
|
||||
for (index, route) in list.iter().enumerate() {
|
||||
if index > 2 {
|
||||
break;
|
||||
}
|
||||
let Ok(mut ping) = NetPacket::new(TransmissionBytes::zeroed_size(
|
||||
HEAD_LENGTH + 8,
|
||||
socket_manager.encrypt_reserve(),
|
||||
)) else {
|
||||
continue;
|
||||
};
|
||||
ping.set_msg_type(MsgType::Ping);
|
||||
ping.set_ttl(1);
|
||||
ping.set_src_id(src.into());
|
||||
ping.set_dest_id(id.into());
|
||||
ping.set_payload(&crate::utils::time::now_ts_ms().to_be_bytes())
|
||||
.unwrap();
|
||||
if socket_manager
|
||||
.send_to(ping, &route.route_key())
|
||||
.await
|
||||
.is_ok()
|
||||
{
|
||||
packet_loss_stats.record_sent(id);
|
||||
}
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
pub async fn route_timeout_task(route_table: RouteTable) {
|
||||
loop {
|
||||
tokio::time::sleep(Duration::from_secs(10)).await;
|
||||
let expired_time = std::time::Instant::now() - Duration::from_secs(10);
|
||||
route_table.remove_oldest_route(expired_time);
|
||||
}
|
||||
}
|
||||
|
||||
/// 客户端中继探测任务
|
||||
/// 每5分钟执行一次,找到所有未直连的目标IP,通过Ping消息发送给已打洞的客户端
|
||||
pub async fn relay_probe_task(
|
||||
network: SharedNetworkAddr,
|
||||
server_info: crate::context::ServerInfoCollection,
|
||||
route_table: RouteTable,
|
||||
socket_manager: P2pOutbound,
|
||||
) {
|
||||
use rand::prelude::*;
|
||||
|
||||
loop {
|
||||
tokio::time::sleep(Duration::from_secs(300)).await; // 5分钟
|
||||
|
||||
let Some(src) = network.ip() else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let online_ips = server_info.client_online_ips();
|
||||
if online_ips.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let mut non_direct_targets = Vec::new();
|
||||
for ip in online_ips {
|
||||
if ip == src {
|
||||
continue;
|
||||
}
|
||||
|
||||
let is_direct = route_table
|
||||
.get_route_by_id(&ip)
|
||||
.ok()
|
||||
.map(|route| route.is_direct())
|
||||
.unwrap_or(false);
|
||||
|
||||
if !is_direct {
|
||||
non_direct_targets.push(ip);
|
||||
if non_direct_targets.len() >= 20 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if non_direct_targets.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let non_direct_count = non_direct_targets.len();
|
||||
|
||||
let targets_to_probe: Vec<Ipv4Addr> = {
|
||||
let mut rng = rand::rng();
|
||||
if non_direct_targets.len() <= 10 {
|
||||
non_direct_targets
|
||||
} else {
|
||||
non_direct_targets
|
||||
.choose_multiple(&mut rng, 10)
|
||||
.copied()
|
||||
.collect()
|
||||
}
|
||||
};
|
||||
|
||||
let all_routes = route_table.route_table();
|
||||
let mut direct_peers = Vec::new();
|
||||
for (ip, routes) in &all_routes {
|
||||
if let Some(best_route) = routes.first() {
|
||||
if best_route.is_direct() {
|
||||
direct_peers.push((*ip, best_route.route_key()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if direct_peers.is_empty() {
|
||||
log::debug!("No direct peers available for relay probe");
|
||||
continue;
|
||||
}
|
||||
|
||||
let max_probes_per_target = 3.min(direct_peers.len());
|
||||
|
||||
for target_ip in &targets_to_probe {
|
||||
let selected_peers: Vec<_> = {
|
||||
let mut rng = rand::rng();
|
||||
direct_peers
|
||||
.iter()
|
||||
.filter(|(ip, _)| ip != target_ip)
|
||||
.choose_multiple(&mut rng, max_probes_per_target)
|
||||
.into_iter()
|
||||
.cloned()
|
||||
.collect()
|
||||
};
|
||||
|
||||
for (relay_ip, route_key) in selected_peers {
|
||||
// 构造Ping消息,目标是target_ip,但发送给relay_ip
|
||||
let Ok(mut ping) = NetPacket::new(TransmissionBytes::zeroed_size(
|
||||
HEAD_LENGTH + 8,
|
||||
socket_manager.encrypt_reserve(),
|
||||
)) else {
|
||||
continue;
|
||||
};
|
||||
|
||||
ping.set_msg_type(MsgType::Ping);
|
||||
ping.set_ttl(2); // TTL设为2,允许中继一次
|
||||
ping.set_src_id(src.into());
|
||||
ping.set_dest_id((*target_ip).into());
|
||||
ping.set_payload(&crate::utils::time::now_ts_ms().to_be_bytes())
|
||||
.unwrap();
|
||||
|
||||
// 发送给已打洞的客户端,让它中继到目标
|
||||
if let Err(e) = socket_manager.send_to(ping, &route_key).await {
|
||||
log::debug!(
|
||||
"Failed to send relay probe to {} for target {}: {:?}",
|
||||
relay_ip,
|
||||
target_ip,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// 控制发送速率,避免网络拥塞
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
}
|
||||
|
||||
log::info!(
|
||||
"Relay probe task completed: {} targets probed (from {} non-direct), {} direct peers",
|
||||
targets_to_probe.len(),
|
||||
non_direct_count,
|
||||
direct_peers.len()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// 隧道收发调度与数据分发
|
||||
pub async fn tunnel_dispatch_task(
|
||||
nat_info: MyNatInfo,
|
||||
task_group: TaskGroup,
|
||||
mut tunnel_factory: TunnelDispatcher,
|
||||
p2p_inbound_handler: P2pInboundHandler,
|
||||
p2p_socket_manager: P2pOutbound,
|
||||
) {
|
||||
loop {
|
||||
let mut tunnel = match tunnel_factory.dispatch().await {
|
||||
Ok(rs) => rs,
|
||||
Err(e) => {
|
||||
log::error!("tunnel disptach :{e:?}");
|
||||
return;
|
||||
}
|
||||
};
|
||||
log::info!("tunnel {:?}-{:?}", tunnel.protocol(), tunnel.remote_addr());
|
||||
let p2p_inbound_handler = p2p_inbound_handler.clone();
|
||||
let nat_info = nat_info.clone();
|
||||
let p2p_socket_manager = p2p_socket_manager.clone();
|
||||
task_group.spawn(async move {
|
||||
let mut buf = vec![0; 65536];
|
||||
while let Some(rs) = tunnel.recv_from(&mut buf).await {
|
||||
let (len, route_key) = match rs {
|
||||
Ok(rs) => rs,
|
||||
Err(e) => {
|
||||
log::warn!("recv_from {e:?}");
|
||||
if tunnel.protocol().is_udp() {
|
||||
continue;
|
||||
}
|
||||
break;
|
||||
}
|
||||
};
|
||||
if tunnel.protocol().is_udp()
|
||||
&& rust_p2p_core::stun::is_stun_response(&buf[..len])
|
||||
&& let Some(pub_addr) = rust_p2p_core::stun::recv_stun_response(&buf[..len])
|
||||
{
|
||||
nat_info.update_public_addr(route_key.index(), pub_addr);
|
||||
continue;
|
||||
}
|
||||
let mut bytes = TransmissionBytes::zeroed(len);
|
||||
bytes.copy_from_slice(&buf[..len]);
|
||||
p2p_inbound_handler
|
||||
.next_handle(bytes, route_key, &p2p_socket_manager, &mut tunnel)
|
||||
.await;
|
||||
}
|
||||
log::info!(
|
||||
"drop tunnel {:?}-{:?}",
|
||||
tunnel.protocol(),
|
||||
tunnel.remote_addr()
|
||||
);
|
||||
if let Tunnel::Tcp(tcp) = tunnel {
|
||||
p2p_inbound_handler.tcp_disconnect(tcp.route_key()).await;
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,346 @@
|
||||
use crate::compression::PacketCompression;
|
||||
use crate::context::config::Config;
|
||||
use crate::context::nat::{MyNatInfo, PunchBackoff};
|
||||
use crate::context::{AppState, NetworkAddr, NetworkRoute, PeerInfoMap, ServerInfoCollection};
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::enhanced_tunnel::inbound::EnhancedInbound;
|
||||
use crate::fec::FecDecoder;
|
||||
use crate::protocol::control_message::{
|
||||
ConfirmRegResponseMsg, RegResponseMsg, RegistrationMode, RequestMessage, ResponseMessage,
|
||||
};
|
||||
use crate::tunnel_core::p2p::transport::punch::NatPuncher;
|
||||
use crate::tunnel_core::server::inbound::ServerTurnInboundHandler;
|
||||
use crate::tunnel_core::server::outbound::ServerOutbound;
|
||||
use crate::tunnel_core::server::rpc::{RpcNotifier, ServerRPC};
|
||||
use crate::tunnel_core::server::transport::TransportClient;
|
||||
use crate::tunnel_core::server::transport::config::ConnectRegConfig;
|
||||
use crate::utils::task_control::TaskGroup;
|
||||
use anyhow::bail;
|
||||
use bytes::Bytes;
|
||||
use std::collections::HashMap;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::sync::mpsc::{Receiver, Sender};
|
||||
|
||||
pub struct InboundHandlerConfig {
|
||||
pub network_route: NetworkRoute,
|
||||
pub server_info: ServerInfoCollection,
|
||||
pub nat_info: MyNatInfo,
|
||||
pub peer_map: PeerInfoMap,
|
||||
pub punch_backoff: PunchBackoff,
|
||||
pub puncher: NatPuncher,
|
||||
pub packet_crypto: PacketCrypto,
|
||||
pub packet_compression: PacketCompression,
|
||||
pub enhanced_inbound: EnhancedInbound,
|
||||
pub fec_decoder: FecDecoder,
|
||||
}
|
||||
|
||||
pub struct ServerTurnManager {
|
||||
server_id: u32,
|
||||
config: ConnectRegConfig,
|
||||
receiver: Option<Receiver<(Bytes, Instant)>>,
|
||||
notifier: RpcNotifier,
|
||||
transport_client: TransportClient,
|
||||
}
|
||||
pub(crate) fn create_server_tunnel(
|
||||
app_state: AppState,
|
||||
config: &Config,
|
||||
packet_crypto: PacketCrypto,
|
||||
) -> (Vec<ServerTurnManager>, ServerOutbound, ServerRPC) {
|
||||
let mut rpc_notifier: HashMap<u32, RpcNotifier> = HashMap::new();
|
||||
let mut sender_map: HashMap<u32, Sender<(Bytes, Instant)>> = HashMap::new();
|
||||
let mut server_manager_list = Vec::with_capacity(config.server_addr.len());
|
||||
let mut server_addr_list = Vec::with_capacity(config.server_addr.len());
|
||||
for (index, server_addr) in config.server_addr.iter().enumerate() {
|
||||
let connect_reg_config = config.to_connect_config(index);
|
||||
|
||||
let server_id = index as u32;
|
||||
|
||||
let (s, r) = tokio::sync::mpsc::channel(1024);
|
||||
|
||||
let notifier = RpcNotifier::new();
|
||||
let manager =
|
||||
ServerTurnManager::new(server_id, connect_reg_config.clone(), r, notifier.clone());
|
||||
server_addr_list.push((server_id, server_addr.clone()));
|
||||
rpc_notifier.insert(server_id, notifier);
|
||||
sender_map.insert(server_id, s);
|
||||
server_manager_list.push(manager);
|
||||
}
|
||||
let server_info_collection = app_state.server_info_collection.clone();
|
||||
server_info_collection.update_server(server_addr_list);
|
||||
let tunnel_to_server =
|
||||
ServerOutbound::new(Arc::new(sender_map), server_info_collection, packet_crypto);
|
||||
|
||||
let server_rpc = ServerRPC::new(tunnel_to_server.clone(), rpc_notifier);
|
||||
|
||||
(server_manager_list, tunnel_to_server, server_rpc)
|
||||
}
|
||||
|
||||
impl ServerTurnManager {
|
||||
pub fn new(
|
||||
server_id: u32,
|
||||
config: ConnectRegConfig,
|
||||
receiver: Receiver<(Bytes, Instant)>,
|
||||
notifier: RpcNotifier,
|
||||
) -> Self {
|
||||
let connector = TransportClient::new();
|
||||
Self {
|
||||
server_id,
|
||||
transport_client: connector,
|
||||
config,
|
||||
receiver: Some(receiver),
|
||||
notifier,
|
||||
}
|
||||
}
|
||||
pub fn disconnect(&mut self) {
|
||||
self.transport_client.disconnect();
|
||||
}
|
||||
|
||||
pub async fn connect_and_reg(
|
||||
&mut self,
|
||||
mode: RegistrationMode,
|
||||
) -> anyhow::Result<ResponseMessage> {
|
||||
let connect_config = self.config.to_connect_config().await?;
|
||||
log::info!(
|
||||
"Connecting to server[{}] {:?} with mode {:?}",
|
||||
self.server_id,
|
||||
connect_config,
|
||||
mode,
|
||||
);
|
||||
|
||||
self.transport_client
|
||||
.connect_timeout(&connect_config, Duration::from_secs(10))
|
||||
.await?;
|
||||
|
||||
let reg_msg = self.config.reg_msg_request(self.server_id, mode);
|
||||
let request_msg = RequestMessage::Reg(reg_msg);
|
||||
let encoded = request_msg.encode();
|
||||
|
||||
self.transport_client
|
||||
.send(encoded.freeze())
|
||||
.await?;
|
||||
let buf = self
|
||||
.transport_client
|
||||
.next_timeout(Duration::from_secs(10))
|
||||
.await?;
|
||||
let response = ResponseMessage::from_slice(&buf)?;
|
||||
match &response {
|
||||
ResponseMessage::Reg(_) => {}
|
||||
ResponseMessage::Error(_e) => {
|
||||
self.disconnect();
|
||||
}
|
||||
ResponseMessage::ConfirmReg(_) => {
|
||||
self.disconnect();
|
||||
}
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
pub async fn send_confirm(&mut self) -> anyhow::Result<ConfirmRegResponseMsg> {
|
||||
self.transport_client
|
||||
.send(RequestMessage::ConfirmReg.encode().freeze())
|
||||
.await?;
|
||||
let buf = self
|
||||
.transport_client
|
||||
.next_timeout(Duration::from_secs(10))
|
||||
.await?;
|
||||
let response = ResponseMessage::from_slice(&buf)?;
|
||||
match response {
|
||||
ResponseMessage::ConfirmReg(msg) => Ok(msg),
|
||||
ResponseMessage::Error(e) => bail!("Confirm failed: {}", e.message),
|
||||
_ => bail!("Unexpected response"),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn set_ip(&mut self, ip: Ipv4Addr) {
|
||||
self.config.ip = Some(ip);
|
||||
}
|
||||
|
||||
/// Start data handling task with an already established connection.
|
||||
pub fn data_handle_task_connected(
|
||||
mut self,
|
||||
task_group: &TaskGroup,
|
||||
config: Box<InboundHandlerConfig>,
|
||||
initial_response: NetworkAddr,
|
||||
) {
|
||||
let data_handler =
|
||||
ServerTurnInboundHandler::new(self.server_id, initial_response, config);
|
||||
let task_group_ = task_group.clone();
|
||||
let Some(mut receiver) = self.receiver.take() else {
|
||||
unreachable!()
|
||||
};
|
||||
|
||||
task_group.spawn(async move {
|
||||
let mut already_connected = true;
|
||||
loop {
|
||||
if !already_connected {
|
||||
self.disconnect();
|
||||
data_handler.handle_disconnected();
|
||||
let msg = match self.connect_and_reg(RegistrationMode::Normal).await {
|
||||
Ok(msg) => msg,
|
||||
Err(e) => {
|
||||
log::error!("连接服务器失败:{e:?}");
|
||||
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
match &msg {
|
||||
ResponseMessage::Reg(reg) => {
|
||||
if reg.ip != initial_response.ip
|
||||
|| reg.prefix_len != initial_response.prefix_len
|
||||
|| reg.gateway != initial_response.gateway
|
||||
{
|
||||
log::error!("虚拟网络发生变化");
|
||||
break;
|
||||
}
|
||||
// 保存服务器版本
|
||||
if !reg.server_version.is_empty() {
|
||||
data_handler.set_server_version(reg.server_version.clone());
|
||||
}
|
||||
}
|
||||
ResponseMessage::Error(e) => {
|
||||
log::error!("注册失败 {e:?}");
|
||||
break;
|
||||
}
|
||||
_ => {
|
||||
log::error!("错误的注册消息");
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
log::info!("已连接服务器:{}", self.config.server_addr);
|
||||
data_handler.handle_connected();
|
||||
|
||||
if let Err(e) = self
|
||||
.data_handle_loop(&mut receiver, &data_handler)
|
||||
.await
|
||||
{
|
||||
log::error!("Error on data_handle_loop: {:?}", e);
|
||||
}
|
||||
already_connected = false;
|
||||
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
|
||||
}
|
||||
self.disconnect();
|
||||
data_handler.handle_disconnected();
|
||||
task_group_.stop();
|
||||
});
|
||||
}
|
||||
|
||||
pub async fn data_handle_loop(
|
||||
&mut self,
|
||||
receiver: &mut Receiver<(Bytes, Instant)>,
|
||||
data_handler: &ServerTurnInboundHandler,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut time = crate::utils::time::now_ts_ms();
|
||||
let mut ping_interval = tokio::time::interval(Duration::from_secs(5));
|
||||
loop {
|
||||
tokio::select! {
|
||||
Some((buf,expired)) = receiver.recv() => {
|
||||
if expired < Instant::now(){
|
||||
continue;
|
||||
}
|
||||
self.transport_client.send(buf).await?;
|
||||
}
|
||||
rs = self.transport_client.next() => {
|
||||
time = crate::utils::time::now_ts_ms();
|
||||
let data = rs?;
|
||||
if let Err(e) = data_handler.handle(&mut self.transport_client,data, &self.notifier,time).await{
|
||||
log::warn!("Error handling data: {:?}", e);
|
||||
}
|
||||
}
|
||||
_ = ping_interval.tick() => {
|
||||
let now = crate::utils::time::now_ts_ms();
|
||||
if now > time + Duration::from_secs(20).as_millis() as i64 {
|
||||
bail!("timeout")
|
||||
}
|
||||
data_handler.handle_ping(&mut self.transport_client,now).await?;
|
||||
}
|
||||
else => {
|
||||
bail!("receiver closed");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Coordinated multi-server pre-registration.
|
||||
/// 1. First server uses PRE_REGISTER mode to get IP
|
||||
/// 2. Other servers pre-register with the obtained IP
|
||||
/// 3. Send confirmation to all servers
|
||||
/// 4. Return the registration response
|
||||
pub async fn coordinated_registration(
|
||||
managers: &mut Vec<ServerTurnManager>,
|
||||
) -> anyhow::Result<RegResponseMsg> {
|
||||
if managers.is_empty() {
|
||||
bail!("No servers to register");
|
||||
}
|
||||
|
||||
// Step 1: First server pre-register to get IP
|
||||
log::info!(
|
||||
"Starting coordinated registration with {} servers",
|
||||
managers.len()
|
||||
);
|
||||
let first_response = managers[0]
|
||||
.connect_and_reg(RegistrationMode::PreRegister)
|
||||
.await?;
|
||||
|
||||
let ip = match &first_response {
|
||||
ResponseMessage::Reg(reg) => reg.ip,
|
||||
ResponseMessage::Error(e) => bail!("First server registration failed: {}", e.message),
|
||||
_ => bail!("Unexpected response from first server"),
|
||||
};
|
||||
log::info!("Got IP {} from first server", ip);
|
||||
|
||||
// Step 2: Set IP and pre-register with other servers
|
||||
for manager in managers.iter_mut().skip(1) {
|
||||
manager.set_ip(ip);
|
||||
}
|
||||
|
||||
if managers.len() > 1 {
|
||||
let other_results: Vec<_> = futures::future::join_all(
|
||||
managers
|
||||
.iter_mut()
|
||||
.skip(1)
|
||||
.map(|m| m.connect_and_reg(RegistrationMode::PreRegister)),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Check all responses
|
||||
for (i, result) in other_results.iter().enumerate() {
|
||||
match result {
|
||||
Ok(ResponseMessage::Reg(_)) => {
|
||||
log::info!("Server {} pre-registered successfully", i + 1);
|
||||
}
|
||||
Ok(ResponseMessage::Error(e)) => {
|
||||
bail!("Server {} registration failed: {}", i + 1, e.message)
|
||||
}
|
||||
Err(e) => bail!("Server {} registration failed: {}", i + 1, e),
|
||||
_ => bail!("Unexpected response from server {}", i + 1),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Step 3: Send confirmation to all servers
|
||||
log::info!("Sending confirmation to all servers");
|
||||
let confirm_results: Vec<_> =
|
||||
futures::future::join_all(managers.iter_mut().map(|m| m.send_confirm())).await;
|
||||
|
||||
// Check all confirmation responses
|
||||
for (i, result) in confirm_results.into_iter().enumerate() {
|
||||
match result {
|
||||
Ok(msg) if msg.success => {
|
||||
log::info!("Server {} confirmed successfully", i);
|
||||
}
|
||||
Ok(_) => bail!("Server {} confirmation failed", i),
|
||||
Err(e) => bail!("Server {} confirmation failed: {}", i, e),
|
||||
}
|
||||
}
|
||||
|
||||
log::info!("Coordinated registration completed successfully");
|
||||
// Return first server's response (contains IP info)
|
||||
match first_response {
|
||||
ResponseMessage::Reg(reg) => Ok(reg),
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
use crate::compression::PacketCompression;
|
||||
use crate::context::nat::{MyNatInfo, PunchBackoff};
|
||||
use crate::context::{NetworkAddr, NetworkRoute, PeerInfoMap, ServerInfoCollection};
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::enhanced_tunnel::inbound::EnhancedInbound;
|
||||
use crate::fec::FecDecoder;
|
||||
use crate::protocol::client_message::PunchInfo;
|
||||
use crate::protocol::control_message::ClientSimpleInfoList;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::rpc_message::RpcMessageResponse;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::p2p::transport::punch::NatPuncher;
|
||||
use crate::tunnel_core::server::rpc::RpcNotifier;
|
||||
use crate::tunnel_core::server::transport::TransportClient;
|
||||
use anyhow::bail;
|
||||
use pnet_packet::Packet;
|
||||
use pnet_packet::icmp::{IcmpPacket, IcmpTypes};
|
||||
use pnet_packet::ipv4::Ipv4Packet;
|
||||
use prost::Message;
|
||||
use rust_p2p_core::nat::NatInfo;
|
||||
use std::net::Ipv4Addr;
|
||||
|
||||
pub(crate) struct ServerTurnInboundHandler {
|
||||
server_id: u32,
|
||||
network_addr: Option<NetworkAddr>,
|
||||
network_route: NetworkRoute,
|
||||
server_info: ServerInfoCollection,
|
||||
nat_info: MyNatInfo,
|
||||
peer_map: PeerInfoMap,
|
||||
punch_backoff: PunchBackoff,
|
||||
puncher: NatPuncher,
|
||||
packet_crypto: PacketCrypto,
|
||||
packet_compression: PacketCompression,
|
||||
enhanced_inbound: EnhancedInbound,
|
||||
fec_decoder: FecDecoder,
|
||||
}
|
||||
impl ServerTurnInboundHandler {
|
||||
pub fn new(
|
||||
server_id: u32,
|
||||
network_addr: NetworkAddr,
|
||||
config: Box<super::connection_manager::InboundHandlerConfig>,
|
||||
) -> Self {
|
||||
let config = *config;
|
||||
Self {
|
||||
server_id,
|
||||
network_addr: Some(network_addr),
|
||||
network_route: config.network_route,
|
||||
server_info: config.server_info,
|
||||
nat_info: config.nat_info,
|
||||
peer_map: config.peer_map,
|
||||
punch_backoff: config.punch_backoff,
|
||||
puncher: config.puncher,
|
||||
packet_crypto: config.packet_crypto,
|
||||
packet_compression: config.packet_compression,
|
||||
enhanced_inbound: config.enhanced_inbound,
|
||||
fec_decoder: config.fec_decoder,
|
||||
}
|
||||
}
|
||||
fn network_contains(&self, ip: &Ipv4Addr) -> bool {
|
||||
self.network_route.network_contains(ip)
|
||||
}
|
||||
fn filter_ip(&self, mut info: NatInfo) -> NatInfo {
|
||||
if self.network_contains(&info.local_ipv4) {
|
||||
info.local_ipv4 = Ipv4Addr::UNSPECIFIED;
|
||||
}
|
||||
info.local_ipv4s.retain(|ip| !self.network_contains(ip));
|
||||
info
|
||||
}
|
||||
fn get_punch_info(&self) -> Option<PunchInfo> {
|
||||
self.nat_info.get().map(|info| PunchInfo {
|
||||
nat_info: self.filter_ip(info),
|
||||
})
|
||||
}
|
||||
fn update_peer_nat_info(&self, ip: Ipv4Addr, nat_info: NatInfo) {
|
||||
self.peer_map.update_nat_info(ip, nat_info);
|
||||
}
|
||||
|
||||
pub async fn handle_server_data(
|
||||
& self,
|
||||
transport_client: &mut TransportClient,
|
||||
network_addr: NetworkAddr,
|
||||
data: TransmissionBytes,
|
||||
rpc_notifier: &RpcNotifier,
|
||||
now: i64,
|
||||
) -> anyhow::Result<()> {
|
||||
let net_packet = NetPacket::new(data)?;
|
||||
let src = net_packet.src_id().into();
|
||||
let msg_type = net_packet.msg_type()?;
|
||||
let mut net_packet = self.packet_compression.decompress(net_packet)?;
|
||||
|
||||
match msg_type {
|
||||
MsgType::Turn => {
|
||||
// 只允许icmp EchoReply
|
||||
let Some(ipv4) = Ipv4Packet::new(net_packet.payload()) else {
|
||||
return Ok(());
|
||||
};
|
||||
if ipv4.get_version() != 4 {
|
||||
return Ok(());
|
||||
}
|
||||
if ipv4.get_next_level_protocol() != pnet_packet::ip::IpNextHeaderProtocols::Icmp {
|
||||
return Ok(());
|
||||
}
|
||||
let Some(icmp) = IcmpPacket::new(ipv4.payload()) else {
|
||||
return Ok(());
|
||||
};
|
||||
if icmp.get_icmp_type() != IcmpTypes::EchoReply {
|
||||
return Ok(());
|
||||
}
|
||||
self.enhanced_inbound
|
||||
.inbound(&network_addr, msg_type, src, net_packet)
|
||||
.await?;
|
||||
}
|
||||
MsgType::Ping => {
|
||||
net_packet.set_ttl(2);
|
||||
net_packet.set_msg_type(MsgType::Pong);
|
||||
net_packet.set_src_id(network_addr.ip.into());
|
||||
net_packet.set_dest_id(src.into());
|
||||
transport_client.send_turn(net_packet).await?;
|
||||
}
|
||||
MsgType::PongTurn => {
|
||||
// 服务端ping 回复,记录延迟
|
||||
if net_packet.payload().len() == 8 + 8 {
|
||||
let time = i64::from_be_bytes(net_packet.payload()[..8].try_into()?);
|
||||
// let data_version = u64::from_be_bytes(net_packet.payload()[8..].try_into()?);
|
||||
if now >= time {
|
||||
self.server_info
|
||||
.set_server_rtt(self.server_id, (now - time) as u32);
|
||||
}
|
||||
}
|
||||
}
|
||||
MsgType::PushClientIps => {
|
||||
let list = ClientSimpleInfoList::from_slice(net_packet.payload())?;
|
||||
self.server_info.update_client_simple_list(
|
||||
self.server_id,
|
||||
network_addr.ip,
|
||||
list,
|
||||
now,
|
||||
);
|
||||
}
|
||||
MsgType::RpcRes => {
|
||||
// 设置rpc响应
|
||||
let response = RpcMessageResponse::decode(net_packet.payload())?;
|
||||
rpc_notifier.notify_response(response);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
pub async fn handle_client_data(
|
||||
&self,
|
||||
network_addr: NetworkAddr,
|
||||
transport_client: &mut TransportClient,
|
||||
data: TransmissionBytes,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut net_packet = NetPacket::new(data)?;
|
||||
let msg_type = net_packet.msg_type()?;
|
||||
let src = Ipv4Addr::from(net_packet.src_id());
|
||||
let dest = Ipv4Addr::from(net_packet.dest_id());
|
||||
|
||||
if msg_type == MsgType::Quic {
|
||||
// QUIC 数据不加密不压缩,但可能有 FEC
|
||||
if net_packet.is_fec() {
|
||||
let packets = self.fec_decoder.receive(net_packet)?;
|
||||
if let Some(packets) = packets {
|
||||
for pkt in packets {
|
||||
self.enhanced_inbound
|
||||
.inbound(&network_addr, msg_type, src, pkt)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
self.enhanced_inbound
|
||||
.inbound(&network_addr, msg_type, src, net_packet)
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// 解密
|
||||
if let Err(e) = self.packet_crypto.decrypt_in_place(&mut net_packet) {
|
||||
log::error!("{},mst_type={msg_type:?},src={src},dst={dest}", e);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// FEC 解码(始终尝试解码,如果有 FEC 标志)
|
||||
if net_packet.is_fec() {
|
||||
let packets = self.fec_decoder.receive(net_packet)?;
|
||||
if let Some(packets) = packets {
|
||||
for pkt in packets {
|
||||
let pkt = self.packet_compression.decompress(pkt)?;
|
||||
self.process_decompressed_packet(
|
||||
network_addr,
|
||||
transport_client,
|
||||
pkt,
|
||||
msg_type,
|
||||
src,
|
||||
dest,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// 解压缩
|
||||
let net_packet = self.packet_compression.decompress(net_packet)?;
|
||||
self.process_decompressed_packet(
|
||||
network_addr,
|
||||
transport_client,
|
||||
net_packet,
|
||||
msg_type,
|
||||
src,
|
||||
dest,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn process_decompressed_packet(
|
||||
&self,
|
||||
network_addr: NetworkAddr,
|
||||
transport_client: &mut TransportClient,
|
||||
net_packet: NetPacket<TransmissionBytes>,
|
||||
msg_type: MsgType,
|
||||
src: Ipv4Addr,
|
||||
dest: Ipv4Addr,
|
||||
) -> anyhow::Result<()> {
|
||||
match msg_type {
|
||||
MsgType::Turn | MsgType::Broadcast => {
|
||||
self.enhanced_inbound
|
||||
.inbound(&network_addr, msg_type, src, net_packet)
|
||||
.await?;
|
||||
}
|
||||
MsgType::PunchStart1 => {
|
||||
// 对方发起打洞
|
||||
let peer_punch_info = PunchInfo::from_slice(net_packet.payload())?;
|
||||
let Some(self_punch_info) = self.get_punch_info() else {
|
||||
return Ok(());
|
||||
};
|
||||
log::info!(
|
||||
"对方主动发起打洞 对方nat信息={peer_punch_info:?},自己nat信息={self_punch_info:?} {src}->{dest}"
|
||||
);
|
||||
self.update_peer_nat_info(src, peer_punch_info.nat_info.clone());
|
||||
let rs = self.puncher.punch(src, peer_punch_info)?;
|
||||
if rs {
|
||||
let bytes_mut = self_punch_info.encode();
|
||||
let mut net_packet = NetPacket::new(TransmissionBytes::zeroed_size(
|
||||
HEAD_LENGTH + bytes_mut.len(),
|
||||
self.packet_crypto.encrypt_reserve(),
|
||||
))?;
|
||||
net_packet.set_msg_type(MsgType::PunchStart2);
|
||||
net_packet.set_ttl(2);
|
||||
net_packet.set_src_id(dest.into());
|
||||
net_packet.set_dest_id(src.into());
|
||||
net_packet.set_payload(&bytes_mut)?;
|
||||
self.packet_crypto.encrypt_in_place(&mut net_packet)?;
|
||||
transport_client.send_turn(net_packet).await?;
|
||||
}else{
|
||||
log::info!("限制打洞频率")
|
||||
}
|
||||
}
|
||||
MsgType::PunchStart2 => {
|
||||
self.punch_backoff.record(src);
|
||||
// 对方回复开始打洞
|
||||
let peer_punch_info = PunchInfo::from_slice(net_packet.payload())?;
|
||||
self.update_peer_nat_info(src, peer_punch_info.nat_info.clone());
|
||||
log::info!("对方回复开始打洞 {:?} {src}->{dest}", peer_punch_info);
|
||||
self.puncher.punch_uncheck(src, peer_punch_info)?;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn handle(
|
||||
&self,
|
||||
transport_client: &mut TransportClient,
|
||||
data: TransmissionBytes,
|
||||
rpc_notifier: &RpcNotifier,
|
||||
now: i64,
|
||||
) -> anyhow::Result<()> {
|
||||
let net_packet = NetPacket::new(&data)?;
|
||||
let Some(network_addr) = self.network_addr else {
|
||||
bail!("未找到自身IP")
|
||||
};
|
||||
|
||||
if net_packet.is_gateway() {
|
||||
// 服务端数据
|
||||
return self
|
||||
.handle_server_data(transport_client, network_addr, data, rpc_notifier, now)
|
||||
.await;
|
||||
}
|
||||
let dest = Ipv4Addr::from(net_packet.dest_id());
|
||||
if !dest.is_broadcast() && !dest.is_unspecified() && network_addr.ip != dest {
|
||||
return Ok(());
|
||||
}
|
||||
self.handle_client_data(network_addr, transport_client, data)
|
||||
.await
|
||||
}
|
||||
pub async fn handle_ping(
|
||||
&self,
|
||||
transport_client: &mut TransportClient,
|
||||
now: i64,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut ping_packet = NetPacket::new(TransmissionBytes::zeroed(HEAD_LENGTH + 8 + 8))?;
|
||||
ping_packet.set_ttl(1);
|
||||
ping_packet.set_msg_type(MsgType::PingTurn);
|
||||
ping_packet.set_gateway_flag(true);
|
||||
ping_packet.set_payload(&now.to_be_bytes())?;
|
||||
ping_packet.payload_mut()[0..8].copy_from_slice(&now.to_be_bytes());
|
||||
ping_packet.payload_mut()[8..]
|
||||
.copy_from_slice(&self.server_info.data_version(self.server_id).to_be_bytes());
|
||||
transport_client
|
||||
.send(ping_packet.into_buffer().into_bytes().freeze())
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
pub fn handle_connected(&self) {
|
||||
self.server_info.set_server_connected(self.server_id, true);
|
||||
self.server_info
|
||||
.set_last_connected_time(self.server_id, Some(crate::utils::time::now_ts_ms()));
|
||||
self.server_info.set_disconnected_time(self.server_id, None);
|
||||
}
|
||||
pub fn set_server_version(&self, version: String) {
|
||||
self.server_info.set_server_version(self.server_id, version);
|
||||
}
|
||||
pub fn handle_disconnected(&self) {
|
||||
if self.server_info.set_server_connected(self.server_id, false) {
|
||||
self.server_info
|
||||
.set_disconnected_time(self.server_id, Some(crate::utils::time::now_ts_ms()));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
pub(crate) mod connection_manager;
|
||||
pub(crate) mod inbound;
|
||||
pub(crate) mod outbound;
|
||||
pub(crate) mod rpc;
|
||||
pub mod transport;
|
||||
@@ -0,0 +1,329 @@
|
||||
use crate::context::ServerInfoCollection;
|
||||
use crate::crypto::PacketCrypto;
|
||||
use crate::protocol::ProtoToBytesMut;
|
||||
use crate::protocol::control_message::SelectiveBroadcast;
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use anyhow::{Context, bail};
|
||||
use bytes::Bytes;
|
||||
use std::collections::HashMap;
|
||||
use std::net::Ipv4Addr;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::sync::mpsc::Sender;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct ServerOutbound {
|
||||
server_id_list: Arc<Vec<u32>>,
|
||||
sender: Arc<HashMap<u32, Sender<(Bytes, Instant)>>>,
|
||||
server_info_collection: ServerInfoCollection,
|
||||
packet_crypto: PacketCrypto,
|
||||
}
|
||||
impl ServerOutbound {
|
||||
pub fn new(
|
||||
sender: Arc<HashMap<u32, Sender<(Bytes, Instant)>>>,
|
||||
server_info_collection: ServerInfoCollection,
|
||||
packet_crypto: PacketCrypto,
|
||||
) -> Self {
|
||||
let server_id_list = Arc::new(sender.keys().copied().collect());
|
||||
Self {
|
||||
server_id_list,
|
||||
sender,
|
||||
server_info_collection,
|
||||
packet_crypto,
|
||||
}
|
||||
}
|
||||
pub fn exists_route(&self, dest: &Ipv4Addr) -> bool {
|
||||
self.server_info_collection.exists_online_client_ip(dest)
|
||||
}
|
||||
pub fn server_id_list(&self) -> &Vec<u32> {
|
||||
&self.server_id_list
|
||||
}
|
||||
pub fn encrypt_reserve(&self) -> usize {
|
||||
self.packet_crypto.encrypt_reserve()
|
||||
}
|
||||
|
||||
pub async fn send_to_gateway_expired(
|
||||
&self,
|
||||
server_id: u32,
|
||||
mut buf: NetPacket<TransmissionBytes>,
|
||||
expired: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
if !self.server_info_collection.is_server_connected(server_id) {
|
||||
bail!("未连接服务器")
|
||||
}
|
||||
buf.set_gateway_flag(true);
|
||||
self.send_expired_impl(server_id, buf, expired).await
|
||||
}
|
||||
|
||||
pub async fn send(
|
||||
&self,
|
||||
dest_ip: Ipv4Addr,
|
||||
buf: NetPacket<TransmissionBytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
self.send_expired(dest_ip, buf, Duration::from_secs(5))
|
||||
.await
|
||||
}
|
||||
pub async fn send_expired(
|
||||
&self,
|
||||
dest_ip: Ipv4Addr,
|
||||
buf: NetPacket<TransmissionBytes>,
|
||||
expired: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
let Some(server_id) = self
|
||||
.server_info_collection
|
||||
.find_ip_to_server(&self.server_id_list, &dest_ip)
|
||||
else {
|
||||
bail!("not found ip route: {dest_ip}")
|
||||
};
|
||||
self.send_expired_impl(server_id, buf, expired).await
|
||||
}
|
||||
async fn send_expired_impl(
|
||||
&self,
|
||||
server_id: u32,
|
||||
mut buf: NetPacket<TransmissionBytes>,
|
||||
expired: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
if !buf.is_gateway() {
|
||||
self.packet_crypto.encrypt_in_place(&mut buf)?;
|
||||
}
|
||||
let Some(sender) = self.sender.get(&server_id) else {
|
||||
bail!("not found server")
|
||||
};
|
||||
sender
|
||||
.send_timeout(
|
||||
(
|
||||
buf.into_buffer().into_bytes().freeze(),
|
||||
Instant::now() + expired,
|
||||
),
|
||||
expired,
|
||||
)
|
||||
.await
|
||||
.context("connect server task failed")
|
||||
}
|
||||
|
||||
pub async fn send_raw(&self, dest_ip: Ipv4Addr, buf: NetPacket<Bytes>) -> anyhow::Result<()> {
|
||||
let Some(server_id) = self
|
||||
.server_info_collection
|
||||
.find_ip_to_server(&self.server_id_list, &dest_ip)
|
||||
else {
|
||||
bail!("not found ip route: {dest_ip}")
|
||||
};
|
||||
let expired = Duration::from_secs(5);
|
||||
let Some(sender) = self.sender.get(&server_id) else {
|
||||
bail!("not found server")
|
||||
};
|
||||
sender
|
||||
.send_timeout((buf.into_buffer(), Instant::now() + expired), expired)
|
||||
.await
|
||||
.context("connect server task failed")
|
||||
}
|
||||
pub async fn send_default_raw(&self, buf: NetPacket<Bytes>) -> anyhow::Result<()> {
|
||||
let Some(server_id) = self
|
||||
.server_info_collection
|
||||
.find_connected_server(&self.server_id_list)
|
||||
else {
|
||||
bail!("not found default route")
|
||||
};
|
||||
let expired = Duration::from_secs(5);
|
||||
let Some(sender) = self.sender.get(&server_id) else {
|
||||
bail!("not found server")
|
||||
};
|
||||
sender
|
||||
.send_timeout((buf.into_buffer(), Instant::now() + expired), expired)
|
||||
.await
|
||||
.context("connect server task failed")
|
||||
}
|
||||
|
||||
pub async fn send_raw_broadcast(
|
||||
&self,
|
||||
exclude_ips: Option<Vec<Ipv4Addr>>,
|
||||
buf: NetPacket<Bytes>,
|
||||
) -> anyhow::Result<()> {
|
||||
let buf = buf.into_buffer();
|
||||
let expired = Duration::from_secs(5);
|
||||
|
||||
let map: HashMap<u32, (Vec<Ipv4Addr>, u32)> =
|
||||
self.server_info_collection.server_client_ip_map();
|
||||
if map.is_empty() {
|
||||
bail!("no connected servers with clients");
|
||||
}
|
||||
|
||||
// 只有一个服务器,直接发送
|
||||
if map.len() == 1 {
|
||||
let (server_id, (ips, _)) = map.iter().next().expect("map has exactly one element");
|
||||
if ips.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let sender = self
|
||||
.sender
|
||||
.get(server_id)
|
||||
.context("server sender not found")?;
|
||||
|
||||
return if let Some(exclude_ips) = exclude_ips {
|
||||
send_exclude_broadcast(sender.clone(), buf, exclude_ips, expired).await
|
||||
} else {
|
||||
send_direct(sender.clone(), buf, expired).await
|
||||
};
|
||||
}
|
||||
|
||||
// 找到最优服务器
|
||||
let (max_server_id, (max_ips, _)) = map
|
||||
.iter()
|
||||
.max_by(|(_, (ips_a, rtt_a)), (_, (ips_b, rtt_b))| {
|
||||
let score_a = ips_a.len() as f64 / (*rtt_a as f64 + 1.0);
|
||||
let score_b = ips_b.len() as f64 / (*rtt_b as f64 + 1.0);
|
||||
score_a
|
||||
.partial_cmp(&score_b)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
})
|
||||
.context("failed to find server with most IPs")?;
|
||||
if max_ips.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let max_ip_set: std::collections::HashSet<_> = max_ips.iter().collect();
|
||||
let exclude_set: std::collections::HashSet<_> = exclude_ips
|
||||
.as_ref()
|
||||
.map(|ips| ips.iter().collect())
|
||||
.unwrap_or_default();
|
||||
|
||||
let mut handles = Vec::new();
|
||||
|
||||
// 任务1: 向主服务器发送
|
||||
let sender = self
|
||||
.sender
|
||||
.get(max_server_id)
|
||||
.cloned()
|
||||
.context("max server sender not found")?;
|
||||
|
||||
if let Some(exclude_ips) = exclude_ips.clone() {
|
||||
let buf_clone = buf.clone();
|
||||
let handle = tokio::spawn(async move {
|
||||
send_exclude_broadcast(sender, buf_clone, exclude_ips, expired).await
|
||||
});
|
||||
handles.push(handle);
|
||||
} else {
|
||||
let buf_clone = buf.clone();
|
||||
let handle = tokio::spawn(async move { send_direct(sender, buf_clone, expired).await });
|
||||
handles.push(handle);
|
||||
}
|
||||
|
||||
// 任务2-N: 向其他服务器发送目标广播
|
||||
for (server_id, (ips, _rtt)) in map.iter() {
|
||||
if *server_id == *max_server_id {
|
||||
continue;
|
||||
}
|
||||
|
||||
// 筛选目标IP:不在最大服务器中,也不在排除列表中
|
||||
let target_ips: Vec<Ipv4Addr> = ips
|
||||
.iter()
|
||||
.filter(|ip| !max_ip_set.contains(ip) && !exclude_set.contains(ip))
|
||||
.copied()
|
||||
.collect();
|
||||
|
||||
if target_ips.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let sender = match self.sender.get(server_id).cloned() {
|
||||
Some(s) => s,
|
||||
None => continue,
|
||||
};
|
||||
let buf_clone = buf.clone();
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
send_target_broadcast(sender, target_ips, buf_clone, expired).await
|
||||
});
|
||||
handles.push(handle);
|
||||
}
|
||||
|
||||
// 等待所有任务完成
|
||||
let mut errors = Vec::new();
|
||||
for handle in handles {
|
||||
match handle.await {
|
||||
Ok(Ok(())) => {}
|
||||
Ok(Err(e)) => errors.push(e),
|
||||
Err(e) => errors.push(anyhow::anyhow!("task join error: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
if !errors.is_empty() {
|
||||
bail!(
|
||||
"broadcast failed with {} errors: {:?}",
|
||||
errors.len(),
|
||||
errors
|
||||
);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
async fn send_exclude_broadcast(
|
||||
sender: Sender<(Bytes, Instant)>,
|
||||
buf: Bytes,
|
||||
exclude_ips: Vec<Ipv4Addr>,
|
||||
expired: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
let broadcast = SelectiveBroadcast::new(&exclude_ips, buf.to_vec());
|
||||
let bytes = broadcast.encode_bytes_mut();
|
||||
|
||||
let mut packet = NetPacket::new(TransmissionBytes::zeroed(HEAD_LENGTH + bytes.len()))?;
|
||||
packet.set_msg_type(MsgType::ExcludeBroadcast);
|
||||
packet.set_ttl(5);
|
||||
packet.payload_mut().copy_from_slice(&bytes);
|
||||
|
||||
sender
|
||||
.send_timeout(
|
||||
(
|
||||
packet.into_buffer().into_bytes().freeze(),
|
||||
Instant::now() + expired,
|
||||
),
|
||||
expired,
|
||||
)
|
||||
.await
|
||||
.context("failed to send exclude broadcast")?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// 直接发送原始数据
|
||||
async fn send_direct(
|
||||
sender: Sender<(Bytes, Instant)>,
|
||||
buf: Bytes,
|
||||
expired: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
sender
|
||||
.send_timeout((buf, Instant::now() + expired), expired)
|
||||
.await
|
||||
.context("failed to send direct broadcast")?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn send_target_broadcast(
|
||||
sender: Sender<(Bytes, Instant)>,
|
||||
target_ips: Vec<Ipv4Addr>,
|
||||
buf: Bytes,
|
||||
expired: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
let target_broadcast = SelectiveBroadcast::new(&target_ips, buf.to_vec());
|
||||
let target_bytes = target_broadcast.encode_bytes_mut();
|
||||
|
||||
let mut packet = NetPacket::new(TransmissionBytes::zeroed(HEAD_LENGTH + target_bytes.len()))?;
|
||||
packet.set_msg_type(MsgType::TargetBroadcast);
|
||||
packet.set_ttl(5);
|
||||
packet.payload_mut().copy_from_slice(&target_bytes);
|
||||
|
||||
sender
|
||||
.send_timeout(
|
||||
(
|
||||
packet.into_buffer().into_bytes().freeze(),
|
||||
Instant::now() + expired,
|
||||
),
|
||||
expired,
|
||||
)
|
||||
.await
|
||||
.context("failed to send target broadcast")?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
use crate::protocol::ip_packet_protocol::{HEAD_LENGTH, MsgType, NetPacket};
|
||||
use crate::protocol::rpc_message::rpc_message_request::RpcReqPayload;
|
||||
use crate::protocol::rpc_message::rpc_message_response::RpcResPayload;
|
||||
use crate::protocol::rpc_message::{
|
||||
ClientInfo, ClientListRequest, ClientListResponse, RpcMessageRequest, RpcMessageResponse,
|
||||
};
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::server::outbound::ServerOutbound;
|
||||
use anyhow::bail;
|
||||
use parking_lot::Mutex;
|
||||
use prost::Message;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tokio::sync::oneshot;
|
||||
use tokio::sync::oneshot::Sender;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ServerRPC {
|
||||
tunnel_to_server: ServerOutbound,
|
||||
rpc_notifier: HashMap<u32, RpcNotifier>,
|
||||
}
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct RpcNotifier {
|
||||
pending_requests: Arc<Mutex<HashMap<u64, Sender<RpcMessageResponse>>>>,
|
||||
rpc_id: Arc<Mutex<u64>>,
|
||||
}
|
||||
|
||||
impl RpcNotifier {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
pending_requests: Arc::new(Mutex::new(HashMap::new())),
|
||||
rpc_id: Arc::new(Mutex::new(0)),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn create_request_and_waiter(&self) -> RpcResponseWaiter {
|
||||
let id: u64 = {
|
||||
let mut id_lock = self.rpc_id.lock();
|
||||
*id_lock += 1;
|
||||
*id_lock
|
||||
};
|
||||
|
||||
let (tx, rx) = oneshot::channel();
|
||||
|
||||
{
|
||||
let mut pending = self.pending_requests.lock();
|
||||
pending.insert(id, tx);
|
||||
}
|
||||
|
||||
RpcResponseWaiter {
|
||||
id,
|
||||
pending_requests_handle: Arc::clone(&self.pending_requests),
|
||||
rx,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn notify_response(&self, response: RpcMessageResponse) {
|
||||
let mut pending = self.pending_requests.lock();
|
||||
|
||||
if let Some(tx) = pending.remove(&response.id) {
|
||||
let _ = tx.send(response);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) struct RpcResponseWaiter {
|
||||
id: u64,
|
||||
pending_requests_handle: Arc<Mutex<HashMap<u64, Sender<RpcMessageResponse>>>>,
|
||||
rx: oneshot::Receiver<RpcMessageResponse>,
|
||||
}
|
||||
|
||||
impl RpcResponseWaiter {
|
||||
pub async fn wait_for_response(
|
||||
mut self,
|
||||
timeout: Duration,
|
||||
) -> anyhow::Result<RpcMessageResponse> {
|
||||
let result = tokio::time::timeout(timeout, &mut self.rx).await;
|
||||
|
||||
match result {
|
||||
Err(_) => bail!("timeout waiting for response"),
|
||||
Ok(Ok(response)) => Ok(response),
|
||||
Ok(Err(_)) => bail!("closed connection"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for RpcResponseWaiter {
|
||||
fn drop(&mut self) {
|
||||
let mut pending = self.pending_requests_handle.lock();
|
||||
let _ = pending.remove(&self.id);
|
||||
}
|
||||
}
|
||||
|
||||
impl ServerRPC {
|
||||
pub(crate) fn new(
|
||||
tunnel_to_server: ServerOutbound,
|
||||
rpc_notifier: HashMap<u32, RpcNotifier>,
|
||||
) -> Self {
|
||||
Self {
|
||||
tunnel_to_server,
|
||||
rpc_notifier,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn client_list(&self) -> anyhow::Result<ClientListResponse> {
|
||||
let mut map: HashMap<String, ClientInfo> = HashMap::new();
|
||||
for server_id in self.tunnel_to_server.server_id_list() {
|
||||
match self.client_list_target(*server_id).await {
|
||||
Ok(rs) => {
|
||||
for client in rs.list {
|
||||
map.entry(client.id.clone()).or_insert(client);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
log::error!("client list target failed: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(ClientListResponse {
|
||||
list: map.into_values().collect(),
|
||||
})
|
||||
}
|
||||
pub async fn client_list_target(&self, server_id: u32) -> anyhow::Result<ClientListResponse> {
|
||||
let Some(rpc_notifier) = self.rpc_notifier.get(&server_id) else {
|
||||
bail!("no RPC notifier");
|
||||
};
|
||||
let waiter = rpc_notifier.create_request_and_waiter();
|
||||
let request = RpcMessageRequest {
|
||||
id: waiter.id,
|
||||
rpc_req_payload: Some(RpcReqPayload::ClientListReq(ClientListRequest::default())),
|
||||
};
|
||||
let buf = request.encode_to_vec();
|
||||
let mut packet = NetPacket::new(TransmissionBytes::zeroed(HEAD_LENGTH + buf.len()))?;
|
||||
packet.set_msg_type(MsgType::RpcReq);
|
||||
packet.set_gateway_flag(true);
|
||||
packet.set_ttl(1);
|
||||
packet.set_payload(&buf)?;
|
||||
|
||||
self.tunnel_to_server
|
||||
.send_to_gateway_expired(server_id, packet, Duration::from_secs(1))
|
||||
.await?;
|
||||
let response = waiter.wait_for_response(Duration::from_secs(3)).await?;
|
||||
if let Some(RpcResPayload::ClientListRes(res)) = response.rpc_res_payload {
|
||||
return Ok(res);
|
||||
}
|
||||
bail!("unexpected response: {:?}", response);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
use crate::protocol::control_message::{RegRequestMsg, RegistrationMode};
|
||||
use crate::tls::verifier::CertValidationMode;
|
||||
use anyhow::Context;
|
||||
use rand::seq::SliceRandom;
|
||||
use std::fmt;
|
||||
use std::net::{Ipv4Addr, SocketAddr};
|
||||
use std::str::FromStr;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct ConnectRegConfig {
|
||||
pub server_addr: ProtocolAddress,
|
||||
pub cert_mode: CertValidationMode,
|
||||
pub network_code: String,
|
||||
pub device_id: String,
|
||||
pub device_name: String,
|
||||
pub ip: Option<Ipv4Addr>,
|
||||
pub key_sign: Option<String>,
|
||||
pub ip_variable: bool,
|
||||
}
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct ConnectConfig {
|
||||
pub protocol_type: ProtocolType,
|
||||
pub server_addr: SocketAddr,
|
||||
pub server_domain: String,
|
||||
pub cert_mode: CertValidationMode,
|
||||
}
|
||||
#[derive(Debug, Copy, Clone, Eq, PartialEq, Default)]
|
||||
pub enum ProtocolType {
|
||||
Quic,
|
||||
#[default]
|
||||
TlsTcp,
|
||||
Wss,
|
||||
Dynamic,
|
||||
}
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ProtocolAddress {
|
||||
pub protocol_type: ProtocolType,
|
||||
pub address: String,
|
||||
}
|
||||
impl Default for ProtocolAddress {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
protocol_type: ProtocolType::default(),
|
||||
address: "127.0.0.1:29872".to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
impl FromStr for ProtocolAddress {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
let (protocol_type, server_addr) = parse_server(s)?;
|
||||
Ok(Self {
|
||||
protocol_type,
|
||||
address: server_addr,
|
||||
})
|
||||
}
|
||||
}
|
||||
impl fmt::Display for ProtocolAddress {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let prefix = match self.protocol_type {
|
||||
ProtocolType::Quic => "quic://",
|
||||
ProtocolType::TlsTcp => "tcp://",
|
||||
ProtocolType::Wss => "wss://",
|
||||
ProtocolType::Dynamic => "dynamic://",
|
||||
};
|
||||
write!(f, "{}{}", prefix, self.address)
|
||||
}
|
||||
}
|
||||
pub fn parse_server(val: &str) -> Result<(ProtocolType, String), String> {
|
||||
let val = val.trim().to_lowercase();
|
||||
if let Some(s) = val.strip_prefix("quic://") {
|
||||
return Ok((ProtocolType::Quic, s.to_string()));
|
||||
}
|
||||
if let Some(s) = val.strip_prefix("tcp://") {
|
||||
return Ok((ProtocolType::TlsTcp, s.to_string()));
|
||||
}
|
||||
if let Some(s) = val.strip_prefix("wss://") {
|
||||
return Ok((ProtocolType::Wss, s.to_string()));
|
||||
}
|
||||
if let Some(s) = val.strip_prefix("dynamic://") {
|
||||
return Ok((ProtocolType::Dynamic, s.to_string()));
|
||||
}
|
||||
if val.contains("://") {
|
||||
return Err(format!("Unknown protocol in server address: {}", val));
|
||||
}
|
||||
Ok((ProtocolType::TlsTcp, val))
|
||||
}
|
||||
impl ConnectRegConfig {
|
||||
pub fn reg_msg_request(
|
||||
&self,
|
||||
server_id: u32,
|
||||
registration_mode: RegistrationMode,
|
||||
) -> RegRequestMsg {
|
||||
RegRequestMsg {
|
||||
network_code: self.network_code.to_string(),
|
||||
device_id: self.device_id.to_string(),
|
||||
ip: self.ip,
|
||||
name: self.device_name.to_string(),
|
||||
version: env!("CARGO_PKG_VERSION").to_string(),
|
||||
key_sign: self.key_sign.clone(),
|
||||
ip_variable: self.ip_variable,
|
||||
server_id,
|
||||
registration_mode,
|
||||
}
|
||||
}
|
||||
pub async fn to_connect_config(&self) -> anyhow::Result<ConnectConfig> {
|
||||
let (protocol_type, server_domain) = match self.server_addr.protocol_type {
|
||||
ProtocolType::Dynamic => {
|
||||
let mut txt = crate::utils::dns_query::dns_query_txt(
|
||||
&self.server_addr.address,
|
||||
vec![],
|
||||
&None,
|
||||
)
|
||||
.await?;
|
||||
txt.shuffle(&mut rand::rng());
|
||||
let x = txt.first().context("DNS query failed")?;
|
||||
let x = x.to_lowercase();
|
||||
let (protocol_type, domain) = if let Some(v) = x.strip_prefix("udp://") {
|
||||
(ProtocolType::Quic, v)
|
||||
} else if let Some(v) = x.strip_prefix("quic://") {
|
||||
(ProtocolType::Quic, v)
|
||||
} else if let Some(v) = x.strip_prefix("tcp://") {
|
||||
(ProtocolType::TlsTcp, v)
|
||||
} else if let Some(v) = x.strip_prefix("ws://") {
|
||||
(ProtocolType::TlsTcp, v)
|
||||
} else if let Some(v) = x.strip_prefix("wss://") {
|
||||
(ProtocolType::TlsTcp, v)
|
||||
} else {
|
||||
(ProtocolType::TlsTcp, x.as_str())
|
||||
};
|
||||
(protocol_type, domain.to_owned())
|
||||
}
|
||||
v => (v, self.server_addr.address.to_string()),
|
||||
};
|
||||
let server_addr =
|
||||
crate::utils::dns_query::dns_query_one(&server_domain, &vec![], &None).await?;
|
||||
let server_domain = strip_port(&server_domain).to_owned();
|
||||
Ok(ConnectConfig {
|
||||
protocol_type,
|
||||
server_addr,
|
||||
server_domain,
|
||||
cert_mode: self.cert_mode.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
fn strip_port(addr: &str) -> &str {
|
||||
if let Some(stripped) = addr.strip_prefix('[')
|
||||
&& let Some(pos) = stripped.find(']')
|
||||
{
|
||||
return &stripped[..pos];
|
||||
}
|
||||
|
||||
if addr.contains(':') && !addr.contains('.') && addr.matches(':').count() > 1 {
|
||||
return addr;
|
||||
}
|
||||
|
||||
if let Some((host, port)) = addr.rsplit_once(':')
|
||||
&& port.chars().all(|c| c.is_ascii_digit())
|
||||
{
|
||||
return host;
|
||||
}
|
||||
|
||||
addr
|
||||
}
|
||||
impl ConnectConfig {
|
||||
pub fn server_addr(&self) -> SocketAddr {
|
||||
self.server_addr
|
||||
}
|
||||
pub fn server_name(&self) -> &String {
|
||||
&self.server_domain
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
use crate::protocol::ip_packet_protocol::NetPacket;
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::server::transport::config::{ConnectConfig, ProtocolType};
|
||||
use crate::tunnel_core::server::transport::quic::QuicTransport;
|
||||
use crate::tunnel_core::server::transport::tcp::TlsTcpTransport;
|
||||
use crate::tunnel_core::server::transport::wss::WssTransport;
|
||||
use anyhow::{Context, bail};
|
||||
use bytes::Bytes;
|
||||
use std::time::Duration;
|
||||
|
||||
pub mod config;
|
||||
pub(crate) mod quic;
|
||||
pub(crate) mod tcp;
|
||||
pub(crate) mod wss;
|
||||
|
||||
#[derive(Default)]
|
||||
pub(crate) enum TransportClient {
|
||||
Quic(QuicTransport),
|
||||
TlsTcp(TlsTcpTransport),
|
||||
Wss(WssTransport),
|
||||
#[default]
|
||||
Pending,
|
||||
}
|
||||
|
||||
impl TransportClient {
|
||||
pub fn new() -> Self {
|
||||
TransportClient::default()
|
||||
}
|
||||
pub fn disconnect(&mut self) {
|
||||
match self {
|
||||
TransportClient::Quic(c) => c.disconnect(),
|
||||
TransportClient::TlsTcp(c) => c.disconnect(),
|
||||
TransportClient::Wss(c) => c.disconnect(),
|
||||
TransportClient::Pending => {}
|
||||
};
|
||||
*self = TransportClient::Pending;
|
||||
}
|
||||
pub async fn connect_timeout(
|
||||
&mut self,
|
||||
config: &ConnectConfig,
|
||||
timeout: Duration,
|
||||
) -> anyhow::Result<()> {
|
||||
tokio::time::timeout(timeout, self.connect(config))
|
||||
.await
|
||||
.context("timeout")?
|
||||
}
|
||||
pub async fn connect(&mut self, config: &ConnectConfig) -> anyhow::Result<()> {
|
||||
match self {
|
||||
TransportClient::Quic(c) => c.connect(config).await?,
|
||||
TransportClient::TlsTcp(c) => c.connect(config).await?,
|
||||
TransportClient::Wss(c) => c.connect(config).await?,
|
||||
TransportClient::Pending => match config.protocol_type {
|
||||
ProtocolType::Quic => {
|
||||
let mut transport = QuicTransport::new();
|
||||
transport.connect(config).await?;
|
||||
*self = TransportClient::Quic(transport);
|
||||
}
|
||||
ProtocolType::TlsTcp => {
|
||||
let mut transport = TlsTcpTransport::new();
|
||||
transport.connect(config).await?;
|
||||
*self = TransportClient::TlsTcp(transport);
|
||||
}
|
||||
ProtocolType::Wss => {
|
||||
let mut transport = WssTransport::new();
|
||||
transport.connect(config).await?;
|
||||
*self = TransportClient::Wss(transport);
|
||||
}
|
||||
ProtocolType::Dynamic => {
|
||||
bail!("unreachable connect")
|
||||
}
|
||||
},
|
||||
};
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn send(&mut self, buf: Bytes) -> anyhow::Result<()> {
|
||||
match self {
|
||||
TransportClient::Quic(c) => c.send(buf).await,
|
||||
TransportClient::TlsTcp(c) => c.send(buf).await,
|
||||
TransportClient::Wss(c) => c.send(buf).await,
|
||||
TransportClient::Pending => {
|
||||
bail!("Not connected");
|
||||
}
|
||||
}
|
||||
}
|
||||
pub async fn next(&mut self) -> anyhow::Result<TransmissionBytes> {
|
||||
match self {
|
||||
TransportClient::Quic(c) => c.next().await,
|
||||
TransportClient::TlsTcp(c) => c.next().await,
|
||||
TransportClient::Wss(c) => c.next().await,
|
||||
|
||||
TransportClient::Pending => {
|
||||
bail!("Not connected");
|
||||
}
|
||||
}
|
||||
}
|
||||
pub async fn next_timeout(&mut self, timeout: Duration) -> anyhow::Result<TransmissionBytes> {
|
||||
tokio::time::timeout(timeout, self.next())
|
||||
.await
|
||||
.context("timeout")?
|
||||
}
|
||||
pub async fn send_turn(&mut self, buf: NetPacket<TransmissionBytes>) -> anyhow::Result<()> {
|
||||
self.send(buf.into_buffer().into_bytes().freeze()).await
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tls::verifier::CertValidationMode;
|
||||
use crate::tunnel_core::server::transport::config::ConnectConfig;
|
||||
use anyhow::{Context, bail};
|
||||
use bytes::Bytes;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use quinn::{ClientConfig, Endpoint, RecvStream, SendStream};
|
||||
use std::sync::Arc;
|
||||
use tokio_util::codec::{FramedRead, FramedWrite, LengthDelimitedCodec};
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct QuicTransport {
|
||||
framed: Option<(
|
||||
FramedWrite<SendStream, LengthDelimitedCodec>,
|
||||
FramedRead<RecvStream, LengthDelimitedCodec>,
|
||||
)>,
|
||||
}
|
||||
impl QuicTransport {
|
||||
pub fn new() -> Self {
|
||||
Default::default()
|
||||
}
|
||||
pub fn disconnect(&mut self) {
|
||||
self.framed = None;
|
||||
}
|
||||
pub async fn connect(&mut self, config: &ConnectConfig) -> anyhow::Result<()> {
|
||||
if self.framed.is_some() {
|
||||
bail!("Already connected");
|
||||
}
|
||||
let (w, r) = connect_quic(config).await?;
|
||||
self.framed = Some((w, r));
|
||||
Ok(())
|
||||
}
|
||||
pub async fn send(&mut self, buf: Bytes) -> anyhow::Result<()> {
|
||||
let Some((w, _r)) = self.framed.as_mut() else {
|
||||
bail!("Not connected");
|
||||
};
|
||||
w.send(buf).await.context("send to server failed")
|
||||
}
|
||||
pub async fn next(&mut self) -> anyhow::Result<TransmissionBytes> {
|
||||
let Some((_w, r)) = self.framed.as_mut() else {
|
||||
bail!("Not connected");
|
||||
};
|
||||
r.next()
|
||||
.await
|
||||
.context("EOF")?
|
||||
.context("receive from server failed")
|
||||
.map(TransmissionBytes::from)
|
||||
}
|
||||
}
|
||||
pub async fn connect_quic(
|
||||
config: &ConnectConfig,
|
||||
) -> anyhow::Result<(
|
||||
FramedWrite<SendStream, LengthDelimitedCodec>,
|
||||
FramedRead<RecvStream, LengthDelimitedCodec>,
|
||||
)> {
|
||||
let server_addr = config.server_addr();
|
||||
let server_name = config.server_name();
|
||||
let quic_config = create_client_config(&config.cert_mode)?;
|
||||
let mut endpoint = match Endpoint::client((std::net::Ipv6Addr::UNSPECIFIED, 0).into()) {
|
||||
Ok(endpoint) => endpoint,
|
||||
Err(e) => {
|
||||
log::warn!("Failed to create QUIC endpoint: {}", e);
|
||||
Endpoint::client((std::net::Ipv4Addr::UNSPECIFIED, 0).into())
|
||||
.context("Failed to create QUIC endpoint")?
|
||||
}
|
||||
};
|
||||
|
||||
endpoint.set_default_client_config(quic_config);
|
||||
let connection = endpoint
|
||||
.connect(server_addr, server_name)?
|
||||
.await
|
||||
.context("Failed to establish QUIC connection")?;
|
||||
let (send_stream, recv_stream) = connection
|
||||
.open_bi()
|
||||
.await
|
||||
.context("Failed to open bidirectional stream")?;
|
||||
let framed_write = FramedWrite::new(send_stream, LengthDelimitedCodec::new());
|
||||
let framed_read = FramedRead::new(recv_stream, LengthDelimitedCodec::new());
|
||||
Ok((framed_write, framed_read))
|
||||
}
|
||||
|
||||
fn create_client_config(cert_mode: &CertValidationMode) -> anyhow::Result<ClientConfig> {
|
||||
let config = cert_mode.create_tls_client_config()?;
|
||||
let client_config = ClientConfig::new(Arc::new(
|
||||
quinn::crypto::rustls::QuicClientConfig::try_from(config)
|
||||
.context("Failed to create QUIC client config")?,
|
||||
));
|
||||
|
||||
Ok(client_config)
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::server::transport::config::ConnectConfig;
|
||||
use anyhow::{Context, bail};
|
||||
use bytes::Bytes;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use std::sync::Arc;
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_rustls::{TlsConnector, client::TlsStream};
|
||||
use tokio_util::codec::{Framed, LengthDelimitedCodec};
|
||||
|
||||
type TlsTcpStream = TlsStream<TcpStream>;
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct TlsTcpTransport {
|
||||
framed: Option<Framed<TlsTcpStream, LengthDelimitedCodec>>,
|
||||
}
|
||||
|
||||
impl TlsTcpTransport {
|
||||
pub fn new() -> Self {
|
||||
Default::default()
|
||||
}
|
||||
pub fn disconnect(&mut self) {
|
||||
self.framed = None;
|
||||
}
|
||||
pub async fn connect(&mut self, config: &ConnectConfig) -> anyhow::Result<()> {
|
||||
if self.framed.is_some() {
|
||||
bail!("Already connected");
|
||||
}
|
||||
let framed = connect_tls_tcp(config).await?;
|
||||
self.framed = Some(framed);
|
||||
Ok(())
|
||||
}
|
||||
pub async fn send(&mut self, buf: Bytes) -> anyhow::Result<()> {
|
||||
let Some(framed) = self.framed.as_mut() else {
|
||||
bail!("Not connected");
|
||||
};
|
||||
framed.send(buf).await.context("send to server failed")
|
||||
}
|
||||
pub async fn next(&mut self) -> anyhow::Result<TransmissionBytes> {
|
||||
let Some(framed) = self.framed.as_mut() else {
|
||||
bail!("Not connected");
|
||||
};
|
||||
framed
|
||||
.next()
|
||||
.await
|
||||
.context("EOF")?
|
||||
.context("receive from server failed")
|
||||
.map(TransmissionBytes::from)
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn connect_tls_tcp(
|
||||
config: &ConnectConfig,
|
||||
) -> anyhow::Result<Framed<TlsTcpStream, LengthDelimitedCodec>> {
|
||||
let server_addr = config.server_addr();
|
||||
let server_name = config.server_name().clone();
|
||||
|
||||
let rustls_config = config.cert_mode.create_tls_client_config()?;
|
||||
let connector = TlsConnector::from(Arc::new(rustls_config));
|
||||
|
||||
let tcp_stream = TcpStream::connect(server_addr)
|
||||
.await
|
||||
.context("Failed to establish underlying TCP connection")?;
|
||||
if let Err(e) = tcp_stream.set_nodelay(true) {
|
||||
log::error!("Failed to set TCP_NODELAY: {}", e);
|
||||
}
|
||||
let dns_name = server_name
|
||||
.try_into()
|
||||
.context("Invalid server name for TLS")?;
|
||||
|
||||
let tls_stream = connector
|
||||
.connect(dns_name, tcp_stream)
|
||||
.await
|
||||
.context("Failed to perform TLS handshake")?;
|
||||
|
||||
let framed = Framed::new(tls_stream, LengthDelimitedCodec::new());
|
||||
|
||||
Ok(framed)
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
use crate::protocol::transmission::TransmissionBytes;
|
||||
use crate::tunnel_core::server::transport::config::ConnectConfig;
|
||||
use anyhow::{Context, bail};
|
||||
use bytes::Bytes;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use std::sync::Arc;
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_rustls::{TlsConnector, client::TlsStream};
|
||||
use tokio_tungstenite::{WebSocketStream, client_async, tungstenite::Message};
|
||||
|
||||
type WssStream = WebSocketStream<TlsStream<TcpStream>>;
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct WssTransport {
|
||||
stream: Option<WssStream>,
|
||||
}
|
||||
|
||||
impl WssTransport {
|
||||
pub fn new() -> Self {
|
||||
Default::default()
|
||||
}
|
||||
pub fn disconnect(&mut self) {
|
||||
self.stream = None;
|
||||
}
|
||||
pub async fn connect(&mut self, config: &ConnectConfig) -> anyhow::Result<()> {
|
||||
if self.stream.is_some() {
|
||||
bail!("Already connected");
|
||||
}
|
||||
let stream = connect_wss(config).await?;
|
||||
self.stream = Some(stream);
|
||||
Ok(())
|
||||
}
|
||||
pub async fn send(&mut self, buf: Bytes) -> anyhow::Result<()> {
|
||||
let Some(framed) = self.stream.as_mut() else {
|
||||
bail!("Not connected");
|
||||
};
|
||||
framed
|
||||
.send(Message::Binary(buf))
|
||||
.await
|
||||
.context("send to server failed")
|
||||
}
|
||||
pub async fn next(&mut self) -> anyhow::Result<TransmissionBytes> {
|
||||
let Some(framed) = self.stream.as_mut() else {
|
||||
bail!("Not connected");
|
||||
};
|
||||
loop {
|
||||
let message = framed
|
||||
.next()
|
||||
.await
|
||||
.context("EOF")?
|
||||
.context("receive from server failed")?;
|
||||
match message {
|
||||
Message::Binary(buf) => {
|
||||
return Ok(TransmissionBytes::from(buf));
|
||||
}
|
||||
Message::Close(_) => {
|
||||
bail!("Disconnected");
|
||||
}
|
||||
_ => {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn connect_wss(config: &ConnectConfig) -> anyhow::Result<WssStream> {
|
||||
let server_addr = config.server_addr();
|
||||
let server_name = config.server_name().clone();
|
||||
|
||||
let rustls_config = config.cert_mode.create_tls_client_config()?;
|
||||
let connector = TlsConnector::from(Arc::new(rustls_config));
|
||||
|
||||
let tcp_stream = TcpStream::connect(server_addr)
|
||||
.await
|
||||
.context("Failed to establish underlying TCP connection")?;
|
||||
if let Err(e) = tcp_stream.set_nodelay(true) {
|
||||
log::error!("Failed to set TCP_NODELAY: {}", e);
|
||||
}
|
||||
let url = format!("wss://{}", server_name);
|
||||
|
||||
let dns_name = server_name
|
||||
.try_into()
|
||||
.context("Invalid server name for TLS")?;
|
||||
|
||||
let tls_stream = connector
|
||||
.connect(dns_name, tcp_stream)
|
||||
.await
|
||||
.context("Failed to perform TLS handshake")?;
|
||||
|
||||
let (ws_stream, _response) = client_async(url, tls_stream)
|
||||
.await
|
||||
.context("Failed to perform WebSocket handshake")?;
|
||||
|
||||
Ok(ws_stream)
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
use anyhow::Context;
|
||||
use std::fs;
|
||||
use std::path::Path;
|
||||
pub fn get_device_id() -> anyhow::Result<String> {
|
||||
match machine_uid::get() {
|
||||
Ok(id) => return Ok(id),
|
||||
Err(e) => {
|
||||
log::warn!("Failed to get system ID: {}. Using fallback.", e);
|
||||
}
|
||||
}
|
||||
|
||||
get_fallback_id()
|
||||
}
|
||||
fn get_fallback_id() -> anyhow::Result<String> {
|
||||
let path = Path::new("device_id");
|
||||
|
||||
if let Ok(content) = fs::read_to_string(path) {
|
||||
let id = content.trim();
|
||||
if !id.is_empty() {
|
||||
return Ok(id.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let new_id = uuid::Uuid::new_v4().to_string();
|
||||
|
||||
fs::write(path, &new_id).context("Failed to write device_id file")?;
|
||||
|
||||
Ok(new_id)
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
use anyhow::{Context, anyhow};
|
||||
use dns_parser::{Builder, Packet, QueryClass, QueryType, RData, ResponseCode};
|
||||
use rand::seq::SliceRandom;
|
||||
use std::io;
|
||||
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
|
||||
use std::str::FromStr;
|
||||
use std::time::Duration;
|
||||
use tokio::net::UdpSocket;
|
||||
|
||||
use rust_p2p_core::socket::LocalInterface;
|
||||
|
||||
pub async fn dns_query_txt(
|
||||
domain: &str,
|
||||
mut name_servers: Vec<String>,
|
||||
default_interface: &Option<LocalInterface>,
|
||||
) -> io::Result<Vec<String>> {
|
||||
let mut err: Option<io::Error> = None;
|
||||
if name_servers.is_empty() {
|
||||
name_servers.push("223.5.5.5:53".into());
|
||||
name_servers.push("114.114.114.114:53".into());
|
||||
}
|
||||
for name_server in name_servers {
|
||||
match txt_dns(domain, name_server, default_interface).await {
|
||||
Ok(addr) => {
|
||||
if !addr.is_empty() {
|
||||
return Ok(addr);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
err.replace(e);
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if let Some(e) = err {
|
||||
Err(e)
|
||||
} else {
|
||||
Err(io::Error::other(format!("DNS query failed {domain:?}")))
|
||||
}
|
||||
}
|
||||
pub async fn dns_query_one(
|
||||
domain: &str,
|
||||
name_servers: &Vec<String>,
|
||||
default_interface: &Option<LocalInterface>,
|
||||
) -> anyhow::Result<SocketAddr> {
|
||||
let mut vec = dns_query_all(domain, name_servers, default_interface).await?;
|
||||
vec.shuffle(&mut rand::rng());
|
||||
vec.pop().context("DNS query failed")
|
||||
}
|
||||
pub async fn dns_query_all(
|
||||
domain: &str,
|
||||
name_servers: &Vec<String>,
|
||||
default_interface: &Option<LocalInterface>,
|
||||
) -> anyhow::Result<Vec<SocketAddr>> {
|
||||
match SocketAddr::from_str(domain) {
|
||||
Ok(addr) => Ok(vec![addr]),
|
||||
Err(_) => {
|
||||
if name_servers.is_empty() {
|
||||
let addrs: Vec<SocketAddr> = tokio::net::lookup_host(domain)
|
||||
.await
|
||||
.map_err(|e| io::Error::other(format!("DNS query failed: {domain:?},{e:?}")))?
|
||||
.collect();
|
||||
return Ok(addrs);
|
||||
}
|
||||
|
||||
let mut err: Option<io::Error> = None;
|
||||
for name_server in name_servers {
|
||||
let end_index = domain
|
||||
.rfind(':')
|
||||
.ok_or_else(|| io::Error::other(format!("not port: {domain:?}")))?;
|
||||
let host = &domain[..end_index];
|
||||
let port = u16::from_str(&domain[end_index + 1..])
|
||||
.map_err(|_| io::Error::other(format!("not port: {domain:?}")))?;
|
||||
let th1 = {
|
||||
let host = host.to_string();
|
||||
let name_server = name_server.clone();
|
||||
let default_interface = default_interface.clone();
|
||||
tokio::spawn(a_dns(host, name_server, default_interface.clone()))
|
||||
};
|
||||
let th2 = {
|
||||
let host = host.to_string();
|
||||
let name_server = name_server.clone();
|
||||
let default_interface = default_interface.clone();
|
||||
tokio::spawn(aaaa_dns(host, name_server, default_interface.clone()))
|
||||
};
|
||||
let mut addr = Vec::new();
|
||||
match th1.await? {
|
||||
Ok(rs) => {
|
||||
for ip in rs {
|
||||
addr.push(SocketAddr::new(ip.into(), port));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
err.replace(e);
|
||||
}
|
||||
}
|
||||
match th2.await? {
|
||||
Ok(rs) => {
|
||||
for ip in rs {
|
||||
addr.push(SocketAddr::new(ip.into(), port));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
if addr.is_empty() {
|
||||
err.replace(e);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
if addr.is_empty() {
|
||||
continue;
|
||||
}
|
||||
return Ok(addr);
|
||||
}
|
||||
if let Some(e) = err {
|
||||
Err(e.into())
|
||||
} else {
|
||||
Err(anyhow!("DNS query failed {domain:?}"))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn query<'a>(
|
||||
udp: &UdpSocket,
|
||||
domain: &str,
|
||||
name_server: SocketAddr,
|
||||
record_type: QueryType,
|
||||
buf: &'a mut [u8],
|
||||
) -> io::Result<Packet<'a>> {
|
||||
let mut builder = Builder::new_query(1, true);
|
||||
builder.add_question(domain, false, record_type, QueryClass::IN);
|
||||
let packet = builder.build().unwrap();
|
||||
|
||||
udp.connect(name_server).await?;
|
||||
let mut count = 0;
|
||||
let len = loop {
|
||||
udp.send(&packet).await?;
|
||||
|
||||
match tokio::time::timeout(Duration::from_secs(3), udp.recv(buf)).await {
|
||||
Ok(len) => {
|
||||
break len?;
|
||||
}
|
||||
Err(_) => {
|
||||
count += 1;
|
||||
if count < 3 {
|
||||
continue;
|
||||
}
|
||||
Err(io::Error::other(format!("DNS {name_server:?} recv error ")))?
|
||||
}
|
||||
};
|
||||
};
|
||||
|
||||
let pkt = Packet::parse(&buf[..len]).map_err(|e| {
|
||||
io::Error::other(format!(
|
||||
"domain {domain:?} DNS {name_server:?} data error: {e}"
|
||||
))
|
||||
})?;
|
||||
if pkt.header.response_code != ResponseCode::NoError {
|
||||
return Err(io::Error::other(format!(
|
||||
"response_code {} DNS {:?} domain {:?}",
|
||||
pkt.header.response_code, name_server, domain
|
||||
)));
|
||||
}
|
||||
if pkt.answers.is_empty() {
|
||||
return Err(io::Error::other(format!(
|
||||
"No records received DNS {name_server:?} domain {domain:?}"
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(pkt)
|
||||
}
|
||||
|
||||
pub async fn txt_dns(
|
||||
domain: &str,
|
||||
name_server: String,
|
||||
default_interface: &Option<LocalInterface>,
|
||||
) -> io::Result<Vec<String>> {
|
||||
let name_server: SocketAddr = name_server
|
||||
.parse()
|
||||
.map_err(|e| io::Error::other(format!("dns {name_server} is error :{e:?}")))?;
|
||||
let udp = bind_udp(name_server, default_interface)?;
|
||||
let mut buf = vec![0u8; 65536];
|
||||
let message = query(&udp, domain, name_server, QueryType::TXT, &mut buf).await?;
|
||||
let mut rs = Vec::new();
|
||||
for record in message.answers {
|
||||
if let RData::TXT(txt) = record.data {
|
||||
for x in txt.iter() {
|
||||
let txt = std::str::from_utf8(x)
|
||||
.map_err(|_| io::Error::other("record type txt is not string"))?;
|
||||
rs.push(txt.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(rs)
|
||||
}
|
||||
|
||||
fn bind_udp(
|
||||
name_server: SocketAddr,
|
||||
default_interface: &Option<LocalInterface>,
|
||||
) -> io::Result<UdpSocket> {
|
||||
let addr: SocketAddr = if name_server.is_ipv4() {
|
||||
"0.0.0.0:0"
|
||||
.parse()
|
||||
.expect("valid IPv4 socket address literal")
|
||||
} else {
|
||||
"[::]:0".parse().expect("valid IPv6 socket address literal")
|
||||
};
|
||||
let socket = rust_p2p_core::socket::bind_udp(addr, default_interface.as_ref())?;
|
||||
UdpSocket::from_std(socket.into())
|
||||
}
|
||||
|
||||
pub async fn a_dns(
|
||||
domain: String,
|
||||
name_server: String,
|
||||
default_interface: Option<LocalInterface>,
|
||||
) -> io::Result<Vec<Ipv4Addr>> {
|
||||
let name_server: SocketAddr = name_server
|
||||
.parse()
|
||||
.map_err(|e| io::Error::other(format!("dns {name_server} is error :{e:?}")))?;
|
||||
let udp = bind_udp(name_server, &default_interface)?;
|
||||
let mut buf = vec![0u8; 65536];
|
||||
let message = query(&udp, &domain, name_server, QueryType::A, &mut buf).await?;
|
||||
let mut rs = Vec::new();
|
||||
for record in message.answers {
|
||||
if let RData::A(a) = record.data {
|
||||
rs.push(a.0);
|
||||
}
|
||||
}
|
||||
Ok(rs)
|
||||
}
|
||||
|
||||
pub async fn aaaa_dns(
|
||||
domain: String,
|
||||
name_server: String,
|
||||
default_interface: Option<LocalInterface>,
|
||||
) -> io::Result<Vec<Ipv6Addr>> {
|
||||
let name_server: SocketAddr = name_server
|
||||
.parse()
|
||||
.map_err(|e| io::Error::other(format!("dns {name_server} is error :{e:?}")))?;
|
||||
let udp = bind_udp(name_server, &default_interface)?;
|
||||
let mut buf = vec![0u8; 65536];
|
||||
let message = query(&udp, &domain, name_server, QueryType::AAAA, &mut buf).await?;
|
||||
let mut rs = Vec::new();
|
||||
for record in message.answers {
|
||||
if let RData::AAAA(a) = record.data {
|
||||
rs.push(a.0);
|
||||
}
|
||||
}
|
||||
Ok(rs)
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
pub mod device_id;
|
||||
pub(crate) mod dns_query;
|
||||
pub mod task_control;
|
||||
pub(crate) mod time {
|
||||
pub fn now_ts_ms() -> i64 {
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_millis() as i64
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,255 @@
|
||||
use parking_lot::Mutex;
|
||||
use std::collections::HashMap;
|
||||
use std::future::Future;
|
||||
use std::sync::{Arc, Weak};
|
||||
use tokio::sync::Notify;
|
||||
use tokio::task::{Id, JoinHandle};
|
||||
|
||||
struct TaskGroupState {
|
||||
stopped: bool,
|
||||
tasks: HashMap<Id, JoinHandle<()>>,
|
||||
}
|
||||
|
||||
struct TaskGroupInner {
|
||||
state: Mutex<TaskGroupState>,
|
||||
all_stopped_notify: Notify,
|
||||
}
|
||||
|
||||
impl TaskGroupInner {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
state: Mutex::new(TaskGroupState {
|
||||
stopped: false,
|
||||
tasks: HashMap::new(),
|
||||
}),
|
||||
all_stopped_notify: Notify::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn spawn<F>(self: &Arc<Self>, f: F) -> Option<Id>
|
||||
where
|
||||
F: Future + Send + 'static,
|
||||
F::Output: Send + 'static,
|
||||
{
|
||||
let mut state = self.state.lock();
|
||||
if state.stopped {
|
||||
return None;
|
||||
}
|
||||
|
||||
let guard = TaskGuard {
|
||||
inner: Arc::downgrade(self),
|
||||
};
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
let _guard = guard;
|
||||
f.await;
|
||||
});
|
||||
|
||||
let task_id = handle.id();
|
||||
state.tasks.insert(task_id, handle);
|
||||
Some(task_id)
|
||||
}
|
||||
|
||||
fn stop(&self) {
|
||||
let mut state = self.state.lock();
|
||||
state.stopped = true;
|
||||
for (_, handle) in state.tasks.drain() {
|
||||
handle.abort();
|
||||
}
|
||||
}
|
||||
|
||||
fn is_stopped(&self) -> bool {
|
||||
self.state.lock().stopped
|
||||
}
|
||||
|
||||
fn remove_task(&self, task_id: Id) {
|
||||
let all_stopped = {
|
||||
let mut state = self.state.lock();
|
||||
state.tasks.remove(&task_id);
|
||||
if state.tasks.is_empty() {
|
||||
state.stopped = true;
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
};
|
||||
if all_stopped {
|
||||
self.all_stopped_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
async fn abort_task(&self, task_id: Id) {
|
||||
let handle = self.state.lock().tasks.remove(&task_id);
|
||||
if let Some(handle) = handle {
|
||||
handle.abort();
|
||||
_ = handle.await;
|
||||
}
|
||||
}
|
||||
|
||||
async fn join_all(&self) {
|
||||
let tasks = std::mem::take(&mut self.state.lock().tasks);
|
||||
for (_, h) in tasks {
|
||||
let _ = h.await;
|
||||
}
|
||||
}
|
||||
|
||||
fn all_tasks_stopped(&self) -> bool {
|
||||
let state = self.state.lock();
|
||||
state.stopped && state.tasks.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TaskGroupInner {
|
||||
fn drop(&mut self) {
|
||||
self.stop();
|
||||
}
|
||||
}
|
||||
|
||||
struct TaskGuard {
|
||||
inner: Weak<TaskGroupInner>,
|
||||
}
|
||||
|
||||
impl Drop for TaskGuard {
|
||||
fn drop(&mut self) {
|
||||
if let Some(inner) = self.inner.upgrade() {
|
||||
let task_id = tokio::task::id();
|
||||
inner.remove_task(task_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct TaskGroup {
|
||||
inner: Arc<TaskGroupInner>,
|
||||
}
|
||||
|
||||
impl TaskGroup {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
inner: Arc::new(TaskGroupInner::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn stop(&self) {
|
||||
self.inner.stop();
|
||||
}
|
||||
|
||||
pub fn is_stopped(&self) -> bool {
|
||||
self.inner.is_stopped()
|
||||
}
|
||||
|
||||
pub fn spawn<F>(&self, f: F) -> SubTask
|
||||
where
|
||||
F: Future + Send + 'static,
|
||||
F::Output: Send + 'static,
|
||||
{
|
||||
match self.inner.spawn(f) {
|
||||
Some(task_id) => SubTask::new(task_id, Arc::downgrade(&self.inner)),
|
||||
None => SubTask::empty(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn join_all(&self) {
|
||||
self.inner.join_all().await;
|
||||
}
|
||||
|
||||
pub async fn wait_all_stopped(&self) {
|
||||
loop {
|
||||
if self.inner.all_tasks_stopped() {
|
||||
return;
|
||||
}
|
||||
self.inner.all_stopped_notify.notified().await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct SubTask {
|
||||
task_id: Option<Id>,
|
||||
inner: Weak<TaskGroupInner>,
|
||||
}
|
||||
|
||||
impl SubTask {
|
||||
fn new(task_id: Id, inner: Weak<TaskGroupInner>) -> Self {
|
||||
Self {
|
||||
task_id: Some(task_id),
|
||||
inner,
|
||||
}
|
||||
}
|
||||
|
||||
fn empty() -> Self {
|
||||
Self {
|
||||
task_id: None,
|
||||
inner: Weak::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn stop(&self) {
|
||||
if let Some(task_id) = self.task_id
|
||||
&& let Some(inner) = self.inner.upgrade()
|
||||
{
|
||||
inner.abort_task(task_id).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_running(&self) -> bool {
|
||||
if let Some(task_id) = self.task_id
|
||||
&& let Some(inner) = self.inner.upgrade()
|
||||
{
|
||||
return inner.state.lock().tasks.contains_key(&task_id);
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
pub fn id(&self) -> Option<Id> {
|
||||
self.task_id
|
||||
}
|
||||
}
|
||||
#[derive(Clone, Default)]
|
||||
pub struct TaskGroupManager {
|
||||
task_group: Arc<Mutex<Option<TaskGroup>>>,
|
||||
}
|
||||
|
||||
impl TaskGroupManager {
|
||||
pub fn new() -> Self {
|
||||
TaskGroupManager::default()
|
||||
}
|
||||
|
||||
pub fn is_running(&self) -> bool {
|
||||
self.task_group.lock().is_some()
|
||||
}
|
||||
|
||||
pub fn is_stopped(&self) -> bool {
|
||||
self.task_group.lock().is_none()
|
||||
}
|
||||
|
||||
pub fn create_task(&self) -> anyhow::Result<(TaskGroup, TaskGroupGuard)> {
|
||||
let mut guard = self.task_group.lock();
|
||||
if guard.is_some() {
|
||||
anyhow::bail!("运行中")
|
||||
}
|
||||
|
||||
let task_group = TaskGroup::new();
|
||||
guard.replace(task_group.clone());
|
||||
let stop_guard = TaskGroupGuard {
|
||||
task_group: self.task_group.clone(),
|
||||
};
|
||||
Ok((task_group, stop_guard))
|
||||
}
|
||||
|
||||
pub fn stop(&self) {
|
||||
let option = self.task_group.lock();
|
||||
if let Some(task_group) = option.as_ref() {
|
||||
task_group.stop();
|
||||
}
|
||||
}
|
||||
}
|
||||
pub struct TaskGroupGuard {
|
||||
task_group: Arc<Mutex<Option<TaskGroup>>>,
|
||||
}
|
||||
impl Drop for TaskGroupGuard {
|
||||
fn drop(&mut self) {
|
||||
if let Some(task_group) = self.task_group.lock().take() {
|
||||
task_group.stop();
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user