This commit is contained in:
lbl
2026-02-10 18:20:39 +08:00
parent 71da72da34
commit b4301a8106
261 changed files with 20179 additions and 21547 deletions
+94
View File
@@ -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()
}
}
+132
View File
@@ -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][..]);
}
}
+45
View File
@@ -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)
}
}
+101
View File
@@ -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(),
}
}
}
+632
View File
@@ -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
}
}
+147
View File
@@ -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();
}
}
+353
View File
@@ -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();
}
}
+169
View File
@@ -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[..]);
}
}
+47
View File
@@ -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(())
}
}
+72
View File
@@ -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(())
}
}
+75
View File
@@ -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))
}
+68
View File
@@ -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(())
}
+302
View File
@@ -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);
}
}
}
}
+247
View File
@@ -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(())
}
+5
View File
@@ -0,0 +1,5 @@
mod decoder;
mod encoder;
pub(crate) use decoder::FecDecoder;
pub(crate) use encoder::FecEncoder;
+15
View File
@@ -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;
+139
View File
@@ -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(())
}
+225
View File
@@ -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
}
}
+81
View File
@@ -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(())
}
+188
View File
@@ -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(())
}
+114
View File
@@ -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
}
}
+100
View File
@@ -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(())
}
+98
View File
@@ -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()
}
}
+275
View File
@@ -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,
})
}
}
+302
View File
@@ -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(),
}
}
}
+21
View File
@@ -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 {}
+4
View File
@@ -0,0 +1,4 @@
mod proto {
include!(concat!(env!("OUT_DIR"), "/protocol.rpc.rs"));
}
pub use proto::*;
+266
View File
@@ -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);
}
}
+114
View File
@@ -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());
}
}
+2
View File
@@ -0,0 +1,2 @@
pub(crate) mod cert;
pub mod verifier;
+228
View File
@@ -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)
}
}
+18
View File
@@ -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,
}
}
}
+236
View File
@@ -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(())
}
}
+6
View File
@@ -0,0 +1,6 @@
mod general;
pub use general::*;
mod sender;
pub use sender::*;
pub mod enhanced_tun;
+40
View File
@@ -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(())
}
}
+4
View File
@@ -0,0 +1,4 @@
pub(crate) mod p2p;
pub mod server;
pub(crate) mod outbound;
+286
View File
@@ -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)
}
}
+268
View File
@@ -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);
}
}
}
+5
View File
@@ -0,0 +1,5 @@
pub(crate) mod inbound;
pub(crate) mod outbound;
pub(crate) mod transport;
pub(crate) mod route_table;
+133
View File
@@ -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
}
}
+265
View File
@@ -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!(),
}
}
+332
View File
@@ -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()));
}
}
}
+5
View File
@@ -0,0 +1,5 @@
pub(crate) mod connection_manager;
pub(crate) mod inbound;
pub(crate) mod outbound;
pub(crate) mod rpc;
pub mod transport;
+329
View File
@@ -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(())
}
+150
View File
@@ -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)
}
+29
View File
@@ -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)
}
+251
View File
@@ -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)
}
+11
View File
@@ -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
}
}
+255
View File
@@ -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();
}
}
}