Compare commits

...
20 Commits
Author SHA1 Message Date
lubeilin ad8fecc319 去除命令sudo 2023-09-04 21:32:38 +08:00
lubeilin 27ae9a89da 避免路由切换时的抖动 2023-09-04 20:35:56 +08:00
lubeilin b3a4a4de5e 完善日志输出等 2023-09-04 20:35:39 +08:00
lubeilin ca76c35f6a linux固定网卡名称,启动时删除网卡 2023-09-04 20:35:06 +08:00
lubeilin 073c820da6 增加打洞选项 2023-09-04 20:34:43 +08:00
lubeilin 954f0d2d05 完善1.2.2 2023-09-03 20:58:48 +08:00
lubeilin a943f5bffc 延迟切换NAT类型 2023-09-03 17:24:25 +08:00
lubeilin c3cff7c5b5 cargo fmt 2023-09-03 11:39:30 +08:00
lubeilin acb5a8a325 优化启动速度 2023-09-03 11:28:40 +08:00
lubeilin 1a4e375dbf 支持ipv6 2023-09-02 23:57:32 +08:00
lubeilin eec7d73ebe 可选数据指纹校验、支持ecb算法 2023-09-01 23:36:04 +08:00
lubeilin 96fb8c881d 修改jni模块 2023-08-30 21:54:30 +08:00
lubeilin aeebbd18fd 增加异常日志 2023-08-30 21:53:49 +08:00
lubeilin baa71a51eb 修复广播问题 2023-08-30 21:53:31 +08:00
lubeilin 959f2aa783 修改版本 2023-08-29 21:15:48 +08:00
lubeilin 99b2aa9522 修改linux上kill不退出的问题 2023-08-29 21:15:39 +08:00
lubeilin 561fa9f8fe Merge remote-tracking branch 'origin/dev' into dev 2023-08-29 21:03:36 +08:00
lubeilin ad9dd6a7f7 增加tun创建失败的说明 2023-08-29 21:03:27 +08:00
lbl8603 65758eb94c Update README.md 2023-08-29 10:39:03 +08:00
lubeilin fb7ccf4d11 增加参数说明 2023-08-29 00:21:25 +08:00
103 changed files with 3703 additions and 2106 deletions
+14 -14
View File
@@ -110,40 +110,40 @@ jobs:
# some additional configuration for cross-compilation on linux # some additional configuration for cross-compilation on linux
cat >>~/.cargo/config <<EOF cat >>~/.cargo/config <<EOF
[target.x86_64-unknown-linux-musl] [target.x86_64-unknown-linux-musl]
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.aarch64-unknown-linux-gnu] [target.aarch64-unknown-linux-gnu]
linker = "aarch64-linux-gnu-gcc" linker = "aarch64-linux-gnu-gcc"
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.aarch64-unknown-linux-musl] [target.aarch64-unknown-linux-musl]
linker = "aarch64-linux-gnu-gcc" linker = "aarch64-linux-gnu-gcc"
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.armv7-unknown-linux-gnueabihf] [target.armv7-unknown-linux-gnueabihf]
linker = "arm-linux-gnueabihf-gcc" linker = "arm-linux-gnueabihf-gcc"
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.armv7-unknown-linux-musleabihf] [target.armv7-unknown-linux-musleabihf]
linker = "arm-linux-gnueabihf-gcc" linker = "arm-linux-gnueabihf-gcc"
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.arm-unknown-linux-gnueabihf] [target.arm-unknown-linux-gnueabihf]
linker = "arm-linux-gnueabihf-gcc" linker = "arm-linux-gnueabihf-gcc"
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.arm-unknown-linux-musleabihf] [target.arm-unknown-linux-musleabihf]
linker = "arm-linux-gnueabihf-gcc" linker = "arm-linux-gnueabihf-gcc"
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.mipsel-unknown-linux-musl] [target.mipsel-unknown-linux-musl]
linker = "mipsel-linux-gnu-gcc" linker = "mipsel-linux-gnu-gcc"
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.x86_64-pc-windows-msvc] [target.x86_64-pc-windows-msvc]
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.i686-pc-windows-msvc] [target.i686-pc-windows-msvc]
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.x86_64-apple-darwin] [target.x86_64-apple-darwin]
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.aarch64-apple-darwin] [target.aarch64-apple-darwin]
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.i686-unknown-linux-musl] [target.i686-unknown-linux-musl]
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
[target.x86_64-unknown-linux-gnu] [target.x86_64-unknown-linux-gnu]
rustflags = ["-C", "target-feature=+crt-static"] rustflags = ["-C", "target-feature=+crt-static","-C", "strip=symbols","--cfg","aes_armv8"]
EOF EOF
- name: Install rust target - name: Install rust target
run: rustup target add $TARGET run: rustup target add $TARGET
+29 -1
View File
@@ -86,11 +86,39 @@ A virtual network tool (VPN)
- p2p组播/广播 - p2p组播/广播
- 客户端数据加密 - 客户端数据加密
- 服务端数据加密 - 服务端数据加密
### 结构
<details> <summary>展开</summary>
<pre>
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
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|e |s |unused| 版本(4) | 协议(8) | 上层协议(8) |初始ttl(4)|生存时间(4) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| 源ip地址(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| 目的ip地址(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| 数据体(n) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| |
| 指纹(96) |
| |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
注:
1. e为是否加密标志,s为服务端通信包标志,unused占两位未使用;
2. 开启加密时,数据体为加密后的密文(加密方式取决于密码长度和加密模式),
且会存在指纹,指纹使用sha256生成,用于对数据包完整性和真实性的校验
</pre>
</details>
### Todo ### Todo
- 桌面UI(测试中) - 桌面UI(测试中)
- 支持Ipv6 - 支持Ipv6(1.2.2已支持客户端之间的ipv6,待支持客户端和服务端之间的ipv6通信)
### 常见问题 ### 常见问题
<details> <summary>展开</summary> <details> <summary>展开</summary>
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "common" name = "common"
version = "1.2.0" version = "1.2.2"
edition = "2021" edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
+1 -1
View File
@@ -76,4 +76,4 @@ pub fn to_ip(mask: &str) -> Result<u32, String> {
} else { } else {
Err("not netmask".to_string()) Err("not netmask".to_string())
} }
} }
+12 -9
View File
@@ -1,4 +1,3 @@
use std::process::Command; use std::process::Command;
#[cfg(target_os = "windows")] #[cfg(target_os = "windows")]
@@ -7,8 +6,9 @@ pub fn get_unique_identifier() -> Option<String> {
let output = match Command::new("wmic") let output = match Command::new("wmic")
.creation_flags(0x08000000) .creation_flags(0x08000000)
.args(&["csproduct", "get", "UUID"]) .args(&["csproduct", "get", "UUID"])
.output() { .output()
Ok(output) => { output } {
Ok(output) => output,
Err(_) => { Err(_) => {
return None; return None;
} }
@@ -27,8 +27,9 @@ pub fn get_unique_identifier() -> Option<String> {
pub fn get_unique_identifier() -> Option<String> { pub fn get_unique_identifier() -> Option<String> {
let output = match Command::new("ioreg") let output = match Command::new("ioreg")
.args(&["-rd1", "-c", "IOPlatformExpertDevice"]) .args(&["-rd1", "-c", "IOPlatformExpertDevice"])
.output() { .output()
Ok(output) => { output } {
Ok(output) => output,
Err(_) => { Err(_) => {
return None; return None;
} }
@@ -38,7 +39,8 @@ pub fn get_unique_identifier() -> Option<String> {
let identifier = result let identifier = result
.lines() .lines()
.find(|line| line.contains("IOPlatformUUID")) .find(|line| line.contains("IOPlatformUUID"))
.unwrap_or("").trim(); .unwrap_or("")
.trim();
if identifier.is_empty() { if identifier.is_empty() {
None None
} else { } else {
@@ -51,8 +53,9 @@ pub fn get_unique_identifier() -> Option<String> {
let output = match Command::new("dmidecode") let output = match Command::new("dmidecode")
.arg("-s") .arg("-s")
.arg("system-uuid") .arg("system-uuid")
.output() { .output()
Ok(output) => { output } {
Ok(output) => output,
Err(_) => { Err(_) => {
return None; return None;
} }
@@ -65,4 +68,4 @@ pub fn get_unique_identifier() -> Option<String> {
} else { } else {
Some(identifier.to_string()) Some(identifier.to_string())
} }
} }
+1 -1
View File
@@ -1,2 +1,2 @@
pub mod identifier;
pub mod args_parse; pub mod args_parse;
pub mod identifier;
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "vnt-cli" name = "vnt-cli"
version = "1.2.0" version = "1.2.2"
edition = "2021" edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
+13 -1
View File
@@ -50,9 +50,21 @@
### --ip `<IP>` ### --ip `<IP>`
指定虚拟ip,指定的ip不能和其他设备重复,必须有效并且在服务端所属网段下,默认情况由服务端分配 指定虚拟ip,指定的ip不能和其他设备重复,必须有效并且在服务端所属网段下,默认情况由服务端分配
### --par `<parallel>` ### --par `<parallel>`
任务并行度(必须为正整数),默认值为2,该值表示处理网卡读写的任务数,组网设备数较多、处理延迟较大时可适当调大此值 任务并行度(必须为正整数),默认值为1,该值表示处理网卡读写的任务数,组网设备数较多、处理延迟较大时可适当调大此值
### --thread `<thread>` ### --thread `<thread>`
线程数(必须为正整数),默认为核心数乘2,该值表示处理网络读写、ip代理、打洞等用到的线程数,组网设备数较多、处理延迟较大时可适当调大此值 线程数(必须为正整数),默认为核心数乘2,该值表示处理网络读写、ip代理、打洞等用到的线程数,组网设备数较多、处理延迟较大时可适当调大此值
### --model `<model>`
加密模式,可选值 aes_gcm/aes_cbc,默认使用aes_gcm,通常情况使用aes_cbc性能更好
| 密码位数 | model | 加密算法 |
|-------|--------|------------|
| 1~8位 | aes_gcm | AES128-GCM |
| `>=`8 | aes_gcm | AES256-GCM |
| 1~8位 | aes_cbc | AES128-CBC |
| `>=`8 | aes_cbc | AES256-CBC |
### --finger
开启数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启,开启会损耗一部分性能
### --relay ### --relay
禁用p2p,在网络环境很差时,只使用服务器中转效果可能更好(可以配合--tcp参数一起使用) 禁用p2p,在网络环境很差时,只使用服务器中转效果可能更好(可以配合--tcp参数一起使用)
### --list ### --list
+1 -1
View File
@@ -7,4 +7,4 @@ fn main() {
// embed_manifest(new_manifest("vnt") // embed_manifest(new_manifest("vnt")
// .requested_execution_level(ExecutionLevel::RequireAdministrator)).expect("unable to embed manifest file"); // .requested_execution_level(ExecutionLevel::RequireAdministrator)).expect("unable to embed manifest file");
// } // }
} }
+12 -15
View File
@@ -3,7 +3,7 @@ use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4, UdpSocket};
use std::str::FromStr; use std::str::FromStr;
use std::time::Duration; use std::time::Duration;
use crate::command::entity::{DeviceItem, RouteItem, Info}; use crate::command::entity::{DeviceItem, Info, RouteItem};
pub struct CommandClient { pub struct CommandClient {
udp: UdpSocket, udp: UdpSocket,
@@ -17,9 +17,12 @@ impl CommandClient {
} }
let port = std::fs::read_to_string(path_buf)?; let port = std::fs::read_to_string(path_buf)?;
let port = match u16::from_str(&port) { let port = match u16::from_str(&port) {
Ok(port) => { port } Ok(port) => port,
Err(_) => { Err(_) => {
return Err(io::Error::new(io::ErrorKind::Other, "'command-port' file error")); return Err(io::Error::new(
io::ErrorKind::Other,
"'command-port' file error",
));
} }
}; };
let udp = UdpSocket::bind("127.0.0.1:0")?; let udp = UdpSocket::bind("127.0.0.1:0")?;
@@ -38,11 +41,9 @@ impl CommandClient {
let mut buf = [0; 10240]; let mut buf = [0; 10240];
let len = self.udp.recv(&mut buf)?; let len = self.udp.recv(&mut buf)?;
match serde_json::from_slice::<Vec<DeviceItem>>(&buf[..len]) { match serde_json::from_slice::<Vec<DeviceItem>>(&buf[..len]) {
Ok(val) => { Ok(val) => Ok(val),
Ok(val)
}
Err(e) => { Err(e) => {
log::error!("{:?}",e); log::error!("{:?}", e);
Err(io::Error::new(io::ErrorKind::Other, "data error")) Err(io::Error::new(io::ErrorKind::Other, "data error"))
} }
} }
@@ -52,11 +53,9 @@ impl CommandClient {
let mut buf = [0; 10240]; let mut buf = [0; 10240];
let len = self.udp.recv(&mut buf)?; let len = self.udp.recv(&mut buf)?;
match serde_json::from_slice::<Vec<RouteItem>>(&buf[..len]) { match serde_json::from_slice::<Vec<RouteItem>>(&buf[..len]) {
Ok(val) => { Ok(val) => Ok(val),
Ok(val)
}
Err(e) => { Err(e) => {
log::error!("{:?}",e); log::error!("{:?}", e);
Err(io::Error::new(io::ErrorKind::Other, "data error")) Err(io::Error::new(io::ErrorKind::Other, "data error"))
} }
} }
@@ -66,11 +65,9 @@ impl CommandClient {
let mut buf = [0; 10240]; let mut buf = [0; 10240];
let len = self.udp.recv(&mut buf)?; let len = self.udp.recv(&mut buf)?;
match serde_json::from_slice::<Info>(&buf[..len]) { match serde_json::from_slice::<Info>(&buf[..len]) {
Ok(val) => { Ok(val) => Ok(val),
Ok(val)
}
Err(e) => { Err(e) => {
log::error!("{:?},{:?}",&buf[..len],e); log::error!("{:?},{:?}", &buf[..len], e);
Err(io::Error::new(io::ErrorKind::Other, "data error")) Err(io::Error::new(io::ErrorKind::Other, "data error"))
} }
} }
+5 -3
View File
@@ -9,7 +9,8 @@ pub struct Info {
pub relay_server: String, pub relay_server: String,
pub nat_type: String, pub nat_type: String,
pub public_ips: String, pub public_ips: String,
pub local_ip: String, pub local_addr: String,
pub ipv6_addr: String,
} }
#[derive(Serialize, Deserialize, Debug)] #[derive(Serialize, Deserialize, Debug)]
@@ -28,9 +29,10 @@ pub struct DeviceItem {
pub nat_type: String, pub nat_type: String,
pub public_ips: String, pub public_ips: String,
pub local_ip: String, pub local_ip: String,
pub ipv6: String,
pub nat_traversal_type: String, pub nat_traversal_type: String,
pub rt: String, pub rt: String,
pub status: String, pub status: String,
pub client_secret: bool, pub client_secret: bool,
pub current_client_secret:bool, pub current_client_secret: bool,
} }
+34 -17
View File
@@ -1,11 +1,11 @@
use crate::command::entity::{DeviceItem, Info, RouteItem};
use crate::console_out;
use std::io; use std::io;
use vnt::core::Vnt; use vnt::core::Vnt;
use crate::command::entity::{DeviceItem, RouteItem, Info};
use crate::console_out;
pub mod client; pub mod client;
pub mod server;
pub mod entity; pub mod entity;
pub mod server;
pub enum CommandEnum { pub enum CommandEnum {
Route, Route,
@@ -51,7 +51,9 @@ pub fn command_route(vnt: &Vnt) -> Vec<RouteItem> {
let route_table = vnt.route_table(); let route_table = vnt.route_table();
let mut route_list = Vec::with_capacity(route_table.len()); let mut route_list = Vec::with_capacity(route_table.len());
for (destination, route) in route_table { for (destination, route) in route_table {
let next_hop = vnt.route_key(&route.route_key()).map_or(String::new(), |v| v.to_string()); let next_hop = vnt
.route_key(&route.route_key())
.map_or(String::new(), |v| v.to_string());
let metric = route.metric.to_string(); let metric = route.metric.to_string();
let rt = if route.rt < 0 { let rt = if route.rt < 0 {
"".to_string() "".to_string()
@@ -79,15 +81,23 @@ pub fn command_list(vnt: &Vnt) -> Vec<DeviceItem> {
for peer in device_list { for peer in device_list {
let name = peer.name; let name = peer.name;
let virtual_ip = peer.virtual_ip.to_string(); let virtual_ip = peer.virtual_ip.to_string();
let (nat_type, public_ips, local_ip) = if let Some(nat_info) = vnt.peer_nat_info(&peer.virtual_ip) { let (nat_type, public_ips, local_ip, ipv6) =
let nat_type = format!("{:?}", nat_info.nat_type); if let Some(nat_info) = vnt.peer_nat_info(&peer.virtual_ip) {
let public_ips: Vec<String> = nat_info.public_ips.iter().map(|v| v.to_string()).collect(); let nat_type = format!("{:?}", nat_info.nat_type);
let public_ips = public_ips.join(","); let public_ips: Vec<String> =
let local_ip = nat_info.local_ip.to_string(); nat_info.public_ips.iter().map(|v| v.to_string()).collect();
(nat_type, public_ips, local_ip) let public_ips = public_ips.join(",");
} else { let local_ip = nat_info.local_ipv4_addr.ip().to_string();
("".to_string(), "".to_string(), "".to_string()) let ipv6 = nat_info.ipv6_addr.ip().to_string();
}; (nat_type, public_ips, local_ip, ipv6)
} else {
(
"".to_string(),
"".to_string(),
"".to_string(),
"".to_string(),
)
};
let (nat_traversal_type, rt) = if let Some(route) = vnt.route(&peer.virtual_ip) { let (nat_traversal_type, rt) = if let Some(route) = vnt.route(&peer.virtual_ip) {
let nat_traversal_type = if route.metric == 1 { let nat_traversal_type = if route.metric == 1 {
"p2p" "p2p"
@@ -95,7 +105,8 @@ pub fn command_list(vnt: &Vnt) -> Vec<DeviceItem> {
"server-relay" "server-relay"
} else { } else {
"client-relay" "client-relay"
}.to_string(); }
.to_string();
let rt = if route.rt < 0 { let rt = if route.rt < 0 {
"".to_string() "".to_string()
} else { } else {
@@ -113,6 +124,7 @@ pub fn command_list(vnt: &Vnt) -> Vec<DeviceItem> {
nat_type, nat_type,
public_ips, public_ips,
local_ip, local_ip,
ipv6,
nat_traversal_type, nat_traversal_type,
rt, rt,
status, status,
@@ -136,7 +148,12 @@ pub fn command_info(vnt: &Vnt) -> Info {
let nat_type = format!("{:?}", nat_info.nat_type); let nat_type = format!("{:?}", nat_info.nat_type);
let public_ips: Vec<String> = nat_info.public_ips.iter().map(|v| v.to_string()).collect(); let public_ips: Vec<String> = nat_info.public_ips.iter().map(|v| v.to_string()).collect();
let public_ips = public_ips.join(","); let public_ips = public_ips.join(",");
let local_ip = nat_info.local_ip.to_string(); let local_addr = nat_info.local_ipv4_addr.to_string();
let ipv6_addr = if nat_info.ipv6_addr.ip().is_unspecified() {
"None".to_string()
} else {
nat_info.ipv6_addr.ip().to_string()
};
Info { Info {
name, name,
virtual_ip, virtual_ip,
@@ -146,7 +163,7 @@ pub fn command_info(vnt: &Vnt) -> Info {
relay_server, relay_server,
nat_type, nat_type,
public_ips, public_ips,
local_ip, local_addr,
ipv6_addr,
} }
} }
+15 -29
View File
@@ -4,7 +4,6 @@ use tokio::net::UdpSocket;
use vnt::core::Vnt; use vnt::core::Vnt;
pub struct CommandServer {} pub struct CommandServer {}
impl CommandServer { impl CommandServer {
@@ -41,39 +40,26 @@ impl CommandServer {
} }
} }
fn command(cmd: &str, vnt: &Vnt) -> io::Result<String> { fn command(cmd: &str, vnt: &Vnt) -> io::Result<String> {
let out_str = match cmd { let out_str = match cmd {
"route" => { "route" => match serde_json::to_string(&crate::command::command_route(vnt)) {
match serde_json::to_string(&crate::command::command_route(vnt)) { Ok(str) => str,
Ok(str) => { Err(e) => {
str format!("{:?}", e)
}
Err(e) => {
format!("{:?}", e)
}
} }
} },
"list" => { "list" => match serde_json::to_string(&crate::command::command_list(vnt)) {
match serde_json::to_string(&crate::command::command_list(vnt)) { Ok(str) => str,
Ok(str) => { Err(e) => {
str format!("{:?}", e)
}
Err(e) => {
format!("{:?}", e)
}
} }
} },
"info" => { "info" => match serde_json::to_string(&crate::command::command_info(vnt)) {
match serde_json::to_string(&crate::command::command_info(vnt)) { Ok(str) => str,
Ok(str) => { Err(e) => {
str format!("{:?}", e)
}
Err(e) => {
format!("{:?}", e)
}
} }
} },
"stop" => { "stop" => {
vnt.stop()?; vnt.stop()?;
"stopped".to_string() "stopped".to_string()
+101 -71
View File
@@ -1,6 +1,6 @@
use console::{style, Style}; use console::{style, Style};
use crate::command::entity::{DeviceItem, RouteItem, Info}; use crate::command::entity::{DeviceItem, Info, RouteItem};
pub mod table; pub mod table;
@@ -9,11 +9,15 @@ pub fn console_info(status: Info) {
println!("Virtual ip: {}", style(status.virtual_ip).green()); println!("Virtual ip: {}", style(status.virtual_ip).green());
println!("Virtual gateway: {}", style(status.virtual_gateway).green()); println!("Virtual gateway: {}", style(status.virtual_gateway).green());
println!("Virtual netmask: {}", style(status.virtual_netmask).green()); println!("Virtual netmask: {}", style(status.virtual_netmask).green());
println!("Connection status: {}", style(status.connect_status).green()); println!(
"Connection status: {}",
style(status.connect_status).green()
);
println!("NAT type: {}", style(status.nat_type).green()); println!("NAT type: {}", style(status.nat_type).green());
println!("Relay server: {}", style(status.relay_server).green()); println!("Relay server: {}", style(status.relay_server).green());
println!("Public ips: {}", style(status.public_ips).green()); println!("Public ips: {}", style(status.public_ips).green());
println!("Local ip: {}", style(status.local_ip).green()); println!("Local addr: {}", style(status.local_addr).green());
println!("IPv6: {}", style(status.ipv6_addr).green());
} }
pub fn console_route_table(mut list: Vec<RouteItem>) { pub fn console_route_table(mut list: Vec<RouteItem>) {
@@ -24,17 +28,21 @@ pub fn console_route_table(mut list: Vec<RouteItem>) {
list.sort_by(|t1, t2| t1.destination.cmp(&t2.destination)); list.sort_by(|t1, t2| t1.destination.cmp(&t2.destination));
let mut out_list = Vec::with_capacity(list.len()); let mut out_list = Vec::with_capacity(list.len());
out_list.push(vec![("Destination".to_string(), Style::new()), out_list.push(vec![
("Next Hop".to_string(), Style::new()), ("Destination".to_string(), Style::new()),
("Metric".to_string(), Style::new()), ("Next Hop".to_string(), Style::new()),
("Rt".to_string(), Style::new()), ("Metric".to_string(), Style::new()),
("Interface".to_string(), Style::new()), ]); ("Rt".to_string(), Style::new()),
("Interface".to_string(), Style::new()),
]);
for item in list { for item in list {
out_list.push(vec![(item.destination, Style::new().green()), out_list.push(vec![
(item.next_hop, Style::new().green()), (item.destination, Style::new().green()),
(item.metric, Style::new().green()), (item.next_hop, Style::new().green()),
(item.rt, Style::new().green()), (item.metric, Style::new().green()),
(item.interface, Style::new().green())]); (item.rt, Style::new().green()),
(item.interface, Style::new().green()),
]);
} }
table::println_table(out_list) table::println_table(out_list)
@@ -49,41 +57,51 @@ pub fn console_device_list(mut list: Vec<DeviceItem>) {
list.sort_by(|t1, t2| t1.status.cmp(&t2.status)); list.sort_by(|t1, t2| t1.status.cmp(&t2.status));
let mut out_list = Vec::with_capacity(list.len()); let mut out_list = Vec::with_capacity(list.len());
//表头 //表头
out_list.push(vec![("Name".to_string(), Style::new()), out_list.push(vec![
("Virtual Ip".to_string(), Style::new()), ("Name".to_string(), Style::new()),
("Status".to_string(), Style::new()), ("Virtual Ip".to_string(), Style::new()),
("P2P/Relay".to_string(), Style::new()), ("Status".to_string(), Style::new()),
("Rt".to_string(), Style::new())]); ("P2P/Relay".to_string(), Style::new()),
("Rt".to_string(), Style::new()),
]);
for item in list { for item in list {
if &item.status == "Online" { if &item.status == "Online" {
if item.client_secret != item.current_client_secret { if item.client_secret != item.current_client_secret {
//加密状态不一致,无法通信的 //加密状态不一致,无法通信的
out_list.push(vec![(item.name, Style::new().red()), out_list.push(vec![
(item.virtual_ip, Style::new().red()), (item.name, Style::new().red()),
(item.status, Style::new().red()), (item.virtual_ip, Style::new().red()),
("".to_string(), Style::new().red()), (item.status, Style::new().red()),
("".to_string(), Style::new().red())]); ("".to_string(), Style::new().red()),
("".to_string(), Style::new().red()),
]);
} else { } else {
if &item.nat_traversal_type == "p2p" { if &item.nat_traversal_type == "p2p" {
out_list.push(vec![(item.name, Style::new().green()), out_list.push(vec![
(item.virtual_ip, Style::new().green()), (item.name, Style::new().green()),
(item.status, Style::new().green()), (item.virtual_ip, Style::new().green()),
(item.nat_traversal_type, Style::new().green()), (item.status, Style::new().green()),
(item.rt, Style::new().green())]); (item.nat_traversal_type, Style::new().green()),
(item.rt, Style::new().green()),
]);
} else { } else {
out_list.push(vec![(item.name, Style::new().yellow()), out_list.push(vec![
(item.virtual_ip, Style::new().yellow()), (item.name, Style::new().yellow()),
(item.status, Style::new().yellow()), (item.virtual_ip, Style::new().yellow()),
(item.nat_traversal_type, Style::new().yellow()), (item.status, Style::new().yellow()),
(item.rt, Style::new().yellow())]); (item.nat_traversal_type, Style::new().yellow()),
(item.rt, Style::new().yellow()),
]);
} }
} }
} else { } else {
out_list.push(vec![(item.name, Style::new().color256(102)), out_list.push(vec![
(item.virtual_ip, Style::new().color256(102)), (item.name, Style::new().color256(102)),
(item.status, Style::new().color256(102)), (item.virtual_ip, Style::new().color256(102)),
("".to_string(), Style::new().color256(102)), (item.status, Style::new().color256(102)),
("".to_string(), Style::new().color256(102))]); ("".to_string(), Style::new().color256(102)),
("".to_string(), Style::new().color256(102)),
]);
} }
} }
table::println_table(out_list) table::println_table(out_list)
@@ -98,45 +116,57 @@ pub fn console_device_list_all(mut list: Vec<DeviceItem>) {
list.sort_by(|t1, t2| t1.status.cmp(&t2.status)); list.sort_by(|t1, t2| t1.status.cmp(&t2.status));
let mut out_list = Vec::with_capacity(list.len()); let mut out_list = Vec::with_capacity(list.len());
//表头 //表头
out_list.push(vec![("Name".to_string(), Style::new()), out_list.push(vec![
("Virtual Ip".to_string(), Style::new()), ("Name".to_string(), Style::new()),
("Status".to_string(), Style::new()), ("Virtual Ip".to_string(), Style::new()),
("P2P/Relay".to_string(), Style::new()), ("Status".to_string(), Style::new()),
("Rt".to_string(), Style::new()), ("P2P/Relay".to_string(), Style::new()),
("NAT Type".to_string(), Style::new()), ("Rt".to_string(), Style::new()),
("Public Ips".to_string(), Style::new()), ("NAT Type".to_string(), Style::new()),
("Local Ip".to_string(), Style::new())]); ("Public Ips".to_string(), Style::new()),
("Local Ip".to_string(), Style::new()),
("IPv6".to_string(), Style::new()),
]);
for item in list { for item in list {
if &item.status == "Online" { if &item.status == "Online" {
if &item.nat_traversal_type == "p2p" { if &item.nat_traversal_type == "p2p" {
out_list.push(vec![(item.name, Style::new().green()), out_list.push(vec![
(item.virtual_ip, Style::new().green()), (item.name, Style::new().green()),
(item.status, Style::new().green()), (item.virtual_ip, Style::new().green()),
(item.nat_traversal_type, Style::new().green()), (item.status, Style::new().green()),
(item.rt, Style::new().green()), (item.nat_traversal_type, Style::new().green()),
(item.nat_type, Style::new().green()), (item.rt, Style::new().green()),
(item.public_ips, Style::new().green()), (item.nat_type, Style::new().green()),
(item.local_ip, Style::new().green())]); (item.public_ips, Style::new().green()),
(item.local_ip, Style::new().green()),
(item.ipv6, Style::new().green()),
]);
} else { } else {
out_list.push(vec![(item.name, Style::new().yellow()), out_list.push(vec![
(item.virtual_ip, Style::new().yellow()), (item.name, Style::new().yellow()),
(item.status, Style::new().yellow()), (item.virtual_ip, Style::new().yellow()),
(item.nat_traversal_type, Style::new().yellow()), (item.status, Style::new().yellow()),
(item.rt, Style::new().yellow()), (item.nat_traversal_type, Style::new().yellow()),
(item.nat_type, Style::new().yellow()), (item.rt, Style::new().yellow()),
(item.public_ips, Style::new().yellow()), (item.nat_type, Style::new().yellow()),
(item.local_ip, Style::new().yellow()), ]); (item.public_ips, Style::new().yellow()),
(item.local_ip, Style::new().yellow()),
(item.ipv6, Style::new().yellow()),
]);
} }
} else { } else {
out_list.push(vec![(item.name, Style::new().color256(102)), out_list.push(vec![
(item.virtual_ip, Style::new().color256(102)), (item.name, Style::new().color256(102)),
(item.status, Style::new().color256(102)), (item.virtual_ip, Style::new().color256(102)),
("".to_string(), Style::new().color256(102)), (item.status, Style::new().color256(102)),
("".to_string(), Style::new().color256(102)), ("".to_string(), Style::new().color256(102)),
("".to_string(), Style::new().color256(102)), ("".to_string(), Style::new().color256(102)),
("".to_string(), Style::new().color256(102)), ("".to_string(), Style::new().color256(102)),
("".to_string(), Style::new().color256(102)), ]); ("".to_string(), Style::new().color256(102)),
("".to_string(), Style::new().color256(102)),
("".to_string(), Style::new().color256(102)),
]);
} }
} }
table::println_table(out_list) table::println_table(out_list)
} }
+1 -1
View File
@@ -20,4 +20,4 @@ pub fn println_table(table: Vec<Vec<(String, Style)>>) {
} }
println!() println!()
} }
} }
+154 -80
View File
@@ -11,6 +11,7 @@ use tokio::signal;
use tokio::signal::unix::{signal, SignalKind}; use tokio::signal::unix::{signal, SignalKind};
use common::args_parse::{ips_parse, out_ips_parse}; use common::args_parse::{ips_parse, out_ips_parse};
use vnt::channel::punch::PunchModel;
use vnt::cipher::CipherModel; use vnt::cipher::CipherModel;
use vnt::core::{Config, Vnt, VntUtil}; use vnt::core::{Config, Vnt, VntUtil};
use vnt::handle::handshake_handler::HandshakeEnum; use vnt::handle::handshake_handler::HandshakeEnum;
@@ -21,7 +22,9 @@ mod console_out;
mod root_check; mod root_check;
pub fn app_home() -> io::Result<PathBuf> { pub fn app_home() -> io::Result<PathBuf> {
let path = dirs::home_dir().ok_or(io::Error::new(io::ErrorKind::Other, "not home"))?.join(".vnt-cli"); let path = dirs::home_dir()
.ok_or(io::Error::new(io::ErrorKind::Other, "not home"))?
.join(".vnt-cli");
if !path.exists() { if !path.exists() {
std::fs::create_dir_all(&path)?; std::fs::create_dir_all(&path)?;
} }
@@ -52,6 +55,13 @@ fn main() {
opts.optopt("", "par", "任务并行度(必须为正整数)", "<parallel>"); opts.optopt("", "par", "任务并行度(必须为正整数)", "<parallel>");
opts.optopt("", "thread", "线程数(必须为正整数)", "<thread>"); opts.optopt("", "thread", "线程数(必须为正整数)", "<thread>");
opts.optopt("", "model", "加密模式", "<model>"); opts.optopt("", "model", "加密模式", "<model>");
opts.optflag("", "finger", "指纹校验");
opts.optopt(
"",
"punch",
"取值ipv4/ipv6,表示仅使用ipv4或ipv6打洞",
"<punch>",
);
//"后台运行时,查看其他设备列表" //"后台运行时,查看其他设备列表"
opts.optflag("", "list", "后台运行时,查看其他设备列表"); opts.optflag("", "list", "后台运行时,查看其他设备列表");
opts.optflag("", "all", "后台运行时,查看其他设备完整信息"); opts.optflag("", "all", "后台运行时,查看其他设备完整信息");
@@ -60,7 +70,7 @@ fn main() {
opts.optflag("", "stop", "停止后台运行"); opts.optflag("", "stop", "停止后台运行");
opts.optflag("h", "help", "帮助"); opts.optflag("h", "help", "帮助");
let matches = match opts.parse(&args[1..]) { let matches = match opts.parse(&args[1..]) {
Ok(m) => { m } Ok(m) => m,
Err(f) => { Err(f) => {
print_usage(&program, opts); print_usage(&program, opts);
println!("{}", f.to_string()); println!("{}", f.to_string());
@@ -122,8 +132,12 @@ fn main() {
println!("parameter -d not found ."); println!("parameter -d not found .");
return; return;
} }
let name = matches.opt_get_default("n", os_info::get().to_string()).unwrap(); let name = matches
let server_address_str = matches.opt_get_default("s", "nat1.wherewego.top:29872".to_string()).unwrap(); .opt_get_default("n", os_info::get().to_string())
.unwrap();
let server_address_str = matches
.opt_get_default("s", "nat1.wherewego.top:29872".to_string())
.unwrap();
let server_address = match server_address_str.to_socket_addrs() { let server_address = match server_address_str.to_socket_addrs() {
Ok(mut addr) => { Ok(mut addr) => {
if let Some(addr) = addr.next() { if let Some(addr) = addr.next() {
@@ -147,7 +161,7 @@ fn main() {
let in_ip = matches.opt_strs("i"); let in_ip = matches.opt_strs("i");
let in_ip = match ips_parse(&in_ip) { let in_ip = match ips_parse(&in_ip) {
Ok(in_ip) => { in_ip } Ok(in_ip) => in_ip,
Err(e) => { Err(e) => {
print_usage(&program, opts); print_usage(&program, opts);
println!(); println!();
@@ -158,7 +172,7 @@ fn main() {
}; };
let out_ip = matches.opt_strs("o"); let out_ip = matches.opt_strs("o");
let out_ip = match out_ips_parse(&out_ip) { let out_ip = match out_ips_parse(&out_ip) {
Ok(out_ip) => { out_ip } Ok(out_ip) => out_ip,
Err(e) => { Err(e) => {
print_usage(&program, opts); print_usage(&program, opts);
println!(); println!();
@@ -174,9 +188,7 @@ fn main() {
let mtu: Option<String> = matches.opt_get("u").unwrap(); let mtu: Option<String> = matches.opt_get("u").unwrap();
let mtu = if let Some(mtu) = mtu { let mtu = if let Some(mtu) = mtu {
match u16::from_str(&mtu) { match u16::from_str(&mtu) {
Ok(mtu) => { Ok(mtu) => Some(mtu),
Some(mtu)
}
Err(e) => { Err(e) => {
print_usage(&program, opts); print_usage(&program, opts);
println!(); println!();
@@ -202,20 +214,51 @@ fn main() {
println!("--par invalid"); println!("--par invalid");
return; return;
} }
let thread_num = matches.opt_get::<usize>("thread").unwrap().unwrap_or(std::thread::available_parallelism().unwrap().get() * 2); let thread_num = matches
let cipher_model = matches.opt_get::<CipherModel>("model").unwrap().unwrap_or(CipherModel::AesGcm); .opt_get::<usize>("thread")
.unwrap()
.unwrap_or(std::thread::available_parallelism().unwrap().get() * 2);
let cipher_model = matches
.opt_get::<CipherModel>("model")
.unwrap()
.unwrap_or(CipherModel::AesGcm);
if thread_num == 0 { if thread_num == 0 {
println!("--thread invalid"); println!("--thread invalid");
return; return;
} }
println!("version 1.2.0"); let finger = matches.opt_present("finger");
let config = Config::new(tap, let punch_model = matches
token, device_id, name, .opt_get::<PunchModel>("punch")
server_address, server_address_str, .unwrap()
stun_server, in_ip, .unwrap_or(PunchModel::All);
out_ip, password, simulate_multicast, mtu, println!("version {}", vnt::VNT_VERSION);
tcp_channel, virtual_ip, relay, server_encrypt, parallel, cipher_model); let config = Config::new(
let runtime = tokio::runtime::Builder::new_multi_thread().enable_all().worker_threads(thread_num).build().unwrap(); tap,
token,
device_id,
name,
server_address,
server_address_str,
stun_server,
in_ip,
out_ip,
password,
simulate_multicast,
mtu,
tcp_channel,
virtual_ip,
relay,
server_encrypt,
parallel,
cipher_model,
finger,
punch_model,
);
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.worker_threads(thread_num)
.build()
.unwrap();
runtime.block_on(main0(config, !unused_cmd)); runtime.block_on(main0(config, !unused_cmd));
std::process::exit(0); std::process::exit(0);
} }
@@ -262,55 +305,51 @@ async fn main0(config: Config, show_cmd: bool) {
Ok(response) => { Ok(response) => {
break response; break response;
} }
Err(e) => { Err(e) => match e {
match e { ReqEnum::TokenError => {
ReqEnum::TokenError => { println!("token error");
println!("token error"); return;
return;
}
ReqEnum::AddressExhausted => {
println!("address exhausted");
return;
}
ReqEnum::Timeout => {
println!("timeout...");
}
ReqEnum::ServerError(str) => {
println!("error:{}", str);
}
ReqEnum::Other(str) => {
println!("error:{}", str);
}
ReqEnum::IpAlreadyExists => {
println!("ip already exists");
return;
}
ReqEnum::InvalidIp => {
println!("invalid ip");
return;
}
} }
} ReqEnum::AddressExhausted => {
println!("address exhausted");
return;
}
ReqEnum::Timeout => {
println!("timeout...");
}
ReqEnum::ServerError(str) => {
println!("error:{}", str);
}
ReqEnum::Other(str) => {
println!("error:{}", str);
}
ReqEnum::IpAlreadyExists => {
println!("ip already exists");
return;
}
ReqEnum::InvalidIp => {
println!("invalid ip");
return;
}
},
} }
} }
Err(e) => { Err(e) => match e {
match e { HandshakeEnum::NotSecret => {
HandshakeEnum::NotSecret => { println!("The server does not support encryption");
println!("The server does not support encryption"); return;
return;
}
HandshakeEnum::KeyError => {}
HandshakeEnum::Timeout => {
println!("handshake timeout")
}
HandshakeEnum::ServerError(str) => {
println!("error:{}", str);
}
HandshakeEnum::Other(str) => {
println!("error:{}", str);
}
} }
} HandshakeEnum::KeyError => {}
HandshakeEnum::Timeout => {
println!("handshake timeout")
}
HandshakeEnum::ServerError(str) => {
println!("error:{}", str);
}
HandshakeEnum::Other(str) => {
println!("error:{}", str);
}
},
} }
}; };
println!(" ====== Connect Successfully ====== "); println!(" ====== Connect Successfully ====== ");
@@ -321,9 +360,7 @@ async fn main0(config: Config, show_cmd: bool) {
println!("name:{}", driver_info.name); println!("name:{}", driver_info.name);
println!("version:{}", driver_info.version); println!("version:{}", driver_info.version);
let mut vnt = match vnt_util.build().await { let mut vnt = match vnt_util.build().await {
Ok(vnt) => { Ok(vnt) => vnt,
vnt
}
Err(e) => { Err(e) => {
println!("error:{}", e); println!("error:{}", e);
return; return;
@@ -336,19 +373,19 @@ async fn main0(config: Config, show_cmd: bool) {
println!("command error :{}", e); println!("command error :{}", e);
} }
}); });
#[cfg(unix)]
let mut sigterm = signal(SignalKind::terminate()).expect("Error setting SIGTERM handler");
if show_cmd { if show_cmd {
let stdin = tokio::io::stdin(); let stdin = tokio::io::stdin();
let mut cmd = String::new(); let mut cmd = String::new();
let mut reader = BufReader::new(stdin); let mut reader = BufReader::new(stdin);
#[cfg(unix)]
let mut sigterm = signal(SignalKind::terminate()).expect("Error setting SIGTERM handler");
loop { loop {
cmd.clear(); cmd.clear();
println!("input:list,info,route,all,stop"); println!("input:list,info,route,all,stop");
#[cfg(unix)] #[cfg(unix)]
tokio::select! { tokio::select! {
_ = vnt.wait_stop()=>{ _ = vnt.wait_stop()=>{
break; return;
} }
_ = signal::ctrl_c()=>{ _ = signal::ctrl_c()=>{
let _ = vnt.stop(); let _ = vnt.stop();
@@ -377,7 +414,7 @@ async fn main0(config: Config, show_cmd: bool) {
#[cfg(windows)] #[cfg(windows)]
tokio::select! { tokio::select! {
_ = vnt.wait_stop()=>{ _ = vnt.wait_stop()=>{
break; return;
} }
_ = signal::ctrl_c()=>{ _ = signal::ctrl_c()=>{
let _ = vnt.stop(); let _ = vnt.stop();
@@ -400,6 +437,23 @@ async fn main0(config: Config, show_cmd: bool) {
} }
} }
} }
#[cfg(unix)]
tokio::select! {
_ = vnt.wait_stop()=>{
return;
}
_ = signal::ctrl_c()=>{
let _ = vnt.stop();
vnt.wait_stop_ms(std::time::Duration::from_secs(3)).await;
return;
}
_ = sigterm.recv()=>{
let _ = vnt.stop();
vnt.wait_stop_ms(std::time::Duration::from_secs(3)).await;
return;
}
}
#[cfg(windows)]
vnt.wait_stop().await; vnt.wait_stop().await;
} }
@@ -436,9 +490,12 @@ fn command(cmd: &str, vnt: &Vnt) -> bool {
fn print_usage(program: &str, _opts: Options) { fn print_usage(program: &str, _opts: Options) {
println!("Usage: {} [options]", program); println!("Usage: {} [options]", program);
println!("version:1.2.0"); println!("version:{}", vnt::VNT_VERSION);
println!("Options:"); println!("Options:");
println!(" -k <token> {}", green("必选,使用相同的token,就能组建一个局域网络".to_string())); println!(
" -k <token> {}",
green("必选,使用相同的token,就能组建一个局域网络".to_string())
);
println!(" -n <name> 给设备一个名字,便于区分不同设备,默认使用系统版本"); println!(" -n <name> 给设备一个名字,便于区分不同设备,默认使用系统版本");
println!(" -d <id> 设备唯一标识符,不使用--ip参数时,服务端凭此参数分配虚拟ip"); println!(" -d <id> 设备唯一标识符,不使用--ip参数时,服务端凭此参数分配虚拟ip");
println!(" -c 关闭交互式命令,使用此参数禁用控制台输入"); println!(" -c 关闭交互式命令,使用此参数禁用控制台输入");
@@ -457,13 +514,31 @@ fn print_usage(program: &str, _opts: Options) {
println!(" --relay 仅使用服务器转发,不使用p2p,默认情况允许使用p2p"); println!(" --relay 仅使用服务器转发,不使用p2p,默认情况允许使用p2p");
println!(" --par <parallel> 任务并行度(必须为正整数),默认值为1"); println!(" --par <parallel> 任务并行度(必须为正整数),默认值为1");
println!(" --thread <thread> 线程数(必须为正整数),默认为核心数乘2"); println!(" --thread <thread> 线程数(必须为正整数),默认为核心数乘2");
println!(" --model <model> 加密模式,可选值 aes_gcm/aes_cbc,默认使用aes_gcm,通常情况使用aes_cbc性能更好"); println!(" --model <model> 加密模式(默认aes_gcm),可选值aes_gcm/aes_cbc/aes_ecb,通常性能aes_ecb>aes_cbc>aes_gcm,安全性则相反");
println!(" --finger 增加数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启");
println!(" --punch <punch> 取值ipv4/ipv6ipv4表示仅使用ipv4打洞");
println!(); println!();
println!(" --list {}", yellow("后台运行时,查看其他设备列表".to_string())); println!(
println!(" --all {}", yellow("后台运行时,查看其他设备完整信息".to_string())); " --list {}",
println!(" --info {}", yellow("后台运行时,查看当前设备信息".to_string())); yellow("后台运行时,查看其他设备列表".to_string())
println!(" --route {}", yellow("后台运行时,查看数据转发路径".to_string())); );
println!(" --stop {}", yellow("停止后台运行".to_string())); println!(
" --all {}",
yellow("后台运行时,查看其他设备完整信息".to_string())
);
println!(
" --info {}",
yellow("后台运行时,查看当前设备信息".to_string())
);
println!(
" --route {}",
yellow("后台运行时,查看数据转发路径".to_string())
);
println!(
" --stop {}",
yellow("停止后台运行".to_string())
);
println!(" -h, --help 帮助"); println!(" -h, --help 帮助");
} }
@@ -474,4 +549,3 @@ fn green(str: String) -> impl std::fmt::Display {
fn yellow(str: String) -> impl std::fmt::Display { fn yellow(str: String) -> impl std::fmt::Display {
style(str).yellow() style(str).yellow()
} }
+1 -1
View File
@@ -8,4 +8,4 @@ pub use windows::is_app_elevated;
mod unix; mod unix;
#[cfg(any(target_os = "linux", target_os = "macos"))] #[cfg(any(target_os = "linux", target_os = "macos"))]
pub use unix::is_app_elevated; pub use unix::is_app_elevated;
+1 -1
View File
@@ -1,3 +1,3 @@
pub fn is_app_elevated() -> bool { pub fn is_app_elevated() -> bool {
sudo::RunningAs::Root == sudo::check() sudo::RunningAs::Root == sudo::check()
} }
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "vnt-jni" name = "vnt-jni"
version = "1.2.0" version = "1.2.2"
edition = "2021" edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
+1 -1
View File
@@ -1,2 +1,2 @@
pub mod vnt;
pub mod vnt_util; pub mod vnt_util;
pub mod vnt;
+19 -15
View File
@@ -1,8 +1,8 @@
use std::ptr;
use jni::errors::Error; use jni::errors::Error;
use jni::JNIEnv;
use jni::objects::{JClass, JObject, JValue}; use jni::objects::{JClass, JObject, JValue};
use jni::sys::{jboolean, jbyte, jint, jlong, jobject, jobjectArray, jsize}; use jni::sys::{jboolean, jbyte, jint, jlong, jobject, jobjectArray, jsize};
use jni::JNIEnv;
use std::ptr;
use vnt::channel::Route; use vnt::channel::Route;
use vnt::core::sync::VntSync; use vnt::core::sync::VntSync;
use vnt::handle::PeerDeviceInfo; use vnt::handle::PeerDeviceInfo;
@@ -67,7 +67,7 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_list0(
"top/wherewego/vnt/jni/PeerDeviceInfo", "top/wherewego/vnt/jni/PeerDeviceInfo",
JObject::null(), JObject::null(),
) { ) {
Ok(arr) => { arr } Ok(arr) => arr,
Err(e) => { Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("error:{:?}", e)) env.throw_new("java/lang/RuntimeException", format!("error:{:?}", e))
.expect("throw"); .expect("throw");
@@ -77,12 +77,8 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_list0(
for (index, peer) in list.into_iter().enumerate() { for (index, peer) in list.into_iter().enumerate() {
let route = if let Some(route) = vnt.route(&peer.virtual_ip) { let route = if let Some(route) = vnt.route(&peer.virtual_ip) {
match route_parse(&mut env, route) { match route_parse(&mut env, route) {
Ok(route) => { Ok(route) => JObject::from_raw(route),
JObject::from_raw(route) Err(_) => JObject::null(),
}
Err(_) => {
JObject::null()
}
} }
} else { } else {
JObject::null() JObject::null()
@@ -115,24 +111,32 @@ fn route_parse(env: &mut JNIEnv, route: Route) -> Result<jobject, Error> {
let rs = env.new_object( let rs = env.new_object(
"top/wherewego/vnt/jni/Route", "top/wherewego/vnt/jni/Route",
"(Ljava/lang/String;BI)V", "(Ljava/lang/String;BI)V",
&[JValue::Object(&env.new_string(address)?.into()), &[
JValue::Object(&env.new_string(address)?.into()),
JValue::Byte(metric as jbyte), JValue::Byte(metric as jbyte),
JValue::Int(rt as jint)], JValue::Int(rt as jint),
],
)?; )?;
Ok(rs.as_raw()) Ok(rs.as_raw())
} }
fn peer_device_info_parse(env: &mut JNIEnv, peer: PeerDeviceInfo, route: JObject) -> Result<jobject, Error> { fn peer_device_info_parse(
env: &mut JNIEnv,
peer: PeerDeviceInfo,
route: JObject,
) -> Result<jobject, Error> {
let virtual_ip = u32::from(peer.virtual_ip); let virtual_ip = u32::from(peer.virtual_ip);
let name = peer.name.to_string(); let name = peer.name.to_string();
let status = format!("{:?}", peer.status); let status = format!("{:?}", peer.status);
let rs = env.new_object( let rs = env.new_object(
"top/wherewego/vnt/jni/PeerDeviceInfo", "top/wherewego/vnt/jni/PeerDeviceInfo",
"(ILjava/lang/String;Ljava/lang/String;Ltop/wherewego/vnt/jni/Route;)V", "(ILjava/lang/String;Ljava/lang/String;Ltop/wherewego/vnt/jni/Route;)V",
&[JValue::Int(virtual_ip as jint), &[
JValue::Int(virtual_ip as jint),
JValue::Object(&env.new_string(name)?.into()), JValue::Object(&env.new_string(name)?.into()),
JValue::Object(&env.new_string(status)?.into()), JValue::Object(&env.new_string(status)?.into()),
JValue::Object(&route)], JValue::Object(&route),
],
)?; )?;
Ok(rs.as_raw()) Ok(rs.as_raw())
} }
+151 -81
View File
@@ -1,5 +1,6 @@
use std::net::ToSocketAddrs; use std::net::ToSocketAddrs;
use std::ptr; use std::ptr;
use std::str::FromStr;
use jni::errors::Error; use jni::errors::Error;
use jni::objects::{JClass, JObject, JString, JValue}; use jni::objects::{JClass, JObject, JString, JValue};
@@ -7,17 +8,22 @@ use jni::objects::{JClass, JObject, JString, JValue};
use jni::sys::jboolean; use jni::sys::jboolean;
use jni::sys::{jint, jlong, jobject}; use jni::sys::{jint, jlong, jobject};
use jni::JNIEnv; use jni::JNIEnv;
use vnt::channel::punch::PunchModel;
use vnt::cipher::CipherModel; use vnt::cipher::CipherModel;
use vnt::core::Config;
use vnt::core::sync::VntUtilSync; use vnt::core::sync::VntUtilSync;
use vnt::core::Config;
use vnt::handle::registration_handler::{RegResponse, ReqEnum}; use vnt::handle::registration_handler::{RegResponse, ReqEnum};
#[cfg(not(target_os = "android"))] #[cfg(not(target_os = "android"))]
use vnt::tun_tap_device::DriverInfo; use vnt::tun_tap_device::DriverInfo;
fn to_string_not_null(env: &mut JNIEnv, config: &JObject, name: &'static str) -> Result<String, Error> { fn to_string_not_null(
env: &mut JNIEnv,
config: &JObject,
name: &'static str,
) -> Result<String, Error> {
let value = env.get_field(config, name, "Ljava/lang/String;")?.l()?; let value = env.get_field(config, name, "Ljava/lang/String;")?.l()?;
if value.is_null() { if value.is_null() {
env.throw_new("Ljava/lang/NullPointerException", name) env.throw_new("java/lang/NullPointerException", name)
.expect("throw"); .expect("throw");
return Err(Error::NullPtr(name)); return Err(Error::NullPtr(name));
} }
@@ -26,7 +32,7 @@ fn to_string_not_null(env: &mut JNIEnv, config: &JObject, name: &'static str) ->
match value.to_str() { match value.to_str() {
Ok(value) => Ok(value.to_string()), Ok(value) => Ok(value.to_string()),
Err(_) => { Err(_) => {
env.throw_new("Ljava/lang/RuntimeException", "not utf-8") env.throw_new("java/lang/RuntimeException", "not utf-8")
.expect("throw"); .expect("throw");
return Err(Error::JavaException); return Err(Error::JavaException);
} }
@@ -43,7 +49,7 @@ fn to_string(env: &mut JNIEnv, config: &JObject, name: &str) -> Result<Option<St
match value.to_str() { match value.to_str() {
Ok(value) => Ok(Some(value.to_string())), Ok(value) => Ok(Some(value.to_string())),
Err(_) => { Err(_) => {
env.throw_new("Ljava/lang/RuntimeException", "not utf-8") env.throw_new("java/lang/RuntimeException", "not utf-8")
.expect("throw"); .expect("throw");
return Err(Error::JavaException); return Err(Error::JavaException);
} }
@@ -56,39 +62,72 @@ fn new_sync(env: &mut JNIEnv, config: JObject) -> Result<VntUtilSync, Error> {
let device_id = to_string_not_null(env, &config, "deviceId")?; let device_id = to_string_not_null(env, &config, "deviceId")?;
let password = to_string(env, &config, "password")?; let password = to_string(env, &config, "password")?;
let server_address_str = to_string_not_null(env, &config, "server")?; let server_address_str = to_string_not_null(env, &config, "server")?;
// let nat_test_server = to_string_not_null(env, &config, "natTestServer")?; let stun_server_str = to_string_not_null(env, &config, "stunServer")?;
let cipher_model = to_string_not_null(env, &config, "cipherModel")?;
let tcp = env.get_field(&config, "tcp", "Z")?.z()?;
let finger = env.get_field(&config, "finger", "Z")?.z()?;
let server_address = match server_address_str.to_socket_addrs() { let server_address = match server_address_str.to_socket_addrs() {
Ok(mut rs) => { Ok(mut rs) => {
if let Some(addr) = rs.next() { if let Some(addr) = rs.next() {
addr addr
} else { } else {
env.throw_new("Ljava/lang/RuntimeException", "server address err") env.throw_new("java/lang/RuntimeException", "server address err")
.expect("throw"); .expect("throw");
return Err(Error::JavaException); return Err(Error::JavaException);
} }
} }
Err(e) => { Err(e) => {
env.throw_new("Ljava/lang/RuntimeException", format!("server address {}", e)) env.throw_new(
"java/lang/RuntimeException",
format!("server address {}", e),
)
.expect("throw");
return Err(Error::JavaException);
}
};
let cipher_model = match CipherModel::from_str(&cipher_model) {
Ok(cipher_model) => cipher_model,
Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("cipher_model {}", e))
.expect("throw"); .expect("throw");
return Err(Error::JavaException); return Err(Error::JavaException);
} }
}; };
let mut stun_server = Vec::new(); let mut stun_server = Vec::new();
stun_server.push("stun1.l.google.com:19302".to_string()); for addr in stun_server_str.split(",") {
stun_server.push("stun2.l.google.com:19302".to_string()); stun_server.push(addr.trim().to_string());
stun_server.push("stun.qq.com:3478".to_string()); }
let config = Config::new(false, let config = Config::new(
token, device_id, name, false,
server_address, server_address_str, token,
stun_server, vec![], device_id,
vec![], password, false, None, false, None, false,false,1,CipherModel::AesGcm); name,
server_address,
server_address_str,
stun_server,
vec![],
vec![],
password,
false,
None,
tcp,
None,
false,
false,
1,
cipher_model,
finger,
PunchModel::All,
);
match VntUtilSync::new(config) { match VntUtilSync::new(config) {
Ok(vnt_util) => { Ok(vnt_util) => Ok(vnt_util),
Ok(vnt_util)
}
Err(e) => { Err(e) => {
env.throw_new("Ljava/lang/RuntimeException", format!("vnt start error {}", e)) env.throw_new(
.expect("throw"); "java/lang/RuntimeException",
format!("vnt start error {}", e),
)
.expect("throw");
return Err(Error::JavaException); return Err(Error::JavaException);
} }
} }
@@ -109,18 +148,22 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_new0(
} }
return 0; return 0;
} }
#[no_mangle] #[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_connect0( pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_connect0(
mut env: JNIEnv, mut env: JNIEnv,
_class: JClass, _class: JClass,
raw_vnt_util: jlong, raw_vnt_util: jlong,
) { ) {
let raw_vnt_util = raw_vnt_util as *mut VntUtilSync; let raw_vnt_util = raw_vnt_util as *mut VntUtilSync;
match (&mut *raw_vnt_util).connect() { match (&mut *raw_vnt_util).connect() {
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("vnt connect error {}", e)) env.throw_new(
.expect("throw"); "java/lang/RuntimeException",
format!("vnt connect error {}", e),
)
.expect("throw");
} }
} }
} }
@@ -133,48 +176,66 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_register0(
) -> jobject { ) -> jobject {
let raw_vnt_util = raw_vnt_util as *mut VntUtilSync; let raw_vnt_util = raw_vnt_util as *mut VntUtilSync;
match (&mut *raw_vnt_util).register() { match (&mut *raw_vnt_util).register() {
Ok(response) => { Ok(response) => match reg_response(&mut env, response) {
match reg_response(&mut env, response) { Ok(res) => {
Ok(res) => { return res;
return res;
}
Err(e) => {
env.throw(format!("vnt register error {}", e)).expect("throw");
}
} }
} Err(e) => {
Err(e) => { env.throw(format!("vnt register error {}", e))
match e { .expect("throw");
ReqEnum::TokenError => {
env.throw_new("top/wherewego/vnt/jni/exception/TokenErrorException", "TokenError")
.expect("throw");
}
ReqEnum::AddressExhausted => {
env.throw_new("top/wherewego/vnt/jni/exception/AddressExhaustedException", "AddressExhausted")
.expect("throw");
}
ReqEnum::Timeout => {
env.throw_new("top/wherewego/vnt/jni/exception/TimeoutException", "Timeout")
.expect("throw");
}
ReqEnum::ServerError(str) => {
env.throw_new("java/lang/RuntimeException", format!("vnt register error {}", str))
.expect("throw");
}
ReqEnum::Other(str) => {
env.throw_new("java/lang/RuntimeException", format!("vnt register error {}", str))
.expect("throw");
}
ReqEnum::IpAlreadyExists => {
env.throw_new("top/wherewego/vnt/jni/exception/IpAlreadyExistsException", "IpAlreadyExists")
.expect("throw");
}
ReqEnum::InvalidIp => {
env.throw_new("top/wherewego/vnt/jni/exception/InvalidIpException", "InvalidIp")
.expect("throw");
}
} }
} },
Err(e) => match e {
ReqEnum::TokenError => {
env.throw_new(
"top/wherewego/vnt/jni/exception/TokenErrorException",
"TokenError",
)
.expect("throw");
}
ReqEnum::AddressExhausted => {
env.throw_new(
"top/wherewego/vnt/jni/exception/AddressExhaustedException",
"AddressExhausted",
)
.expect("throw");
}
ReqEnum::Timeout => {
env.throw_new(
"top/wherewego/vnt/jni/exception/TimeoutException",
"Timeout",
)
.expect("throw");
}
ReqEnum::ServerError(str) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt register error {}", str),
)
.expect("throw");
}
ReqEnum::Other(str) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt register error {}", str),
)
.expect("throw");
}
ReqEnum::IpAlreadyExists => {
env.throw_new(
"top/wherewego/vnt/jni/exception/IpAlreadyExistsException",
"IpAlreadyExists",
)
.expect("throw");
}
ReqEnum::InvalidIp => {
env.throw_new(
"top/wherewego/vnt/jni/exception/InvalidIpException",
"InvalidIp",
)
.expect("throw");
}
},
} }
return ptr::null_mut(); return ptr::null_mut();
} }
@@ -202,19 +263,21 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_createIface0(
let raw_vnt_util = raw_vnt_util as *mut VntUtilSync; let raw_vnt_util = raw_vnt_util as *mut VntUtilSync;
let rs = (&mut *raw_vnt_util).create_iface(); let rs = (&mut *raw_vnt_util).create_iface();
match rs { match rs {
Ok(driver_info) => { Ok(driver_info) => match driver_info_e(&mut env, driver_info) {
match driver_info_e(&mut env, driver_info) { Ok(res) => {
Ok(res) => { return res;
return res;
}
Err(e) => {
env.throw(format!("vnt create iface error {}", e)).expect("throw");
}
} }
} Err(e) => {
env.throw(format!("vnt create iface error {}", e))
.expect("throw");
}
},
Err(e) => { Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("vnt create iface error {}", e)) env.throw_new(
.expect("throw"); "java/lang/RuntimeException",
format!("vnt create iface error {}", e),
)
.expect("throw");
} }
} }
return ptr::null_mut(); return ptr::null_mut();
@@ -232,8 +295,11 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_build0(
return Box::into_raw(Box::new(rs)) as jlong; return Box::into_raw(Box::new(rs)) as jlong;
} }
Err(e) => { Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("vnt start error:{:?}", e)) env.throw_new(
.expect("throw"); "java/lang/RuntimeException",
format!("vnt start error:{:?}", e),
)
.expect("throw");
} }
} }
return 0; return 0;
@@ -246,9 +312,11 @@ fn reg_response(env: &mut JNIEnv, response: RegResponse) -> Result<jobject, Erro
let response = env.new_object( let response = env.new_object(
"top/wherewego/vnt/jni/RegResponse", "top/wherewego/vnt/jni/RegResponse",
"(III)V", "(III)V",
&[JValue::Int(virtual_ip as jint), &[
JValue::Int(virtual_ip as jint),
JValue::Int(virtual_gateway as jint), JValue::Int(virtual_gateway as jint),
JValue::Int(virtual_netmask as jint)], JValue::Int(virtual_netmask as jint),
],
)?; )?;
Ok(response.into_raw()) Ok(response.into_raw())
} }
@@ -262,10 +330,12 @@ fn driver_info_e(env: &mut JNIEnv, driver_info: DriverInfo) -> Result<jobject, E
let response = env.new_object( let response = env.new_object(
"top/wherewego/vnt/jni/DriverInfo", "top/wherewego/vnt/jni/DriverInfo",
"(ZLjava/lang/String;Ljava/lang/String;Ljava/lang/String;)V", "(ZLjava/lang/String;Ljava/lang/String;Ljava/lang/String;)V",
&[JValue::Bool(is_tun as jboolean), &[
JValue::Bool(is_tun as jboolean),
JValue::Object(&env.new_string(name)?.into()), JValue::Object(&env.new_string(name)?.into()),
JValue::Object(&env.new_string(version)?.into()), JValue::Object(&env.new_string(version)?.into()),
JValue::Object(&env.new_string(mac)?.into()), ], JValue::Object(&env.new_string(mac)?.into()),
],
)?; )?;
Ok(response.into_raw()) Ok(response.into_raw())
} }
+2 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "vnt" name = "vnt"
version = "1.2.0" version = "1.2.2"
edition = "2021" edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
@@ -24,6 +24,7 @@ tokio = { version = "1.28.1", features = ["full"] }
aes-gcm = {version="0.10.2", optional = true} aes-gcm = {version="0.10.2", optional = true}
ring = {version="0.16.20", optional = true} ring = {version="0.16.20", optional = true}
cbc = "0.1.2" cbc = "0.1.2"
ecb = "0.1.2"
aes = "0.8.3" aes = "0.8.3"
stun-format = {version="1.0.1",features=["fmt","rfc3489"]} stun-format = {version="1.0.1",features=["fmt","rfc3489"]}
rsa = {version="0.7.2", features = [] } rsa = {version="0.7.2", features = [] }
+7 -7
View File
@@ -3,12 +3,12 @@ use std::{fmt, io};
/// 地址解析协议,由IP地址找到MAC地址 /// 地址解析协议,由IP地址找到MAC地址
/// https://www.ietf.org/rfc/rfc6747.txt /// https://www.ietf.org/rfc/rfc6747.txt
/* /*
0 2 4 5 6 8 10 (字节) 0 2 4 5 6 8 10 (字节)
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| 硬件类型|协议类型|硬件地址长度|协议地址长度|操作类型| | 硬件类型|协议类型|硬件地址长度|协议地址长度|操作类型|
| 源MAC地址 | 源ip地址 | | 源MAC地址 | 源ip地址 |
| 目的MAC地址 | 目的ip地址 | | 目的MAC地址 | 目的ip地址 |
*/ */
pub struct ArpPacket<B> { pub struct ArpPacket<B> {
buffer: B, buffer: B,
@@ -119,4 +119,4 @@ impl<B: AsRef<[u8]>> fmt::Debug for ArpPacket<B> {
.field("target_protocol_addr", &self.target_protocol_addr()) .field("target_protocol_addr", &self.target_protocol_addr())
.finish() .finish()
} }
} }
+1 -1
View File
@@ -1 +1 @@
pub mod arp; pub mod arp;
+1 -1
View File
@@ -1,2 +1,2 @@
pub mod packet; pub mod packet;
pub mod protocol; pub mod protocol;
+6 -6
View File
@@ -1,13 +1,13 @@
use std::{fmt, io};
use crate::ethernet::protocol::Protocol; use crate::ethernet::protocol::Protocol;
use std::{fmt, io};
/// 以太网帧协议 /// 以太网帧协议
/// https://www.ietf.org/rfc/rfc894.txt /// https://www.ietf.org/rfc/rfc894.txt
/* /*
0 6 12 14 (字节) 0 6 12 14 (字节)
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| 目的地址 | 源地址 | 类型 | | 目的地址 | 源地址 | 类型 |
*/ */
pub struct EthernetPacket<B> { pub struct EthernetPacket<B> {
pub buffer: B, pub buffer: B,
} }
@@ -74,4 +74,4 @@ impl<B: AsRef<[u8]>> fmt::Debug for EthernetPacket<B> {
.field("payload", &self.payload()) .field("payload", &self.payload())
.finish() .finish()
} }
} }
+24 -24
View File
@@ -102,7 +102,7 @@ impl From<u16> for Protocol {
0x88f7 => Ptp, 0x88f7 => Ptp,
0x8902 => Cfm, 0x8902 => Cfm,
0x9100 => QinQ, 0x9100 => QinQ,
n => Unknown(n), n => Unknown(n),
} }
} }
} }
@@ -112,30 +112,30 @@ impl Into<u16> for Protocol {
use self::Protocol::*; use self::Protocol::*;
match self { match self {
Ipv4 => 0x0800, Ipv4 => 0x0800,
Arp => 0x0806, Arp => 0x0806,
WakeOnLan => 0x0842, WakeOnLan => 0x0842,
Trill => 0x22f3, Trill => 0x22f3,
DecNet => 0x6003, DecNet => 0x6003,
Rarp => 0x8035, Rarp => 0x8035,
AppleTalk => 0x809b, AppleTalk => 0x809b,
Aarp => 0x80f3, Aarp => 0x80f3,
Ipx => 0x8137, Ipx => 0x8137,
Qnx => 0x8204, Qnx => 0x8204,
Ipv6 => 0x86dd, Ipv6 => 0x86dd,
FlowControl => 0x8808, FlowControl => 0x8808,
CobraNet => 0x8819, CobraNet => 0x8819,
Mpls => 0x8847, Mpls => 0x8847,
MplsMulticast => 0x8848, MplsMulticast => 0x8848,
PppoeDiscovery => 0x8863, PppoeDiscovery => 0x8863,
PppoeSession => 0x8864, PppoeSession => 0x8864,
Vlan => 0x8100, Vlan => 0x8100,
PBridge => 0x88a8, PBridge => 0x88a8,
Lldp => 0x88cc, Lldp => 0x88cc,
Ptp => 0x88f7, Ptp => 0x88f7,
Cfm => 0x8902, Cfm => 0x8902,
QinQ => 0x9100, QinQ => 0x9100,
Unknown(n) => n, Unknown(n) => n,
} }
} }
} }
+8 -8
View File
@@ -1,8 +1,8 @@
use std::{fmt, io};
use byteorder::{BigEndian, ReadBytesExt};
use crate::cal_checksum; use crate::cal_checksum;
use crate::icmp::{Code, Kind}; use crate::icmp::{Code, Kind};
use crate::ip::ipv4::packet::IpV4Packet; use crate::ip::ipv4::packet::IpV4Packet;
use byteorder::{BigEndian, ReadBytesExt};
use std::{fmt, io};
/// icmp 协议 /// icmp 协议
/* https://www.rfc-editor.org/rfc/rfc792 /* https://www.rfc-editor.org/rfc/rfc792
@@ -67,7 +67,7 @@ impl<B: AsRef<[u8]>> IcmpPacket<B> {
| Kind::TimestampReply | Kind::TimestampReply
| Kind::InformationRequest | Kind::InformationRequest
| Kind::InformationReply => { | Kind::InformationReply => {
let ide =u16::from_be_bytes(self.buffer.as_ref()[4..6].try_into().unwrap()); let ide = u16::from_be_bytes(self.buffer.as_ref()[4..6].try_into().unwrap());
let seq = u16::from_be_bytes(self.buffer.as_ref()[6..8].try_into().unwrap()); let seq = u16::from_be_bytes(self.buffer.as_ref()[6..8].try_into().unwrap());
HeaderOther::Identifier(ide, seq) HeaderOther::Identifier(ide, seq)
} }
@@ -121,11 +121,11 @@ impl<B: AsRef<[u8]>> fmt::Debug for IcmpPacket<B> {
} else { } else {
"icmp::Packet!" "icmp::Packet!"
}) })
.field("kind", &self.kind()) .field("kind", &self.kind())
.field("code", &self.code()) .field("code", &self.code())
.field("checksum", &self.checksum()) .field("checksum", &self.checksum())
.field("payload", &self.payload()) .field("payload", &self.payload())
.finish() .finish()
} }
} }
+12 -12
View File
@@ -1,17 +1,17 @@
use std::{fmt, io};
use std::net::Ipv4Addr;
use crate::cal_checksum; use crate::cal_checksum;
use std::net::Ipv4Addr;
use std::{fmt, io};
/// igmp v1 /// igmp v1
/* https://datatracker.ietf.org/doc/html/rfc1112 /* https://datatracker.ietf.org/doc/html/rfc1112
0 1 2 3 0 1 2 3
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 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
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
|Version| Type | Unused | Checksum | |Version| Type | Unused | Checksum |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| Group Address | | Group Address |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
*/ */
/// v1版本的报文 /// v1版本的报文
pub struct IgmpV1Packet<B> { pub struct IgmpV1Packet<B> {
pub buffer: B, pub buffer: B,
@@ -43,7 +43,7 @@ impl Into<u8> for IgmpV1Type {
match self { match self {
IgmpV1Type::Query => 0x11, IgmpV1Type::Query => 0x11,
IgmpV1Type::ReportV1 => 0x12, IgmpV1Type::ReportV1 => 0x12,
IgmpV1Type::Unknown(v) => v IgmpV1Type::Unknown(v) => v,
} }
} }
} }
@@ -114,4 +114,4 @@ impl<B: AsRef<[u8]>> fmt::Debug for IgmpV1Packet<B> {
.field("group_address", &self.group_address()) .field("group_address", &self.group_address())
.finish() .finish()
} }
} }
+11 -11
View File
@@ -1,18 +1,18 @@
use std::{fmt, io};
use std::net::Ipv4Addr;
use crate::cal_checksum; use crate::cal_checksum;
use std::net::Ipv4Addr;
use std::{fmt, io};
/// igmp v2 /// igmp v2
/* https://www.rfc-editor.org/rfc/rfc2236.html /* https://www.rfc-editor.org/rfc/rfc2236.html
0 1 2 3 0 1 2 3
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 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
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| Type | Max Resp Time | Checksum | | Type | Max Resp Time | Checksum |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| Group Address | | Group Address |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
*/ */
/// v2版本的报文 /// v2版本的报文
pub struct IgmpV2Packet<B> { pub struct IgmpV2Packet<B> {
@@ -48,7 +48,7 @@ impl Into<u8> for IgmpV2Type {
IgmpV2Type::Query => 0x11, IgmpV2Type::Query => 0x11,
IgmpV2Type::ReportV2 => 0x16, IgmpV2Type::ReportV2 => 0x16,
IgmpV2Type::LeaveV2 => 0x17, IgmpV2Type::LeaveV2 => 0x17,
IgmpV2Type::Unknown(v) => v IgmpV2Type::Unknown(v) => v,
} }
} }
} }
+8 -6
View File
@@ -1,5 +1,5 @@
use std::{fmt, io};
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::{fmt, io};
use crate::cal_checksum; use crate::cal_checksum;
@@ -116,7 +116,7 @@ impl Into<u8> for IgmpV3Type {
match self { match self {
IgmpV3Type::Query => 0x11, IgmpV3Type::Query => 0x11,
IgmpV3Type::ReportV3 => 0x22, IgmpV3Type::ReportV3 => 0x22,
IgmpV3Type::Unknown(v) => v IgmpV3Type::Unknown(v) => v,
} }
} }
} }
@@ -203,7 +203,7 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> IgmpV3QueryPacket<B> {
self.buffer.as_mut()[2..4].copy_from_slice(&checksum.to_be_bytes()) self.buffer.as_mut()[2..4].copy_from_slice(&checksum.to_be_bytes())
} }
pub fn set_qrv(&mut self, qrv: u8) { pub fn set_qrv(&mut self, qrv: u8) {
self.buffer.as_mut()[8] = (self.buffer.as_ref()[8]&(!0x07)) | (qrv & 0x07) self.buffer.as_mut()[8] = (self.buffer.as_ref()[8] & (!0x07)) | (qrv & 0x07)
} }
pub fn set_qqic(&mut self, qqic: u8) { pub fn set_qqic(&mut self, qqic: u8) {
self.buffer.as_mut()[9] = qqic self.buffer.as_mut()[9] = qqic
@@ -349,7 +349,10 @@ impl<B: AsRef<[u8]>> IgmpV3ReportPacket<B> {
return None; return None;
} }
if let Ok(record) = IgmpV3RecordPacket::new(&buf[start..]) { if let Ok(record) = IgmpV3RecordPacket::new(&buf[start..]) {
let end = start + 8 + record.aux_data_len() as usize * 4 + record.source_number() as usize * 4; let end = start
+ 8
+ record.aux_data_len() as usize * 4
+ record.source_number() as usize * 4;
if end > len { if end > len {
return None; return None;
} }
@@ -364,7 +367,6 @@ impl<B: AsRef<[u8]>> IgmpV3ReportPacket<B> {
} }
} }
/// group record /// group record
pub struct IgmpV3RecordPacket<B> { pub struct IgmpV3RecordPacket<B> {
pub buffer: B, pub buffer: B,
@@ -488,4 +490,4 @@ impl<B: AsRef<[u8]>> fmt::Debug for IgmpV3RecordPacket<B> {
.field("auxiliary_data", &self.auxiliary_data()) .field("auxiliary_data", &self.auxiliary_data())
.finish() .finish()
} }
} }
+3 -3
View File
@@ -2,7 +2,7 @@ pub mod igmp_v1;
pub mod igmp_v2; pub mod igmp_v2;
pub mod igmp_v3; pub mod igmp_v3;
#[derive(Debug,Copy, Clone,Eq, PartialEq)] #[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub enum IgmpType { pub enum IgmpType {
/// 0x11 所有组224.0.0.1或者特定组 /// 0x11 所有组224.0.0.1或者特定组
Query, Query,
@@ -40,7 +40,7 @@ impl Into<u8> for IgmpType {
IgmpType::ReportV2 => 0x16, IgmpType::ReportV2 => 0x16,
IgmpType::ReportV3 => 0x22, IgmpType::ReportV3 => 0x22,
IgmpType::LeaveV2 => 0x17, IgmpType::LeaveV2 => 0x17,
IgmpType::Unknown(v) => v IgmpType::Unknown(v) => v,
} }
} }
} }
+1 -2
View File
@@ -1,6 +1,5 @@
use std::{fmt, io};
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::{fmt, io};
use crate::cal_checksum; use crate::cal_checksum;
use crate::ip::ipv4::protocol::Protocol; use crate::ip::ipv4::protocol::Protocol;
+1 -1
View File
@@ -1,4 +1,4 @@
#[derive(Eq, PartialEq,Ord, PartialOrd, Copy, Clone, Debug)] #[derive(Eq, PartialEq, Ord, PartialOrd, Copy, Clone, Debug)]
pub enum Protocol { pub enum Protocol {
/// ///
Hopopt, Hopopt,
+1 -1
View File
@@ -1,5 +1,5 @@
use std::io;
use ipv4::packet::IpV4Packet; use ipv4::packet::IpV4Packet;
use std::io;
pub mod ipv4; pub mod ipv4;
+2 -2
View File
@@ -3,13 +3,13 @@ use std::net::Ipv4Addr;
use byteorder::BigEndian; use byteorder::BigEndian;
use byteorder::ReadBytesExt; use byteorder::ReadBytesExt;
pub mod arp;
pub mod ethernet;
pub mod icmp; pub mod icmp;
pub mod igmp; pub mod igmp;
pub mod ip; pub mod ip;
pub mod tcp; pub mod tcp;
pub mod udp; pub mod udp;
pub mod ethernet;
pub mod arp;
// pub enum IpUpperLayer<B> { // pub enum IpUpperLayer<B> {
// UDP(UdpPacket<B>), // UDP(UdpPacket<B>),
// Unknown(B), // Unknown(B),
+6 -2
View File
@@ -1,5 +1,5 @@
use std::{fmt, io};
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::{fmt, io};
use crate::tcp::Flags; use crate::tcp::Flags;
@@ -58,7 +58,11 @@ impl<B: AsRef<[u8]>> TcpPacket<B> {
buffer, buffer,
} }
} }
pub fn new(source_ip: Ipv4Addr, destination_ip: Ipv4Addr, buffer: B) -> io::Result<TcpPacket<B>> { pub fn new(
source_ip: Ipv4Addr,
destination_ip: Ipv4Addr,
buffer: B,
) -> io::Result<TcpPacket<B>> {
let packet = TcpPacket::unchecked(source_ip, destination_ip, buffer); let packet = TcpPacket::unchecked(source_ip, destination_ip, buffer);
if packet.buffer.as_ref().len() < 20 { if packet.buffer.as_ref().len() < 20 {
+6 -2
View File
@@ -1,5 +1,5 @@
use std::{fmt, io};
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::{fmt, io};
/// udp协议 /// udp协议
/// ///
@@ -60,7 +60,11 @@ impl<B: AsRef<[u8]>> UdpPacket<B> {
buffer, buffer,
} }
} }
pub fn new(source_ip: Ipv4Addr, destination_ip: Ipv4Addr, buffer: B) -> io::Result<UdpPacket<B>> { pub fn new(
source_ip: Ipv4Addr,
destination_ip: Ipv4Addr,
buffer: B,
) -> io::Result<UdpPacket<B>> {
if buffer.as_ref().len() < 8 { if buffer.as_ref().len() < 8 {
Err(io::Error::from(io::ErrorKind::InvalidData))?; Err(io::Error::from(io::ErrorKind::InvalidData))?;
} }
+2 -1
View File
@@ -54,7 +54,8 @@ message PunchInfo{
bool reply = 6; bool reply = 6;
fixed32 local_ip = 7; fixed32 local_ip = 7;
uint32 local_port = 8; uint32 local_port = 8;
repeated bytes public_ipv6_list = 9; bytes ipv6 = 9;
uint32 ipv6_port = 10;
} }
enum PunchNatType{ enum PunchNatType{
Symmetric = 0; Symmetric = 0;
+1 -1
View File
@@ -19,7 +19,7 @@ use crate::error::*;
/// A TUN device. /// A TUN device.
pub trait Device { pub trait Device {
type Queue ; type Queue;
/// Reconfigure the device. /// Reconfigure the device.
fn configure(&mut self, config: &Configuration) -> Result<()> { fn configure(&mut self, config: &Configuration) -> Result<()> {
+4 -4
View File
@@ -77,10 +77,10 @@ impl Device {
req.ifru.flags = device_type req.ifru.flags = device_type
| if config.platform.packet_information { | if config.platform.packet_information {
0 0
} else { } else {
IFF_NO_PI IFF_NO_PI
} }
| if queues_num > 1 { IFF_MULTI_QUEUE } else { 0 }; | if queues_num > 1 { IFF_MULTI_QUEUE } else { 0 };
for _ in 0..queues_num { for _ in 0..queues_num {
+1 -1
View File
@@ -22,7 +22,7 @@ use std::ptr;
use std::sync::Arc; use std::sync::Arc;
use libc; use libc;
use libc::{AF_INET, c_char, c_uint, c_void, SOCK_DGRAM, sockaddr, socklen_t}; use libc::{c_char, c_uint, c_void, sockaddr, socklen_t, AF_INET, SOCK_DGRAM};
use crate::configuration::{Configuration, Layer}; use crate::configuration::{Configuration, Layer};
use crate::device::Device as D; use crate::device::Device as D;
-1
View File
@@ -27,7 +27,6 @@ pub mod macos;
#[cfg(target_os = "macos")] #[cfg(target_os = "macos")]
pub use self::macos::{create, Configuration, Device, Queue}; pub use self::macos::{create, Configuration, Device, Queue};
#[cfg(test)] #[cfg(test)]
mod test { mod test {
use crate::configuration::Configuration; use crate::configuration::Configuration;
+1 -2
View File
@@ -14,7 +14,7 @@
use std::io; use std::io;
use std::mem; use std::mem;
use std::os::unix::io::{AsRawFd,RawFd}; use std::os::unix::io::{AsRawFd, RawFd};
use std::sync::Arc; use std::sync::Arc;
use crate::platform::posix::Fd; use crate::platform::posix::Fd;
@@ -72,7 +72,6 @@ impl Writer {
} }
} }
pub fn write_vectored(&self, bufs: &[io::IoSlice<'_>]) -> io::Result<usize> { pub fn write_vectored(&self, bufs: &[io::IoSlice<'_>]) -> io::Result<usize> {
unsafe { unsafe {
let mut msg: libc::msghdr = mem::zeroed(); let mut msg: libc::msghdr = mem::zeroed();
+194 -99
View File
@@ -2,28 +2,32 @@ use std::io;
use std::net::{Ipv4Addr, SocketAddr}; use std::net::{Ipv4Addr, SocketAddr};
use std::sync::Arc; use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use byte_pool::{Block, BytePool};
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap; use dashmap::DashMap;
use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpStream, UdpSocket};
use tokio::net::tcp::OwnedReadHalf; use tokio::net::tcp::OwnedReadHalf;
use tokio::net::{TcpStream, UdpSocket};
use tokio::sync::watch::{channel, Receiver, Sender}; use tokio::sync::watch::{channel, Receiver, Sender};
use crate::channel::{Route, RouteKey, Status};
use crate::channel::punch::NatType; use crate::channel::punch::NatType;
use crate::channel::{Route, RouteKey, Status};
use crate::core::status::VntWorker; use crate::core::status::VntWorker;
use crate::handle::CurrentDeviceInfo;
use crate::handle::recv_handler::ChannelDataHandler; use crate::handle::recv_handler::ChannelDataHandler;
use byte_pool::{Block, BytePool}; use crate::handle::CurrentDeviceInfo;
lazy_static::lazy_static! { lazy_static::lazy_static! {
static ref POOL:BytePool = BytePool::new(); static ref POOL:BytePool = BytePool::new();
} }
pub struct ContextInner { pub struct ContextInner {
//udp用于打洞、服务端通信(可选) //udp用于打洞、服务端通信(可选)
pub(crate) main_channel: Arc<UdpSocket>, pub(crate) main_channel: Arc<UdpSocket>,
pub(crate) main_channel_ipv6: Option<Arc<UdpSocket>>,
//在udp的基础上,可以选择使用tcp和服务端通信 //在udp的基础上,可以选择使用tcp和服务端通信
pub(crate) main_tcp_channel: Option<tokio::sync::mpsc::Sender<Vec<u8>>>, pub(crate) main_tcp_channel: Option<tokio::sync::mpsc::Sender<Vec<u8>>>,
pub(crate) route_table: DashMap<Ipv4Addr, Vec<Route>>, pub(crate) route_table: DashMap<Ipv4Addr, Vec<Route>>,
pub(crate) route_table_time: DashMap<(RouteKey, Ipv4Addr), AtomicCell<Instant>>, pub(crate) route_table_time: DashMap<(RouteKey, Ipv4Addr), Instant>,
pub(crate) status_receiver: Receiver<Status>, pub(crate) status_receiver: Receiver<Status>,
pub(crate) status_sender: Sender<Status>, pub(crate) status_sender: Sender<Status>,
pub(crate) udp_map: DashMap<usize, Arc<UdpSocket>>, pub(crate) udp_map: DashMap<usize, Arc<UdpSocket>>,
@@ -37,12 +41,19 @@ pub struct Context {
} }
impl Context { impl Context {
pub fn new(main_channel: Arc<UdpSocket>, main_tcp_channel: Option<tokio::sync::mpsc::Sender<Vec<u8>>>, current_device: Arc<AtomicCell<CurrentDeviceInfo>>, _channel_num: usize) -> Self { pub fn new(
main_channel: Arc<UdpSocket>,
main_channel_ipv6: Option<Arc<UdpSocket>>,
main_tcp_channel: Option<tokio::sync::mpsc::Sender<Vec<u8>>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
_channel_num: usize,
) -> Self {
//当前版本只支持一个通道 //当前版本只支持一个通道
let channel_num = 1; let channel_num = 1;
let (status_sender, status_receiver) = channel(Status::Cone); let (status_sender, status_receiver) = channel(Status::Cone);
let inner = Arc::new(ContextInner { let inner = Arc::new(ContextInner {
main_channel, main_channel,
main_channel_ipv6,
main_tcp_channel, main_tcp_channel,
route_table: DashMap::with_capacity(16), route_table: DashMap::with_capacity(16),
route_table_time: DashMap::with_capacity(16), route_table_time: DashMap::with_capacity(16),
@@ -52,9 +63,7 @@ impl Context {
channel_num, channel_num,
current_device, current_device,
}); });
Self { Self { inner }
inner
}
} }
} }
@@ -87,11 +96,37 @@ impl Context {
pub fn switch_to_symmetric(&self) { pub fn switch_to_symmetric(&self) {
let _ = self.inner.status_sender.send(Status::Symmetric); let _ = self.inner.status_sender.send(Status::Symmetric);
} }
pub fn main_local_port(&self) -> io::Result<u16> { pub fn main_local_ipv4_port(&self) -> io::Result<u16> {
self.inner.main_channel.local_addr().map(|k| k.port()) self.inner.main_channel.local_addr().map(|k| k.port())
} }
pub fn main_local_ipv6_port(&self) -> io::Result<u16> {
if let Some(ipv6) = &self.inner.main_channel_ipv6 {
ipv6.local_addr().map(|k| k.port())
} else {
Err(io::Error::new(io::ErrorKind::Other, "not ipv6"))
}
}
pub async fn send_main_udp(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> { pub async fn send_main_udp(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> {
self.inner.main_channel.send_to(buf, addr).await if addr.is_ipv6() {
if let Some(udp_ipv6) = &self.inner.main_channel_ipv6 {
udp_ipv6.send_to(buf, addr).await
} else {
Err(io::Error::new(io::ErrorKind::Other, "not ipv6"))
}
} else {
self.inner.main_channel.send_to(buf, addr).await
}
}
pub fn try_send_main_udp(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> {
if addr.is_ipv6() {
if let Some(udp_ipv6) = &self.inner.main_channel_ipv6 {
udp_ipv6.try_send_to(buf, addr)
} else {
Err(io::Error::new(io::ErrorKind::Other, "not ipv6"))
}
} else {
self.inner.main_channel.try_send_to(buf, addr)
}
} }
pub async fn send_main(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> { pub async fn send_main(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> {
if let Some(sender) = &self.inner.main_tcp_channel { if let Some(sender) = &self.inner.main_tcp_channel {
@@ -101,7 +136,7 @@ impl Context {
Err(io::Error::new(io::ErrorKind::Other, "send_main err")) Err(io::Error::new(io::ErrorKind::Other, "send_main err"))
} }
} else { } else {
self.inner.main_channel.send_to(buf, addr).await self.send_main_udp(buf, addr).await
} }
} }
pub fn try_send_main(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> { pub fn try_send_main(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> {
@@ -112,7 +147,7 @@ impl Context {
Err(io::Error::new(io::ErrorKind::Other, "try_send_main err")) Err(io::Error::new(io::ErrorKind::Other, "try_send_main err"))
} }
} else { } else {
self.inner.main_channel.try_send_to(buf, addr) self.try_send_main_udp(buf, addr)
} }
} }
@@ -120,7 +155,8 @@ impl Context {
for udp_ref in self.inner.udp_map.iter() { for udp_ref in self.inner.udp_map.iter() {
let udp = udp_ref.clone(); let udp = udp_ref.clone();
drop(udp_ref); drop(udp_ref);
udp.send_to(buf, addr).await?; //使用ipv6的udp发送ipv4报文会出错
let _ = udp.send_to(buf, addr).await;
} }
Ok(()) Ok(())
} }
@@ -132,20 +168,19 @@ impl Context {
} }
let route = v.value()[0]; let route = v.value()[0];
drop(v); drop(v);
if route.rt == 199 {
//这通常是刚加入路由,直接放弃使用,避免抖动
return Err(io::Error::new(io::ErrorKind::NotFound, "route not found"));
}
if !route.is_p2p() { if !route.is_p2p() {
if let Some(time) = self.inner.route_table_time.get(&(route.route_key(), *id)) { if let Some(time) = self.inner.route_table_time.get(&(route.route_key(), *id)) {
//借道传输时,长时间不通信的通道不使用 //借道传输时,长时间不通信的通道不使用
if time.value().load().elapsed() > Duration::from_secs(6) { if time.value().elapsed() > Duration::from_secs(6) {
return Err(io::Error::new(io::ErrorKind::NotFound, "route time out")); return Err(io::Error::new(io::ErrorKind::NotFound, "route time out"));
} }
} }
} }
return self.send_by_key(buf, &route.route_key()).await;
if let Some(udp_ref) = self.inner.udp_map.get(&route.index) {
let udp = udp_ref.value().clone();
drop(udp_ref);
return udp.send_to(buf, route.addr).await;
}
} }
Err(io::Error::new(io::ErrorKind::NotFound, "route not found")) Err(io::Error::new(io::ErrorKind::NotFound, "route not found"))
} }
@@ -207,7 +242,11 @@ impl Context {
} }
fn add_route_(&self, id: Ipv4Addr, route: Route, only_if_absent: bool) { fn add_route_(&self, id: Ipv4Addr, route: Route, only_if_absent: bool) {
let key = route.route_key(); let key = route.route_key();
let mut list = self.inner.route_table.entry(id).or_insert_with(|| Vec::with_capacity(4)); let mut list = self
.inner
.route_table
.entry(id)
.or_insert_with(|| Vec::with_capacity(4));
let mut exist = false; let mut exist = false;
for x in list.iter_mut() { for x in list.iter_mut() {
if x.metric < route.metric { if x.metric < route.metric {
@@ -238,7 +277,9 @@ impl Context {
list.truncate(max_len); list.truncate(max_len);
} }
} }
self.inner.route_table_time.insert((key, id), AtomicCell::new(Instant::now())); self.inner
.route_table_time
.insert((key, id), Instant::now());
} }
pub fn route(&self, id: &Ipv4Addr) -> Option<Vec<Route>> { pub fn route(&self, id: &Ipv4Addr) -> Option<Vec<Route>> {
if let Some(v) = self.inner.route_table.get(id) { if let Some(v) = self.inner.route_table.get(id) {
@@ -273,7 +314,11 @@ impl Context {
true true
} }
pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec<Route>)> { pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec<Route>)> {
self.inner.route_table.iter().map(|k| (k.key().clone(), k.value().clone())).collect() self.inner
.route_table
.iter()
.map(|k| (k.key().clone(), k.value().clone()))
.collect()
} }
pub fn route_table_one(&self) -> Vec<(Ipv4Addr, Route)> { pub fn route_table_one(&self) -> Vec<(Ipv4Addr, Route)> {
let mut v = Vec::with_capacity(8); let mut v = Vec::with_capacity(8);
@@ -312,8 +357,8 @@ impl Context {
self.inner.route_table_time.remove(&(route_key, *id)); self.inner.route_table_time.remove(&(route_key, *id));
} }
pub fn update_read_time(&self, id: &Ipv4Addr, route_key: &RouteKey) { pub fn update_read_time(&self, id: &Ipv4Addr, route_key: &RouteKey) {
if let Some(time) = self.inner.route_table_time.get(&(*route_key, *id)) { if let Some(mut time) = self.inner.route_table_time.get_mut(&(*route_key, *id)) {
time.value().store(Instant::now()); *time.value_mut() = Instant::now();
} }
} }
} }
@@ -324,17 +369,16 @@ pub struct Channel {
} }
impl Channel { impl Channel {
pub fn new(context: Context, pub fn new(context: Context, handler: ChannelDataHandler) -> Self {
handler: ChannelDataHandler, ) -> Self { Self { context, handler }
Self {
context,
handler,
}
} }
} }
#[derive(Clone)] #[derive(Clone)]
struct BufSenderGroup(usize, Vec<tokio::sync::mpsc::Sender<(Block<'static>, usize, usize, RouteKey)>>); struct BufSenderGroup(
usize,
Vec<tokio::sync::mpsc::Sender<(Block<'static>, usize, usize, RouteKey)>>,
);
struct BufReceiverGroup(Vec<tokio::sync::mpsc::Receiver<(Block<'static>, usize, usize, RouteKey)>>); struct BufReceiverGroup(Vec<tokio::sync::mpsc::Receiver<(Block<'static>, usize, usize, RouteKey)>>);
@@ -350,15 +394,23 @@ fn buf_channel_group(size: usize) -> (BufSenderGroup, BufReceiverGroup) {
let mut buf_sender_group = Vec::with_capacity(size); let mut buf_sender_group = Vec::with_capacity(size);
let mut buf_receiver_group = Vec::with_capacity(size); let mut buf_receiver_group = Vec::with_capacity(size);
for _ in 0..size { for _ in 0..size {
let (buf_sender, buf_receiver) = tokio::sync::mpsc::channel::<(Block<'static, Vec<u8>>, usize, usize, RouteKey)>(10); let (buf_sender, buf_receiver) =
tokio::sync::mpsc::channel::<(Block<'static, Vec<u8>>, usize, usize, RouteKey)>(10);
buf_sender_group.push(buf_sender); buf_sender_group.push(buf_sender);
buf_receiver_group.push(buf_receiver); buf_receiver_group.push(buf_receiver);
} }
(BufSenderGroup(0, buf_sender_group), BufReceiverGroup(buf_receiver_group)) (
BufSenderGroup(0, buf_sender_group),
BufReceiverGroup(buf_receiver_group),
)
} }
impl Channel { impl Channel {
async fn tcp_handle(mut tcp_r: OwnedReadHalf, mut buf_sender: BufSenderGroup, head_reserve: usize) -> io::Result<()> { async fn tcp_handle(
mut tcp_r: OwnedReadHalf,
mut buf_sender: BufSenderGroup,
head_reserve: usize,
) -> io::Result<()> {
let mut head = [0; 4]; let mut head = [0; 4];
let addr = tcp_r.peer_addr()?; let addr = tcp_r.peer_addr()?;
let key = RouteKey::new(0, addr); let key = RouteKey::new(0, addr);
@@ -372,21 +424,34 @@ impl Channel {
"length overflow", "length overflow",
)); ));
} }
tcp_r.read_exact(&mut buf[head_reserve..head_reserve + len]).await?; tcp_r
if !buf_sender.send((buf, head_reserve, head_reserve + len, key)).await { .read_exact(&mut buf[head_reserve..head_reserve + len])
return Err(io::Error::new(io::ErrorKind::Other, "buf_sender发送数据失败")); .await?;
if !buf_sender
.send((buf, head_reserve, head_reserve + len, key))
.await
{
return Err(io::Error::new(
io::ErrorKind::Other,
"buf_sender发送数据失败",
));
} }
} }
} }
async fn start_tcp(mut worker: VntWorker, tcp_stream: TcpStream, mut receiver: tokio::sync::mpsc::Receiver<Vec<u8>>, async fn start_tcp(
current_device: Arc<AtomicCell<CurrentDeviceInfo>>, mut worker: VntWorker,
buf_sender: BufSenderGroup, head_reserve: usize) { tcp_stream: TcpStream,
mut receiver: tokio::sync::mpsc::Receiver<Vec<u8>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
buf_sender: BufSenderGroup,
head_reserve: usize,
) {
let (tcp_r, mut tcp_w) = tcp_stream.into_split(); let (tcp_r, mut tcp_w) = tcp_stream.into_split();
{ {
let buf_sender = buf_sender.clone(); let buf_sender = buf_sender.clone();
tokio::spawn(async move { tokio::spawn(async move {
if let Err(e) = Self::tcp_handle(tcp_r, buf_sender, head_reserve).await { if let Err(e) = Self::tcp_handle(tcp_r, buf_sender, head_reserve).await {
log::info!("tcp链接断开:{:?}",e); log::info!("tcp链接断开:{:?}", e);
} }
}); });
} }
@@ -436,13 +501,14 @@ impl Channel {
worker.stop_all(); worker.stop_all();
} }
pub async fn start(self, pub async fn start(
mut worker: VntWorker, self,
tcp: Option<(TcpStream, tokio::sync::mpsc::Receiver<Vec<u8>>)>, mut worker: VntWorker,
head_reserve: usize,//头部预留字节 tcp: Option<(TcpStream, tokio::sync::mpsc::Receiver<Vec<u8>>)>,
symmetric_channel_num: usize,//对称网络,则再加一组监听,提升打洞成功率 head_reserve: usize, //头部预留字节
relay: bool, symmetric_channel_num: usize, //对称网络,则再加一组监听,提升打洞成功率
parallel: usize, relay: bool,
parallel: usize,
) { ) {
let handler = self.handler.clone(); let handler = self.handler.clone();
let context = self.context; let context = self.context;
@@ -454,7 +520,9 @@ impl Channel {
let handler = handler.clone(); let handler = handler.clone();
tokio::spawn(async move { tokio::spawn(async move {
while let Some((mut buf, start, end, route_key)) = buf_receiver.recv().await { while let Some((mut buf, start, end, route_key)) = buf_receiver.recv().await {
handler.handle(&mut buf, start, end, route_key, &context).await; handler
.handle(&mut buf, start, end, route_key, &context)
.await;
} }
}); });
} }
@@ -463,9 +531,35 @@ impl Channel {
None None
}; };
if let Some((tcp_stream, receiver)) = tcp { if let Some((tcp_stream, receiver)) = tcp {
tokio::spawn(Self::start_tcp(worker.worker("main_channel_tcp"), tcp_stream, receiver, context.inner.current_device.clone(), buf_sender.clone().unwrap(), head_reserve)); tokio::spawn(Self::start_tcp(
worker.worker("main_channel_tcp"),
tcp_stream,
receiver,
context.inner.current_device.clone(),
buf_sender.clone().unwrap(),
head_reserve,
));
} }
tokio::spawn(Self::start_(worker.worker("main_channel_1"), context.clone(), main_channel.clone(), handler.clone(), buf_sender.clone(), head_reserve, true)); if let Some(main_channel_ipv6) = &context.inner.main_channel_ipv6 {
tokio::spawn(Self::start_(
worker.worker("main_channel_ipv6"),
context.clone(),
main_channel_ipv6.clone(),
handler.clone(),
buf_sender.clone(),
head_reserve,
true,
));
}
tokio::spawn(Self::start_(
worker.worker("main_channel_1"),
context.clone(),
main_channel.clone(),
handler.clone(),
buf_sender.clone(),
head_reserve,
true,
));
if relay { if relay {
worker.stop_wait().await; worker.stop_wait().await;
return; return;
@@ -517,21 +611,24 @@ impl Channel {
} }
worker.stop_all(); worker.stop_all();
} }
async fn start_(mut worker: VntWorker, context: Context, async fn start_(
udp: Arc<UdpSocket>, mut worker: VntWorker,
handler: ChannelDataHandler, context: Context,
buf_sender: Option<BufSenderGroup>, udp: Arc<UdpSocket>,
head_reserve: usize, handler: ChannelDataHandler,
is_core: bool) { buf_sender: Option<BufSenderGroup>,
head_reserve: usize,
is_core: bool,
) {
let mut status_receiver = context.inner.status_receiver.clone(); let mut status_receiver = context.inner.status_receiver.clone();
#[cfg(target_os = "windows")] #[cfg(target_os = "windows")]
use std::os::windows::io::AsRawSocket; use std::os::windows::io::AsRawSocket;
#[cfg(target_os = "windows")] #[cfg(target_os = "windows")]
let id = 1 + udp.as_raw_socket() as usize; let id = 1 + udp.as_raw_socket() as usize;
#[cfg(any(unix))] #[cfg(any(unix))]
use std::os::fd::AsRawFd; use std::os::fd::AsRawFd;
#[cfg(any(unix))] #[cfg(any(unix))]
let id = 1 + udp.as_raw_fd() as usize; let id = 1 + udp.as_raw_fd() as usize;
context.inner.udp_map.insert(id, udp.clone()); context.inner.udp_map.insert(id, udp.clone());
match buf_sender { match buf_sender {
None => { None => {
@@ -574,49 +671,47 @@ impl Channel {
} }
} }
} }
Some(mut buf_sender) => { Some(mut buf_sender) => loop {
loop { let mut buf = POOL.alloc(4096);
let mut buf = POOL.alloc(4096); tokio::select! {
tokio::select! { rs=udp.recv_from(&mut buf[head_reserve..])=>{
rs=udp.recv_from(&mut buf[head_reserve..])=>{ match rs {
match rs { Ok((len, addr)) => {
Ok((len, addr)) => { if !buf_sender.send((buf,head_reserve,head_reserve+len,RouteKey::new(id, addr))).await{
if !buf_sender.send((buf,head_reserve,head_reserve+len,RouteKey::new(id, addr))).await{ log::error!("udp buf_sender发送数据失败");
log::error!("udp buf_sender发送数据失败"); break;
break;
}
}
Err(e) => {
log::error!("{:?}",e)
} }
} }
} Err(e) => {
changed=status_receiver.changed()=>{ log::error!("{:?}",e)
match changed { }
Ok(_) => {
match *status_receiver.borrow() {
Status::Cone => {
if !is_core{
break;
}
}
Status::Close=>{
break;
}
Status::Symmetric => {}
}
}
Err(_) => {
break;
}
}
}
_=worker.stop_wait()=>{
break;
} }
} }
changed=status_receiver.changed()=>{
match changed {
Ok(_) => {
match *status_receiver.borrow() {
Status::Cone => {
if !is_core{
break;
}
}
Status::Close=>{
break;
}
Status::Symmetric => {}
}
}
Err(_) => {
break;
}
}
}
_=worker.stop_wait()=>{
break;
}
} }
} },
} }
context.inner.udp_map.remove(&id); context.inner.udp_map.remove(&id);
if is_core { if is_core {
+6 -11
View File
@@ -1,10 +1,9 @@
use crate::channel::channel::Context;
use crate::channel::RouteKey;
use std::io; use std::io;
use std::io::{Error, ErrorKind}; use std::io::{Error, ErrorKind};
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::time::Duration; use std::time::Duration;
use crate::channel::channel::Context;
use crate::channel::RouteKey;
pub struct Idle { pub struct Idle {
read_idle: Duration, read_idle: Duration,
@@ -12,12 +11,8 @@ pub struct Idle {
} }
impl Idle { impl Idle {
pub fn new(read_idle: Duration, pub fn new(read_idle: Duration, context: Context) -> Self {
context: Context, ) -> Self { Self { read_idle, context }
Self {
read_idle,
context,
}
} }
} }
@@ -27,7 +22,7 @@ impl Idle {
loop { loop {
let mut max = Duration::from_secs(0); let mut max = Duration::from_secs(0);
for entry in self.context.inner.route_table_time.iter() { for entry in self.context.inner.route_table_time.iter() {
let last_read = entry.value().load().elapsed(); let last_read = entry.value().elapsed();
if last_read >= self.read_idle { if last_read >= self.read_idle {
return Ok((entry.key().1.clone(), entry.key().0.clone())); return Ok((entry.key().1.clone(), entry.key().0.clone()));
} else { } else {
@@ -45,4 +40,4 @@ impl Idle {
} }
} }
} }
} }
+5 -10
View File
@@ -1,8 +1,8 @@
use std::net::SocketAddr; use std::net::SocketAddr;
pub mod channel; pub mod channel;
pub mod punch;
pub mod idle; pub mod idle;
pub mod punch;
pub mod sender; pub mod sender;
#[derive(Copy, Clone, Eq, PartialEq)] #[derive(Copy, Clone, Eq, PartialEq)]
@@ -27,8 +27,7 @@ pub struct RouteSortKey {
} }
impl Route { impl Route {
pub fn new(index: usize, pub fn new(index: usize, addr: SocketAddr, metric: u8, rt: i64) -> Self {
addr: SocketAddr, metric: u8, rt: i64, ) -> Self {
Self { Self {
index, index,
addr, addr,
@@ -68,11 +67,7 @@ pub struct RouteKey {
} }
impl RouteKey { impl RouteKey {
pub(crate) fn new(index: usize, pub(crate) fn new(index: usize, addr: SocketAddr) -> Self {
addr: SocketAddr, ) -> Self { Self { index, addr }
Self {
index,
addr,
}
} }
} }
+75 -26
View File
@@ -1,19 +1,39 @@
use std::collections::HashMap; use std::collections::HashMap;
use std::io; use std::io;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4, SocketAddrV6};
use std::str::FromStr;
use std::time::Duration; use std::time::Duration;
use rand::prelude::SliceRandom; use rand::prelude::SliceRandom;
use crate::channel::channel::Context; use crate::channel::channel::Context;
#[derive(Copy, Clone, Eq, PartialEq, Debug)]
pub enum PunchModel {
IPv4,
IPv6,
All,
}
impl FromStr for PunchModel {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().trim() {
"ipv4" => Ok(PunchModel::IPv4),
"ipv6" => Ok(PunchModel::IPv6),
_ => Ok(PunchModel::All),
}
}
}
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
pub struct NatInfo { pub struct NatInfo {
pub public_ips: Vec<Ipv4Addr>, pub public_ips: Vec<Ipv4Addr>,
pub public_port: u16, pub public_port: u16,
pub public_port_range: u16, pub public_port_range: u16,
pub local_ip: Ipv4Addr, pub local_ipv4_addr: SocketAddrV4,
pub local_port: u16, pub ipv6_addr: SocketAddrV6,
pub nat_type: NatType, pub nat_type: NatType,
} }
@@ -24,21 +44,21 @@ pub enum NatType {
} }
impl NatInfo { impl NatInfo {
pub fn new(mut public_ips: Vec<Ipv4Addr>, pub fn new(
public_port: u16, mut public_ips: Vec<Ipv4Addr>,
public_port_range: u16, public_port: u16,
local_ip: Ipv4Addr, public_port_range: u16,
local_port: u16, local_ipv4_addr: SocketAddrV4,
nat_type: NatType, ) -> Self { ipv6_addr: SocketAddrV6,
public_ips.retain(|ip| { nat_type: NatType,
!ip.is_loopback() && !ip.is_private() ) -> Self {
}); public_ips.retain(|ip| !ip.is_loopback() && !ip.is_private());
Self { Self {
public_ips, public_ips,
public_port, public_port,
public_port_range, public_port_range,
local_ip, local_ipv4_addr,
local_port, ipv6_addr,
nat_type, nat_type,
} }
} }
@@ -49,10 +69,11 @@ pub struct Punch {
context: Context, context: Context,
port_vec: Vec<u16>, port_vec: Vec<u16>,
port_index: HashMap<Ipv4Addr, usize>, port_index: HashMap<Ipv4Addr, usize>,
punch_model: PunchModel,
} }
impl Punch { impl Punch {
pub fn new(context: Context) -> Self { pub fn new(context: Context, punch_model: PunchModel) -> Self {
let mut port_vec: Vec<u16> = (1..65535).collect(); let mut port_vec: Vec<u16> = (1..65535).collect();
port_vec.push(65535); port_vec.push(65535);
let mut rng = rand::thread_rng(); let mut rng = rand::thread_rng();
@@ -61,6 +82,7 @@ impl Punch {
context, context,
port_vec, port_vec,
port_index: HashMap::new(), port_index: HashMap::new(),
punch_model,
} }
} }
} }
@@ -70,8 +92,24 @@ impl Punch {
if !self.context.need_punch(&id) { if !self.context.need_punch(&id) {
return Ok(()); return Ok(());
} }
if !nat_info.local_ip.is_unspecified() || nat_info.local_port != 0 { if !nat_info.local_ipv4_addr.ip().is_unspecified() && nat_info.local_ipv4_addr.port() != 0 {
let _ = self.context.send_main_udp(buf, SocketAddr::V4(SocketAddrV4::new(nat_info.local_ip, nat_info.local_port))).await; let _ = self
.context
.send_main_udp(buf, SocketAddr::V4(nat_info.local_ipv4_addr))
.await;
}
if self.punch_model != PunchModel::IPv4
&& !nat_info.ipv6_addr.ip().is_unspecified()
&& nat_info.ipv6_addr.port() != 0
{
let rs = self
.context
.send_main_udp(buf, SocketAddr::V6(nat_info.ipv6_addr))
.await;
log::info!("发送到ipv6地址:{:?},rs={:?}", nat_info.ipv6_addr, rs);
if rs.is_ok() && self.punch_model == PunchModel::IPv6 {
return Ok(());
}
} }
match nat_info.nat_type { match nat_info.nat_type {
NatType::Symmetric => { NatType::Symmetric => {
@@ -91,12 +129,10 @@ impl Punch {
} else { } else {
1 1
}; };
let (max_port, overflow) = nat_info.public_port.overflowing_add(nat_info.public_port_range); let (max_port, overflow) = nat_info
let max_port = if overflow { .public_port
65535 .overflowing_add(nat_info.public_port_range);
} else { let max_port = if overflow { 65535 } else { max_port };
max_port
};
let k = if max_port - min_port + 1 > max_k1 { let k = if max_port - min_port + 1 > max_k1 {
max_k1 as usize max_k1 as usize
} else { } else {
@@ -108,7 +144,8 @@ impl Punch {
let mut rng = rand::thread_rng(); let mut rng = rand::thread_rng();
nums.shuffle(&mut rng); nums.shuffle(&mut rng);
} }
self.punch_symmetric(&nums[..k], buf, &nat_info.public_ips, max_k1 as usize).await?; self.punch_symmetric(&nums[..k], buf, &nat_info.public_ips, max_k1 as usize)
.await?;
} }
let start = *self.port_index.entry(id.clone()).or_insert(0); let start = *self.port_index.entry(id.clone()).or_insert(0);
let mut end = start + max_k2; let mut end = start + max_k2;
@@ -117,7 +154,13 @@ impl Punch {
end = self.port_vec.len(); end = self.port_vec.len();
index = 0 index = 0
} }
self.punch_symmetric(&self.port_vec[start..end], buf, &nat_info.public_ips, max_k2).await?; self.punch_symmetric(
&self.port_vec[start..end],
buf,
&nat_info.public_ips,
max_k2,
)
.await?;
self.port_index.insert(id, index); self.port_index.insert(id, index);
} }
NatType::Cone => { NatType::Cone => {
@@ -137,7 +180,13 @@ impl Punch {
Ok(()) Ok(())
} }
async fn punch_symmetric(&self, ports: &[u16], buf: &[u8], ips: &Vec<Ipv4Addr>, max: usize) -> io::Result<()> { async fn punch_symmetric(
&self,
ports: &[u16],
buf: &[u8],
ips: &Vec<Ipv4Addr>,
max: usize,
) -> io::Result<()> {
let mut count = 0; let mut count = 0;
for port in ports { for port in ports {
for pub_ip in ips { for pub_ip in ips {
+2 -4
View File
@@ -1,5 +1,5 @@
use std::ops::Deref;
use crate::channel::channel::Context; use crate::channel::channel::Context;
use std::ops::Deref;
#[derive(Clone)] #[derive(Clone)]
pub struct ChannelSender { pub struct ChannelSender {
@@ -8,9 +8,7 @@ pub struct ChannelSender {
impl ChannelSender { impl ChannelSender {
pub fn new(context: Context) -> Self { pub fn new(context: Context) -> Self {
Self { Self { context }
context,
}
} }
} }
+62 -35
View File
@@ -5,7 +5,7 @@ use rand::RngCore;
use crate::cipher::Finger; use crate::cipher::Finger;
use crate::protocol::body::AesCbcSecretBody; use crate::protocol::body::AesCbcSecretBody;
use crate::protocol::{HEAD_LEN, NetPacket}; use crate::protocol::{NetPacket, HEAD_LEN};
type Aes128CbcEnc = cbc::Encryptor<aes::Aes128>; type Aes128CbcEnc = cbc::Encryptor<aes::Aes128>;
type Aes128CbcDec = cbc::Decryptor<aes::Aes128>; type Aes128CbcDec = cbc::Decryptor<aes::Aes128>;
@@ -15,7 +15,7 @@ type Aes256CbcDec = cbc::Decryptor<aes::Aes256>;
#[derive(Clone)] #[derive(Clone)]
pub struct AesCbcCipher { pub struct AesCbcCipher {
pub(crate) cipher: AesCbcEnum, pub(crate) cipher: AesCbcEnum,
pub(crate) finger: Finger, pub(crate) finger: Option<Finger>,
} }
#[derive(Clone)] #[derive(Clone)]
@@ -27,33 +27,36 @@ pub enum AesCbcEnum {
impl AesCbcCipher { impl AesCbcCipher {
pub fn key(&self) -> &[u8] { pub fn key(&self) -> &[u8] {
match &self.cipher { match &self.cipher {
AesCbcEnum::AES128CBC(key) => { key } AesCbcEnum::AES128CBC(key) => key,
AesCbcEnum::AES256CBC(key) => { key } AesCbcEnum::AES256CBC(key) => key,
} }
} }
} }
impl AesCbcCipher { impl AesCbcCipher {
pub fn new_128(key: [u8; 16], finger: Finger) -> Self { pub fn new_128(key: [u8; 16], finger: Option<Finger>) -> Self {
Self { Self {
cipher: AesCbcEnum::AES128CBC(key), cipher: AesCbcEnum::AES128CBC(key),
finger, finger,
} }
} }
pub fn new_256(key: [u8; 32], finger: Finger) -> Self { pub fn new_256(key: [u8; 32], finger: Option<Finger>) -> Self {
Self { Self {
cipher: AesCbcEnum::AES256CBC(key), cipher: AesCbcEnum::AES256CBC(key),
finger, finger,
} }
} }
pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
if !net_packet.is_encrypt() { if !net_packet.is_encrypt() {
//未加密的数据直接丢弃 //未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
} }
if net_packet.payload().len() < 12 + 16 { if net_packet.payload().len() < 16 {
log::error!("数据异常,长度{}小于{}",net_packet.payload().len(),12+16); log::error!("数据异常,长度{}小于{}", net_packet.payload().len(), 16);
return Err(io::Error::new(io::ErrorKind::Other, "data err")); return Err(io::Error::new(io::ErrorKind::Other, "data err"));
} }
let mut iv = [0; 16]; let mut iv = [0; 16];
@@ -63,16 +66,23 @@ impl AesCbcCipher {
iv[9] = net_packet.transport_protocol(); iv[9] = net_packet.transport_protocol();
iv[10] = net_packet.is_gateway() as u8; iv[10] = net_packet.is_gateway() as u8;
iv[11] = net_packet.source_ttl(); iv[11] = net_packet.source_ttl();
iv[12..16].copy_from_slice(&self.finger.hash[0..4]); if let Some(finger) = &self.finger {
iv[12..16].copy_from_slice(&finger.hash[0..4]);
}
let mut secret_body = AesCbcSecretBody::new(net_packet.payload_mut())?; let mut secret_body =
let finger = self.finger.calculate_finger(&iv[..12], secret_body.en_body()); AesCbcSecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
if &finger != secret_body.finger() { if let Some(finger) = &self.finger {
return Err(io::Error::new(io::ErrorKind::Other, "finger err")); let finger = finger.calculate_finger(&iv[..12], secret_body.en_body());
if &finger != secret_body.finger() {
return Err(io::Error::new(io::ErrorKind::Other, "finger err"));
}
} }
let rs = match &self.cipher { let rs = match &self.cipher {
AesCbcEnum::AES128CBC(key) => { Aes128CbcDec::new(&(*key).into(), &iv.into()).decrypt_padded_mut::<Pkcs7>(secret_body.en_body_mut()) } AesCbcEnum::AES128CBC(key) => Aes128CbcDec::new(&(*key).into(), &iv.into())
AesCbcEnum::AES256CBC(key) => { Aes256CbcDec::new(&(*key).into(), &iv.into()).decrypt_padded_mut::<Pkcs7>(secret_body.en_body_mut()) } .decrypt_padded_mut::<Pkcs7>(secret_body.en_body_mut()),
AesCbcEnum::AES256CBC(key) => Aes256CbcDec::new(&(*key).into(), &iv.into())
.decrypt_padded_mut::<Pkcs7>(secret_body.en_body_mut()),
}; };
match rs { match rs {
Ok(buf) => { Ok(buf) => {
@@ -82,14 +92,19 @@ impl AesCbcCipher {
net_packet.set_data_len(HEAD_LEN + len - 4)?; net_packet.set_data_len(HEAD_LEN + len - 4)?;
Ok(()) Ok(())
} }
Err(e) => { Err(e) => Err(io::Error::new(
Err(io::Error::new(io::ErrorKind::Other, format!("解密失败:{}", e))) io::ErrorKind::Other,
} format!("解密失败:{}", e),
)),
} }
} }
/// net_packet 必须预留足够长度 /// net_packet 必须预留足够长度
/// data_len是有效载荷的长度 /// data_len是有效载荷的长度
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
let data_len = net_packet.data_len();
let mut iv = [0; 16]; let mut iv = [0; 16];
iv[0..4].copy_from_slice(&net_packet.source().octets()); iv[0..4].copy_from_slice(&net_packet.source().octets());
iv[4..8].copy_from_slice(&net_packet.destination().octets()); iv[4..8].copy_from_slice(&net_packet.destination().octets());
@@ -97,32 +112,44 @@ impl AesCbcCipher {
iv[9] = net_packet.transport_protocol(); iv[9] = net_packet.transport_protocol();
iv[10] = net_packet.is_gateway() as u8; iv[10] = net_packet.is_gateway() as u8;
iv[11] = net_packet.source_ttl(); iv[11] = net_packet.source_ttl();
iv[12..16].copy_from_slice(&self.finger.hash[0..4]); if let Some(finger) = &self.finger {
iv[12..16].copy_from_slice(&finger.hash[0..4]);
net_packet.set_data_len(data_len + 16)?;
} else {
net_packet.set_data_len(data_len + 4)?;
}
//先扩充随机数 //先扩充随机数
let data_len = net_packet.data_len(); let mut secret_body =
net_packet.set_data_len(data_len + 16)?; AesCbcSecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
let mut secret_body = AesCbcSecretBody::new(net_packet.payload_mut())?;
secret_body.set_random(rand::thread_rng().next_u32()); secret_body.set_random(rand::thread_rng().next_u32());
let p_len = secret_body.en_body().len(); let p_len = secret_body.en_body().len();
net_packet.set_data_len_max(); net_packet.set_data_len_max();
let rs = match &self.cipher { let rs = match &self.cipher {
AesCbcEnum::AES128CBC(key) => { Aes128CbcEnc::new(&(*key).into(), &iv.into()).encrypt_padded_mut::<Pkcs7>(net_packet.payload_mut(), p_len) } AesCbcEnum::AES128CBC(key) => Aes128CbcEnc::new(&(*key).into(), &iv.into())
AesCbcEnum::AES256CBC(key) => { Aes256CbcEnc::new(&(*key).into(), &iv.into()).encrypt_padded_mut::<Pkcs7>(net_packet.payload_mut(), p_len) } .encrypt_padded_mut::<Pkcs7>(net_packet.payload_mut(), p_len),
AesCbcEnum::AES256CBC(key) => Aes256CbcEnc::new(&(*key).into(), &iv.into())
.encrypt_padded_mut::<Pkcs7>(net_packet.payload_mut(), p_len),
}; };
return match rs { return match rs {
Ok(buf) => { Ok(buf) => {
let len = buf.len(); let len = buf.len();
let finger = self.finger.calculate_finger(&iv[..12], buf); if let Some(finger) = &self.finger {
//设置实际长度 let finger = finger.calculate_finger(&iv[..12], buf);
net_packet.set_data_len(HEAD_LEN + len + finger.len())?; //设置实际长度
let mut secret_body = AesCbcSecretBody::new(net_packet.payload_mut())?; net_packet.set_data_len(HEAD_LEN + len + finger.len())?;
secret_body.set_finger(&finger)?; let mut secret_body = AesCbcSecretBody::new(net_packet.payload_mut(), true)?;
secret_body.set_finger(&finger)?;
} else {
net_packet.set_data_len(HEAD_LEN + len)?;
}
net_packet.set_encrypt_flag(true); net_packet.set_encrypt_flag(true);
Ok(()) Ok(())
} }
Err(e) => { Err(e) => Err(io::Error::new(
Err(io::Error::new(io::ErrorKind::Other, format!("加密失败:{}", e))) io::ErrorKind::Other,
} format!("加密失败:{}", e),
)),
}; };
} }
} }
+154
View File
@@ -0,0 +1,154 @@
use crate::cipher::Finger;
use crate::protocol::body::AesCbcSecretBody;
use crate::protocol::{NetPacket, HEAD_LEN};
use aes::cipher::{block_padding::Pkcs7, BlockDecryptMut, BlockEncryptMut, KeyInit};
use rand::RngCore;
use std::io;
type Aes128EcbEnc = ecb::Encryptor<aes::Aes128>;
type Aes128EcbDec = ecb::Decryptor<aes::Aes128>;
type Aes256EcbEnc = ecb::Encryptor<aes::Aes256>;
type Aes256EcbDec = ecb::Decryptor<aes::Aes256>;
#[derive(Clone)]
pub struct AesEcbCipher {
pub(crate) cipher: AesEcbEnum,
pub(crate) finger: Option<Finger>,
}
#[derive(Clone)]
pub enum AesEcbEnum {
AES128ECB([u8; 16]),
AES256ECB([u8; 32]),
}
impl AesEcbCipher {
pub fn key(&self) -> &[u8] {
match &self.cipher {
AesEcbEnum::AES128ECB(key) => key,
AesEcbEnum::AES256ECB(key) => key,
}
}
}
impl AesEcbCipher {
pub fn new_128(key: [u8; 16], finger: Option<Finger>) -> Self {
Self {
cipher: AesEcbEnum::AES128ECB(key),
finger,
}
}
pub fn new_256(key: [u8; 32], finger: Option<Finger>) -> Self {
Self {
cipher: AesEcbEnum::AES256ECB(key),
finger,
}
}
pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
if !net_packet.is_encrypt() {
//未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
}
if net_packet.payload().len() < 16 {
log::error!("数据异常,长度{}小于{}", net_packet.payload().len(), 16);
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
let mut iv = [0; 16];
iv[0..4].copy_from_slice(&net_packet.source().octets());
iv[4..8].copy_from_slice(&net_packet.destination().octets());
iv[8] = net_packet.protocol().into();
iv[9] = net_packet.transport_protocol();
iv[10] = net_packet.is_gateway() as u8;
iv[11] = net_packet.source_ttl();
if let Some(finger) = &self.finger {
iv[12..16].copy_from_slice(&finger.hash[0..4]);
}
let mut secret_body =
AesCbcSecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
if let Some(finger) = &self.finger {
let finger = finger.calculate_finger(&iv[..12], secret_body.en_body());
if &finger != secret_body.finger() {
return Err(io::Error::new(io::ErrorKind::Other, "finger err"));
}
}
let rs = match &self.cipher {
AesEcbEnum::AES128ECB(key) => Aes128EcbDec::new(&(*key).into())
.decrypt_padded_mut::<Pkcs7>(secret_body.en_body_mut()),
AesEcbEnum::AES256ECB(key) => Aes256EcbDec::new(&(*key).into())
.decrypt_padded_mut::<Pkcs7>(secret_body.en_body_mut()),
};
match rs {
Ok(buf) => {
let len = buf.len();
net_packet.set_encrypt_flag(false);
//减去末尾的随机数
net_packet.set_data_len(HEAD_LEN + len - 4)?;
Ok(())
}
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("解密失败:{}", e),
)),
}
}
/// net_packet 必须预留足够长度
/// data_len是有效载荷的长度
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
let data_len = net_packet.data_len();
let mut iv = [0; 16];
iv[0..4].copy_from_slice(&net_packet.source().octets());
iv[4..8].copy_from_slice(&net_packet.destination().octets());
iv[8] = net_packet.protocol().into();
iv[9] = net_packet.transport_protocol();
iv[10] = net_packet.is_gateway() as u8;
iv[11] = net_packet.source_ttl();
if let Some(finger) = &self.finger {
iv[12..16].copy_from_slice(&finger.hash[0..4]);
net_packet.set_data_len(data_len + 16)?;
} else {
net_packet.set_data_len(data_len + 4)?;
}
//先扩充随机数
let mut secret_body =
AesCbcSecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
secret_body.set_random(rand::thread_rng().next_u32());
let p_len = secret_body.en_body().len();
net_packet.set_data_len_max();
let rs = match &self.cipher {
AesEcbEnum::AES128ECB(key) => Aes128EcbEnc::new(&(*key).into())
.encrypt_padded_mut::<Pkcs7>(net_packet.payload_mut(), p_len),
AesEcbEnum::AES256ECB(key) => Aes256EcbEnc::new(&(*key).into())
.encrypt_padded_mut::<Pkcs7>(net_packet.payload_mut(), p_len),
};
return match rs {
Ok(buf) => {
let len = buf.len();
if let Some(finger) = &self.finger {
let finger = finger.calculate_finger(&iv[..12], buf);
//设置实际长度
net_packet.set_data_len(HEAD_LEN + len + finger.len())?;
let mut secret_body = AesCbcSecretBody::new(net_packet.payload_mut(), true)?;
secret_body.set_finger(&finger)?;
} else {
net_packet.set_data_len(HEAD_LEN + len)?;
}
net_packet.set_encrypt_flag(true);
Ok(())
}
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("加密失败:{}", e),
)),
};
}
}
+46 -26
View File
@@ -1,19 +1,17 @@
use std::io; use std::io;
use aes_gcm::{AeadInPlace, Aes128Gcm, Aes256Gcm, Key, KeyInit, Nonce, Tag};
use aes_gcm::aead::consts::{U12, U16}; use aes_gcm::aead::consts::{U12, U16};
use aes_gcm::aead::generic_array::GenericArray; use aes_gcm::aead::generic_array::GenericArray;
use aes_gcm::{AeadInPlace, Aes128Gcm, Aes256Gcm, Key, KeyInit, Nonce, Tag};
use rand::RngCore; use rand::RngCore;
use crate::cipher::finger::Finger; use crate::cipher::finger::Finger;
use crate::protocol::{body::ENCRYPTION_RESERVED, body::SecretBody, NetPacket}; use crate::protocol::{body::SecretBody, body::ENCRYPTION_RESERVED, NetPacket};
#[derive(Clone)] #[derive(Clone)]
pub struct AesGcmCipher { pub struct AesGcmCipher {
pub(crate) cipher: AesGcmEnum, pub(crate) cipher: AesGcmEnum,
pub(crate) finger: Finger, pub(crate) finger: Option<Finger>,
} }
#[derive(Clone)] #[derive(Clone)]
@@ -23,14 +21,14 @@ pub enum AesGcmEnum {
} }
impl AesGcmCipher { impl AesGcmCipher {
pub fn new_128(key: [u8; 16], finger: Finger) -> Self { pub fn new_128(key: [u8; 16], finger: Option<Finger>) -> Self {
let key: &Key<Aes128Gcm> = &key.into(); let key: &Key<Aes128Gcm> = &key.into();
Self { Self {
cipher: AesGcmEnum::AES128GCM(Aes128Gcm::new(key)), cipher: AesGcmEnum::AES128GCM(Aes128Gcm::new(key)),
finger, finger,
} }
} }
pub fn new_256(key: [u8; 32], finger: Finger) -> Self { pub fn new_256(key: [u8; 32], finger: Option<Finger>) -> Self {
let key: &Key<Aes256Gcm> = &key.into(); let key: &Key<Aes256Gcm> = &key.into();
Self { Self {
cipher: AesGcmEnum::AES256GCM(Aes256Gcm::new(key)), cipher: AesGcmEnum::AES256GCM(Aes256Gcm::new(key)),
@@ -38,13 +36,16 @@ impl AesGcmCipher {
} }
} }
pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
if !net_packet.is_encrypt() { if !net_packet.is_encrypt() {
//未加密的数据直接丢弃 //未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
} }
if net_packet.payload().len() < ENCRYPTION_RESERVED { if net_packet.payload().len() < ENCRYPTION_RESERVED {
log::error!("数据异常,长度小于{}",ENCRYPTION_RESERVED); log::error!("数据异常,长度小于{}", ENCRYPTION_RESERVED);
return Err(io::Error::new(io::ErrorKind::Other, "data err")); return Err(io::Error::new(io::ErrorKind::Other, "data err"));
} }
let mut nonce_raw = [0; 12]; let mut nonce_raw = [0; 12];
@@ -56,19 +57,28 @@ impl AesGcmCipher {
nonce_raw[11] = net_packet.source_ttl(); nonce_raw[11] = net_packet.source_ttl();
let nonce: &GenericArray<u8, U12> = Nonce::from_slice(&nonce_raw); let nonce: &GenericArray<u8, U12> = Nonce::from_slice(&nonce_raw);
let mut secret_body = SecretBody::new(net_packet.payload_mut())?; let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
let tag = secret_body.tag(); let tag = secret_body.tag();
let finger = self.finger.calculate_finger(&nonce_raw, secret_body.en_body()); if let Some(finger) = &self.finger {
if &finger != secret_body.finger() { let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body());
return Err(io::Error::new(io::ErrorKind::Other, "finger err")); if &finger != secret_body.finger() {
return Err(io::Error::new(io::ErrorKind::Other, "finger err"));
}
} }
let tag: GenericArray<u8, U16> = Tag::clone_from_slice(tag); let tag: GenericArray<u8, U16> = Tag::clone_from_slice(tag);
let rs = match &self.cipher { let rs = match &self.cipher {
AesGcmEnum::AES128GCM(aes_gcm) => { aes_gcm.decrypt_in_place_detached(nonce, &[], secret_body.body_mut(), &tag) } AesGcmEnum::AES128GCM(aes_gcm) => {
AesGcmEnum::AES256GCM(aes_gcm) => { aes_gcm.decrypt_in_place_detached(nonce, &[], secret_body.body_mut(), &tag) } aes_gcm.decrypt_in_place_detached(nonce, &[], secret_body.body_mut(), &tag)
}
AesGcmEnum::AES256GCM(aes_gcm) => {
aes_gcm.decrypt_in_place_detached(nonce, &[], secret_body.body_mut(), &tag)
}
}; };
if let Err(e) = rs { if let Err(e) = rs {
return Err(io::Error::new(io::ErrorKind::Other, format!("解密失败:{}", e))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("解密失败:{}", e),
));
} }
net_packet.set_encrypt_flag(false); net_packet.set_encrypt_flag(false);
net_packet.set_data_len(net_packet.data_len() - ENCRYPTION_RESERVED)?; net_packet.set_data_len(net_packet.data_len() - ENCRYPTION_RESERVED)?;
@@ -76,7 +86,10 @@ impl AesGcmCipher {
} }
/// net_packet 必须预留足够长度 /// net_packet 必须预留足够长度
/// data_len是有效载荷的长度 /// data_len是有效载荷的长度
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
if net_packet.reserve() < ENCRYPTION_RESERVED { if net_packet.reserve() < ENCRYPTION_RESERVED {
return Err(io::Error::new(io::ErrorKind::Other, "too short")); return Err(io::Error::new(io::ErrorKind::Other, "too short"));
} }
@@ -90,23 +103,30 @@ impl AesGcmCipher {
let nonce: &GenericArray<u8, U12> = Nonce::from_slice(&nonce_raw); let nonce: &GenericArray<u8, U12> = Nonce::from_slice(&nonce_raw);
let data_len = net_packet.data_len() + ENCRYPTION_RESERVED; let data_len = net_packet.data_len() + ENCRYPTION_RESERVED;
net_packet.set_data_len(data_len)?; net_packet.set_data_len(data_len)?;
let mut secret_body = SecretBody::new(net_packet.payload_mut())?; let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
secret_body.set_random(rand::thread_rng().next_u32()); secret_body.set_random(rand::thread_rng().next_u32());
let rs = match &self.cipher { let rs = match &self.cipher {
AesGcmEnum::AES128GCM(aes_gcm) => { aes_gcm.encrypt_in_place_detached(nonce, &[], secret_body.body_mut()) } AesGcmEnum::AES128GCM(aes_gcm) => {
AesGcmEnum::AES256GCM(aes_gcm) => { aes_gcm.encrypt_in_place_detached(nonce, &[], secret_body.body_mut()) } aes_gcm.encrypt_in_place_detached(nonce, &[], secret_body.body_mut())
}
AesGcmEnum::AES256GCM(aes_gcm) => {
aes_gcm.encrypt_in_place_detached(nonce, &[], secret_body.body_mut())
}
}; };
return match rs { return match rs {
Ok(tag) => { Ok(tag) => {
secret_body.set_tag(tag.as_slice())?; secret_body.set_tag(tag.as_slice())?;
let finger = self.finger.calculate_finger(&nonce_raw, secret_body.en_body()); if let Some(finger) = &self.finger {
secret_body.set_finger(&finger)?; let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body());
secret_body.set_finger(&finger)?;
}
net_packet.set_encrypt_flag(true); net_packet.set_encrypt_flag(true);
Ok(()) Ok(())
} }
Err(e) => { Err(e) => Err(io::Error::new(
Err(io::Error::new(io::ErrorKind::Other, format!("加密失败:{}", e))) io::ErrorKind::Other,
} format!("加密失败:{}", e),
)),
}; };
} }
} }
+62 -60
View File
@@ -1,32 +1,31 @@
use std::io; use crate::cipher::aes_ecb::AesEcbCipher;
use std::str::FromStr;
use crate::cipher::{aes_cbc, Finger};
use crate::protocol::NetPacket;
use sha2::Digest;
#[cfg(feature = "ring-cipher")]
use crate::cipher::ring_aes_gcm_cipher::AesGcmCipher;
#[cfg(not(feature = "ring-cipher"))] #[cfg(not(feature = "ring-cipher"))]
use crate::cipher::aes_gcm_cipher::AesGcmCipher; use crate::cipher::aes_gcm_cipher::AesGcmCipher;
#[cfg(feature = "ring-cipher")]
use crate::cipher::ring_aes_gcm_cipher::AesGcmCipher;
use crate::cipher::{aes_cbc, Finger};
use crate::protocol::NetPacket;
use aes_cbc::AesCbcCipher; use aes_cbc::AesCbcCipher;
use sha2::Digest;
use std::io;
use std::str::FromStr;
#[derive(Copy, Clone, Eq, PartialEq, Debug)] #[derive(Copy, Clone, Eq, PartialEq, Debug)]
pub enum CipherModel { pub enum CipherModel {
AesGcm, AesGcm,
AesCbc, AesCbc,
AesEcb,
} }
impl FromStr for CipherModel { impl FromStr for CipherModel {
type Err = String; type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> { fn from_str(s: &str) -> Result<Self, Self::Err> {
match s { match s.to_lowercase().trim() {
"aes_gcm" => { "aes_gcm" => Ok(CipherModel::AesGcm),
Ok(CipherModel::AesGcm) "aes_cbc" => Ok(CipherModel::AesCbc),
} "aes_ecb" => Ok(CipherModel::AesEcb),
"aes_cbc" => { Ok(CipherModel::AesCbc) } _ => Err(format!("not match '{}'", s)),
_ => {
Err(format!("not match '{}'", s))
}
} }
} }
} }
@@ -35,12 +34,17 @@ impl FromStr for CipherModel {
pub enum Cipher { pub enum Cipher {
AesGcm((AesGcmCipher, Vec<u8>)), AesGcm((AesGcmCipher, Vec<u8>)),
AesCbc(AesCbcCipher), AesCbc(AesCbcCipher),
AesEcb(AesEcbCipher),
None, None,
} }
impl Cipher { impl Cipher {
pub fn new_password(model: CipherModel, password: Option<String>, token: String) -> Self { pub fn new_password(
let finger = Finger::new(&token); model: CipherModel,
password: Option<String>,
token: Option<String>,
) -> Self {
let finger = token.map(|token| Finger::new(&token));
if let Some(password) = password { if let Some(password) = password {
let mut hasher = sha2::Sha256::new(); let mut hasher = sha2::Sha256::new();
hasher.update(password.as_bytes()); hasher.update(password.as_bytes());
@@ -64,13 +68,22 @@ impl Cipher {
Cipher::AesCbc(aes) Cipher::AesCbc(aes)
} }
} }
CipherModel::AesEcb => {
if password.len() < 8 {
let aes = AesEcbCipher::new_128(key[..16].try_into().unwrap(), finger);
Cipher::AesEcb(aes)
} else {
let aes = AesEcbCipher::new_256(key, finger);
Cipher::AesEcb(aes)
}
}
} }
} else { } else {
Cipher::None Cipher::None
} }
} }
pub fn new_key(key: [u8; 32], token: String) -> io::Result<Self> { pub fn new_key(key: [u8; 32], token: String) -> io::Result<Self> {
let finger = Finger::new(&token); let finger = Some(Finger::new(&token));
match key.len() { match key.len() {
16 => { 16 => {
let aes = AesGcmCipher::new_128(key[..16].try_into().unwrap(), finger); let aes = AesGcmCipher::new_128(key[..16].try_into().unwrap(), finger);
@@ -80,20 +93,17 @@ impl Cipher {
let aes = AesGcmCipher::new_256(key, finger); let aes = AesGcmCipher::new_256(key, finger);
Ok(Cipher::AesGcm((aes, key.to_vec()))) Ok(Cipher::AesGcm((aes, key.to_vec())))
} }
_ => { _ => Err(io::Error::new(io::ErrorKind::Other, "key error")),
Err(io::Error::new(io::ErrorKind::Other, "key error"))
}
} }
} }
pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
match self { match self {
Cipher::AesGcm((aes_gcm, _)) => { Cipher::AesGcm((aes_gcm, _)) => aes_gcm.decrypt_ipv4(net_packet),
aes_gcm.decrypt_ipv4(net_packet) Cipher::AesCbc(aes_cbc) => aes_cbc.decrypt_ipv4(net_packet),
} Cipher::AesEcb(aes_ecb) => aes_ecb.decrypt_ipv4(net_packet),
Cipher::AesCbc(aes_cbc) => {
aes_cbc.decrypt_ipv4(net_packet).unwrap();
Ok(())
}
Cipher::None => { Cipher::None => {
if net_packet.is_encrypt() { if net_packet.is_encrypt() {
return Err(io::Error::new(io::ErrorKind::Other, "not key")); return Err(io::Error::new(io::ErrorKind::Other, "not key"));
@@ -102,44 +112,36 @@ impl Cipher {
} }
} }
} }
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
match self { match self {
Cipher::AesGcm((aes_gcm, _)) => { Cipher::AesGcm((aes_gcm, _)) => aes_gcm.encrypt_ipv4(net_packet),
aes_gcm.encrypt_ipv4(net_packet) Cipher::AesCbc(aes_cbc) => aes_cbc.encrypt_ipv4(net_packet),
} Cipher::AesEcb(aes_ecb) => aes_ecb.encrypt_ipv4(net_packet),
Cipher::AesCbc(aes_cbc) => { Cipher::None => Ok(()),
aes_cbc.encrypt_ipv4(net_packet).unwrap();
Ok(())
}
Cipher::None => {
Ok(())
}
} }
} }
pub fn check_finger<B: AsRef<[u8]>>(&self, net_packet: &NetPacket<B>) -> io::Result<()> { pub fn check_finger<B: AsRef<[u8]>>(&self, net_packet: &NetPacket<B>) -> io::Result<()> {
match self { let finger = match self {
Cipher::AesGcm((aes_gcm, _)) => { Cipher::AesGcm((aes_gcm, _)) => aes_gcm.finger.as_ref(),
aes_gcm.finger.check_finger(net_packet) Cipher::AesCbc(aes_cbc) => aes_cbc.finger.as_ref(),
} Cipher::AesEcb(aes_ecb) => aes_ecb.finger.as_ref(),
Cipher::AesCbc(aes_cbc) => { Cipher::None => None,
aes_cbc.finger.check_finger(net_packet) };
} if let Some(finger) = finger {
Cipher::None => { finger.check_finger(net_packet)
Ok(()) } else {
} Ok(())
} }
} }
pub fn key(&self) -> Option<&[u8]> { pub fn key(&self) -> Option<&[u8]> {
match self { match self {
Cipher::AesGcm((_, key)) => { Cipher::AesGcm((_, key)) => Some(key),
Some(key) Cipher::AesCbc(aes_cbc) => Some(aes_cbc.key()),
} Cipher::AesEcb(aes_ecb) => Some(aes_ecb.key()),
Cipher::AesCbc(aes_cbc) => { Cipher::None => None,
Some(aes_cbc.key())
}
Cipher::None => {
None
}
} }
} }
} }
+2 -2
View File
@@ -23,7 +23,7 @@ impl Finger {
} }
let payload_len = net_packet.payload().len(); let payload_len = net_packet.payload().len();
if payload_len < 12 { if payload_len < 12 {
log::error!("数据异常,长度小于{}",12); log::error!("数据异常,长度小于{}", 12);
return Err(io::Error::new(io::ErrorKind::Other, "data err")); return Err(io::Error::new(io::ErrorKind::Other, "data err"));
} }
let mut nonce_raw = [0; 12]; let mut nonce_raw = [0; 12];
@@ -48,4 +48,4 @@ impl Finger {
let key: [u8; 32] = hasher.finalize().into(); let key: [u8; 32] = hasher.finalize().into();
return key[20..].try_into().unwrap(); return key[20..].try_into().unwrap();
} }
} }
+7 -7
View File
@@ -1,14 +1,14 @@
#[cfg(feature = "ring-cipher")] mod aes_cbc;
mod ring_aes_gcm_cipher; mod aes_ecb;
#[cfg(not(feature = "ring-cipher"))] #[cfg(not(feature = "ring-cipher"))]
mod aes_gcm_cipher; mod aes_gcm_cipher;
mod rsa_cipher;
mod aes_cbc;
mod finger;
mod cipher; mod cipher;
mod finger;
#[cfg(feature = "ring-cipher")]
mod ring_aes_gcm_cipher;
mod rsa_cipher;
pub use cipher::Cipher; pub use cipher::Cipher;
pub use cipher::CipherModel;
pub use finger::Finger; pub use finger::Finger;
pub use rsa_cipher::RsaCipher; pub use rsa_cipher::RsaCipher;
pub use cipher::CipherModel;
+43 -25
View File
@@ -1,16 +1,16 @@
use std::io; use crate::cipher::Finger;
use rand::RngCore; use rand::RngCore;
use ring::aead; use ring::aead;
use ring::aead::{LessSafeKey, UnboundKey}; use ring::aead::{LessSafeKey, UnboundKey};
use crate::cipher::Finger; use std::io;
use crate::protocol::body::{SecretBody, ENCRYPTION_RESERVED};
use crate::protocol::NetPacket; use crate::protocol::NetPacket;
use crate::protocol::body::{ENCRYPTION_RESERVED, SecretBody};
#[derive(Clone)] #[derive(Clone)]
pub struct AesGcmCipher { pub struct AesGcmCipher {
pub(crate) cipher: AesGcmEnum, pub(crate) cipher: AesGcmEnum,
pub(crate) finger: Finger, pub(crate) finger: Option<Finger>,
} }
pub enum AesGcmEnum { pub enum AesGcmEnum {
@@ -22,11 +22,13 @@ impl Clone for AesGcmEnum {
fn clone(&self) -> Self { fn clone(&self) -> Self {
match &self { match &self {
AesGcmEnum::AesGCM128(_, key) => { AesGcmEnum::AesGCM128(_, key) => {
let c = LessSafeKey::new(UnboundKey::new(&aead::AES_128_GCM, key.as_slice()).unwrap()); let c =
LessSafeKey::new(UnboundKey::new(&aead::AES_128_GCM, key.as_slice()).unwrap());
AesGcmEnum::AesGCM128(c, *key) AesGcmEnum::AesGCM128(c, *key)
} }
AesGcmEnum::AesGCM256(_, key) => { AesGcmEnum::AesGCM256(_, key) => {
let c = LessSafeKey::new(UnboundKey::new(&aead::AES_256_GCM, key.as_slice()).unwrap()); let c =
LessSafeKey::new(UnboundKey::new(&aead::AES_256_GCM, key.as_slice()).unwrap());
AesGcmEnum::AesGCM256(c, *key) AesGcmEnum::AesGCM256(c, *key)
} }
} }
@@ -34,27 +36,30 @@ impl Clone for AesGcmEnum {
} }
impl AesGcmCipher { impl AesGcmCipher {
pub fn new_128(key: [u8; 16], finger: Finger) -> Self { pub fn new_128(key: [u8; 16], finger: Option<Finger>) -> Self {
let cipher = LessSafeKey::new(UnboundKey::new(&aead::AES_128_GCM, &key).unwrap()); let cipher = LessSafeKey::new(UnboundKey::new(&aead::AES_128_GCM, &key).unwrap());
Self { Self {
cipher: AesGcmEnum::AesGCM128(cipher, key), cipher: AesGcmEnum::AesGCM128(cipher, key),
finger, finger,
} }
} }
pub fn new_256(key: [u8; 32], finger: Finger) -> Self { pub fn new_256(key: [u8; 32], finger: Option<Finger>) -> Self {
let cipher = LessSafeKey::new(UnboundKey::new(&aead::AES_256_GCM, &key).unwrap()); let cipher = LessSafeKey::new(UnboundKey::new(&aead::AES_256_GCM, &key).unwrap());
Self { Self {
cipher: AesGcmEnum::AesGCM256(cipher, key), cipher: AesGcmEnum::AesGCM256(cipher, key),
finger, finger,
} }
} }
pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
if !net_packet.is_encrypt() { if !net_packet.is_encrypt() {
//未加密的数据直接丢弃 //未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
} }
if net_packet.payload().len() < ENCRYPTION_RESERVED { if net_packet.payload().len() < ENCRYPTION_RESERVED {
log::error!("数据异常,长度小于{}",ENCRYPTION_RESERVED); log::error!("数据异常,长度小于{}", ENCRYPTION_RESERVED);
return Err(io::Error::new(io::ErrorKind::Other, "data err")); return Err(io::Error::new(io::ErrorKind::Other, "data err"));
} }
let mut nonce_raw = [0; 12]; let mut nonce_raw = [0; 12];
@@ -65,11 +70,12 @@ impl AesGcmCipher {
nonce_raw[10] = net_packet.is_gateway() as u8; nonce_raw[10] = net_packet.is_gateway() as u8;
nonce_raw[11] = net_packet.source_ttl(); nonce_raw[11] = net_packet.source_ttl();
let nonce = aead::Nonce::assume_unique_for_key(nonce_raw); let nonce = aead::Nonce::assume_unique_for_key(nonce_raw);
let mut secret_body = SecretBody::new(net_packet.payload_mut())?; let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
let tag = secret_body.tag(); if let Some(finger) = &self.finger {
let finger = self.finger.calculate_finger(&nonce_raw, secret_body.en_body()); let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body());
if &finger != secret_body.finger() { if &finger != secret_body.finger() {
return Err(io::Error::new(io::ErrorKind::Other, "finger err")); return Err(io::Error::new(io::ErrorKind::Other, "ring aes finger err"));
}
} }
let rs = match &self.cipher { let rs = match &self.cipher {
@@ -81,7 +87,10 @@ impl AesGcmCipher {
} }
}; };
if let Err(e) = rs { if let Err(e) = rs {
return Err(io::Error::new(io::ErrorKind::Other, format!("解密失败:{}", e))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("解密失败:{}", e),
));
} }
net_packet.set_encrypt_flag(false); net_packet.set_encrypt_flag(false);
net_packet.set_data_len(net_packet.data_len() - ENCRYPTION_RESERVED)?; net_packet.set_data_len(net_packet.data_len() - ENCRYPTION_RESERVED)?;
@@ -90,7 +99,10 @@ impl AesGcmCipher {
/// net_packet 必须预留足够长度 /// net_packet 必须预留足够长度
/// data_len是有效载荷的长度 /// data_len是有效载荷的长度
/// 返回加密后载荷的长度 /// 返回加密后载荷的长度
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<()> { pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
let mut nonce_raw = [0; 12]; let mut nonce_raw = [0; 12];
nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); nonce_raw[0..4].copy_from_slice(&net_packet.source().octets());
nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets()); nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets());
@@ -101,7 +113,7 @@ impl AesGcmCipher {
let nonce = aead::Nonce::assume_unique_for_key(nonce_raw); let nonce = aead::Nonce::assume_unique_for_key(nonce_raw);
let data_len = net_packet.data_len() + ENCRYPTION_RESERVED; let data_len = net_packet.data_len() + ENCRYPTION_RESERVED;
net_packet.set_data_len(data_len)?; net_packet.set_data_len(data_len)?;
let mut secret_body = SecretBody::new(net_packet.payload_mut())?; let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?;
secret_body.set_random(rand::thread_rng().next_u32()); secret_body.set_random(rand::thread_rng().next_u32());
let rs = match &self.cipher { let rs = match &self.cipher {
@@ -116,17 +128,23 @@ impl AesGcmCipher {
Ok(tag) => { Ok(tag) => {
let tag = tag.as_ref(); let tag = tag.as_ref();
if tag.len() != 16 { if tag.len() != 16 {
return Err(io::Error::new(io::ErrorKind::Other, format!("加密tag长度错误:{}", tag.len()))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("加密tag长度错误:{}", tag.len()),
));
} }
secret_body.set_tag(tag)?; secret_body.set_tag(tag)?;
let finger = self.finger.calculate_finger(&nonce_raw, secret_body.en_body()); if let Some(finger) = &self.finger {
secret_body.set_finger(&finger)?; let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body());
secret_body.set_finger(&finger)?;
}
net_packet.set_encrypt_flag(true); net_packet.set_encrypt_flag(true);
Ok(()) Ok(())
} }
Err(e) => { Err(e) => Err(io::Error::new(
Err(io::Error::new(io::ErrorKind::Other, format!("加密失败:{}", e))) io::ErrorKind::Other,
} format!("加密失败:{}", e),
)),
}; };
} }
} }
+41 -39
View File
@@ -1,11 +1,11 @@
use std::io; use crate::protocol::body::{RsaSecretBody, ENCRYPTION_RESERVED};
use crate::protocol::NetPacket;
use rand::Rng; use rand::Rng;
use rsa::pkcs8::der::Decode; use rsa::pkcs8::der::Decode;
use rsa::{PublicKey, RsaPublicKey}; use rsa::{PublicKey, RsaPublicKey};
use spki::{DecodePublicKey, EncodePublicKey};
use crate::protocol::body::{ENCRYPTION_RESERVED, RsaSecretBody};
use crate::protocol::NetPacket;
use sha2::Digest; use sha2::Digest;
use spki::{DecodePublicKey, EncodePublicKey};
use std::io;
#[derive(Clone)] #[derive(Clone)]
pub struct RsaCipher { pub struct RsaCipher {
@@ -21,47 +21,44 @@ impl RsaCipher {
pub fn new(der: &[u8]) -> io::Result<Self> { pub fn new(der: &[u8]) -> io::Result<Self> {
match RsaPublicKey::from_public_key_der(der) { match RsaPublicKey::from_public_key_der(der) {
Ok(public_key) => { Ok(public_key) => {
let inner = Inner { let inner = Inner { public_key };
public_key, Ok(Self { inner })
};
Ok(Self {
inner
})
}
Err(e) => {
Err(io::Error::new(io::ErrorKind::Other, format!("from_public_key_der failed {}", e)))
} }
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("from_public_key_der failed {}", e),
)),
} }
} }
pub fn finger(&self) -> io::Result<String> { pub fn finger(&self) -> io::Result<String> {
match self.inner.public_key.to_public_key_der() { match self.inner.public_key.to_public_key_der() {
Ok(der) => { Ok(der) => match rsa::pkcs8::SubjectPublicKeyInfo::from_der(der.as_bytes()) {
match rsa::pkcs8::SubjectPublicKeyInfo::from_der(der.as_bytes()) { Ok(spki) => match spki.fingerprint_base64() {
Ok(spki) => { Ok(finger) => Ok(finger),
match spki.fingerprint_base64() { Err(e) => Err(io::Error::new(
Ok(finger) => { io::ErrorKind::Other,
Ok(finger) format!("fingerprint_base64 error {}", e),
} )),
Err(e) => { },
Err(io::Error::new(io::ErrorKind::Other, format!("fingerprint_base64 error {}", e))) Err(e) => Err(io::Error::new(
} io::ErrorKind::Other,
} format!("from_der error {}", e),
} )),
Err(e) => { },
Err(io::Error::new(io::ErrorKind::Other, format!("from_der error {}", e))) Err(e) => Err(io::Error::new(
} io::ErrorKind::Other,
} format!("to_public_key_der error {}", e),
} )),
Err(e) => {
Err(io::Error::new(io::ErrorKind::Other, format!("to_public_key_der error {}", e)))
}
} }
} }
} }
impl RsaCipher { impl RsaCipher {
/// net_packet 必须预留足够长度 /// net_packet 必须预留足够长度
pub fn encrypt<B: AsRef<[u8]> + AsMut<[u8]>>(&self, net_packet: &mut NetPacket<B>) -> io::Result<NetPacket<Vec<u8>>> { pub fn encrypt<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<NetPacket<Vec<u8>>> {
if net_packet.reserve() < ENCRYPTION_RESERVED { if net_packet.reserve() < ENCRYPTION_RESERVED {
return Err(io::Error::new(io::ErrorKind::Other, "too short")); return Err(io::Error::new(io::ErrorKind::Other, "too short"));
} }
@@ -84,16 +81,21 @@ impl RsaCipher {
hasher.update(nonce_raw); hasher.update(nonce_raw);
let key: [u8; 32] = hasher.finalize().into(); let key: [u8; 32] = hasher.finalize().into();
secret_body.set_finger(&key[16..])?; secret_body.set_finger(&key[16..])?;
match self.inner.public_key.encrypt(&mut rng, rsa::PaddingScheme::PKCS1v15Encrypt, secret_body.buffer()) { match self.inner.public_key.encrypt(
&mut rng,
rsa::PaddingScheme::PKCS1v15Encrypt,
secret_body.buffer(),
) {
Ok(enc_data) => { Ok(enc_data) => {
let mut net_packet_e = NetPacket::new(vec![0; 12 + enc_data.len()])?; let mut net_packet_e = NetPacket::new(vec![0; 12 + enc_data.len()])?;
net_packet_e.buffer_mut()[..12].copy_from_slice(&net_packet.buffer()[..12]); net_packet_e.buffer_mut()[..12].copy_from_slice(&net_packet.buffer()[..12]);
net_packet_e.set_payload(&enc_data)?; net_packet_e.set_payload(&enc_data)?;
Ok(net_packet_e) Ok(net_packet_e)
} }
Err(e) => { Err(e) => Err(io::Error::new(
Err(io::Error::new(io::ErrorKind::Other, format!("encrypt failed {}", e))) io::ErrorKind::Other,
} format!("encrypt failed {}", e),
)),
} }
} }
} }
+265 -91
View File
@@ -10,22 +10,25 @@ use rand::Rng;
use tokio::net::{TcpStream, UdpSocket}; use tokio::net::{TcpStream, UdpSocket};
use tokio::sync::mpsc::channel; use tokio::sync::mpsc::channel;
use crate::channel::{Route, RouteKey};
use crate::channel::channel::{Channel, Context}; use crate::channel::channel::{Channel, Context};
use crate::channel::idle::Idle; use crate::channel::idle::Idle;
use crate::channel::punch::{NatInfo, Punch}; use crate::channel::punch::{NatInfo, Punch, PunchModel};
use crate::channel::sender::ChannelSender; use crate::channel::sender::ChannelSender;
use crate::channel::{Route, RouteKey};
use crate::cipher::{Cipher, CipherModel, RsaCipher}; use crate::cipher::{Cipher, CipherModel, RsaCipher};
use crate::core::status::VntStatusManger; use crate::core::status::VntStatusManger;
use crate::error::Error; use crate::error::Error;
use crate::external_route::{AllowExternalRoute, ExternalRoute}; use crate::external_route::{AllowExternalRoute, ExternalRoute};
use crate::handle::{ConnectStatus, CurrentDeviceInfo, handshake_handler, heartbeat_handler, PeerDeviceInfo, punch_handler, registration_handler};
use crate::handle::handshake_handler::HandshakeEnum; use crate::handle::handshake_handler::HandshakeEnum;
use crate::handle::recv_handler::ChannelDataHandler; use crate::handle::recv_handler::ChannelDataHandler;
use crate::handle::registration_handler::{RegResponse, ReqEnum}; use crate::handle::registration_handler::{RegResponse, ReqEnum};
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
use crate::handle::tun_tap::tap_handler; use crate::handle::tun_tap::tap_handler;
use crate::handle::tun_tap::tun_handler; use crate::handle::tun_tap::tun_handler;
use crate::handle::{
handshake_handler, heartbeat_handler, punch_handler, registration_handler, ConnectStatus,
CurrentDeviceInfo, PeerDeviceInfo,
};
use crate::igmp_server::IgmpServer; use crate::igmp_server::IgmpServer;
use crate::nat::NatTest; use crate::nat::NatTest;
use crate::tun_tap_device; use crate::tun_tap_device;
@@ -34,7 +37,6 @@ use crate::tun_tap_device::{DeviceReader, DeviceWriter};
pub mod status; pub mod status;
pub mod sync; pub mod sync;
#[derive(Clone)] #[derive(Clone)]
pub struct Vnt { pub struct Vnt {
config: Config, config: Config,
@@ -54,6 +56,7 @@ pub struct Vnt {
pub struct VntUtil { pub struct VntUtil {
config: Config, config: Config,
main_channel: UdpSocket, main_channel: UdpSocket,
main_channel_ipv6: Option<UdpSocket>,
main_tcp_channel: Option<TcpStream>, main_tcp_channel: Option<TcpStream>,
response: Option<RegResponse>, response: Option<RegResponse>,
iface: Option<(DeviceWriter, DeviceReader)>, iface: Option<(DeviceWriter, DeviceReader)>,
@@ -64,6 +67,13 @@ pub struct VntUtil {
impl VntUtil { impl VntUtil {
pub async fn new(config: Config) -> io::Result<VntUtil> { pub async fn new(config: Config) -> io::Result<VntUtil> {
let main_channel = UdpSocket::bind("0.0.0.0:0").await?; let main_channel = UdpSocket::bind("0.0.0.0:0").await?;
let main_channel_ipv6 = match UdpSocket::bind("[::]:0").await {
Ok(main_channel_ipv6) => Some(main_channel_ipv6),
Err(e) => {
log::warn!("绑定ipv6地址失败:{}", e);
None
}
};
let server_cipher = if config.server_encrypt { let server_cipher = if config.server_encrypt {
let mut key = [0 as u8; 32]; let mut key = [0 as u8; 32];
rand::thread_rng().fill(&mut key); rand::thread_rng().fill(&mut key);
@@ -74,6 +84,7 @@ impl VntUtil {
Ok(VntUtil { Ok(VntUtil {
config, config,
main_channel, main_channel,
main_channel_ipv6,
main_tcp_channel: None, main_tcp_channel: None,
response: None, response: None,
iface: None, iface: None,
@@ -92,26 +103,48 @@ impl VntUtil {
///握手 用于获取公钥 ///握手 用于获取公钥
pub async fn handshake(&mut self) -> Result<Option<RsaCipher>, HandshakeEnum> { pub async fn handshake(&mut self) -> Result<Option<RsaCipher>, HandshakeEnum> {
let rsa_cipher = handshake_handler::handshake(&self.main_channel, self.main_tcp_channel.as_mut(), self.config.server_address, self.config.server_encrypt).await?; let rsa_cipher = handshake_handler::handshake(
&self.main_channel,
self.main_tcp_channel.as_mut(),
self.config.server_address,
self.config.server_encrypt,
)
.await?;
self.rsa_cipher = rsa_cipher.clone(); self.rsa_cipher = rsa_cipher.clone();
Ok(rsa_cipher) Ok(rsa_cipher)
} }
/// 加密握手 用于同步密钥 /// 加密握手 用于同步密钥
pub async fn secret_handshake(&mut self) -> Result<(), HandshakeEnum> { pub async fn secret_handshake(&mut self) -> Result<(), HandshakeEnum> {
handshake_handler::secret_handshake(&self.main_channel, self.main_tcp_channel.as_mut(), self.config.server_address, self.rsa_cipher.as_ref().unwrap(), &self.server_cipher, self.config.token.clone()).await handshake_handler::secret_handshake(
&self.main_channel,
self.main_tcp_channel.as_mut(),
self.config.server_address,
self.rsa_cipher.as_ref().unwrap(),
&self.server_cipher,
self.config.token.clone(),
)
.await
} }
/// 注册 /// 注册
pub async fn register(&mut self) -> Result<RegResponse, ReqEnum> { pub async fn register(&mut self) -> Result<RegResponse, ReqEnum> {
match registration_handler::registration(&self.main_channel, self.main_tcp_channel.as_mut(), &self.server_cipher, self.config.server_address, match registration_handler::registration(
self.config.token.clone(), self.config.device_id.clone(), &self.main_channel,
self.config.name.clone(), self.config.ip.unwrap_or(Ipv4Addr::UNSPECIFIED), self.config.password.is_some()).await { self.main_tcp_channel.as_mut(),
&self.server_cipher,
self.config.server_address,
self.config.token.clone(),
self.config.device_id.clone(),
self.config.name.clone(),
self.config.ip.unwrap_or(Ipv4Addr::UNSPECIFIED),
self.config.password.is_some(),
)
.await
{
Ok(res) => { Ok(res) => {
let _ = self.response.insert(res.clone()); let _ = self.response.insert(res.clone());
Ok(res) Ok(res)
} }
Err(e) => { Err(e) => Err(e),
Err(e)
}
} }
} }
#[cfg(any(target_os = "android"))] #[cfg(any(target_os = "android"))]
@@ -128,19 +161,15 @@ impl VntUtil {
None => { None => {
return Err(io::Error::from(io::ErrorKind::AlreadyExists)); return Err(io::Error::from(io::ErrorKind::AlreadyExists));
} }
Some(res) => { Some(res) => res,
res
}
}; };
let device_type = if self.config.tap { let device_type = if self.config.tap {
#[cfg(windows)]
{ {
//删除tun网卡避免ip冲突,因为非正常退出会保留网卡 //删除tun网卡避免ip冲突,因为非正常退出会保留网卡
tun_tap_device::delete_device(tun_tap_device::DeviceType::Tun); tun_tap_device::delete_device(tun_tap_device::DeviceType::Tun);
} }
tun_tap_device::DeviceType::Tap tun_tap_device::DeviceType::Tap
} else { } else {
#[cfg(windows)]
{ {
//删除tap网卡避免ip冲突,非正常退出会保留网卡 //删除tap网卡避免ip冲突,非正常退出会保留网卡
tun_tap_device::delete_device(tun_tap_device::DeviceType::Tap); tun_tap_device::delete_device(tun_tap_device::DeviceType::Tap);
@@ -150,19 +179,28 @@ impl VntUtil {
let mtu = match self.config.mtu { let mtu = match self.config.mtu {
None => { None => {
if self.config.password.is_none() { if self.config.password.is_none() {
1430 1450
} else { } else {
1410 1420
} }
} }
Some(mtu) => { Some(mtu) => mtu,
mtu
}
}; };
let in_ips = self.config.in_ips.iter().map(|(dest, mask, _)| { (Ipv4Addr::from(*dest & *mask), Ipv4Addr::from(*mask)) }).collect::<Vec<(Ipv4Addr, Ipv4Addr)>>(); let in_ips = self
.config
.in_ips
.iter()
.map(|(dest, mask, _)| (Ipv4Addr::from(*dest & *mask), Ipv4Addr::from(*mask)))
.collect::<Vec<(Ipv4Addr, Ipv4Addr)>>();
let (device_writer, device_reader, driver_info) = tun_tap_device::create_device(device_type, response.virtual_ip, let (device_writer, device_reader, driver_info) = tun_tap_device::create_device(
response.virtual_netmask, response.virtual_gateway, in_ips, mtu)?; device_type,
response.virtual_ip,
response.virtual_netmask,
response.virtual_gateway,
in_ips,
mtu,
)?;
let _ = self.iface.insert((device_writer, device_reader)); let _ = self.iface.insert((device_writer, device_reader));
Ok(driver_info) Ok(driver_info)
} }
@@ -171,25 +209,32 @@ impl VntUtil {
None => { None => {
return Err(Error::Stop("response None".to_string())); return Err(Error::Stop("response None".to_string()));
} }
Some(res) => { Some(res) => res,
res
}
}; };
let (device_writer, device_reader) = match self.iface { let (device_writer, device_reader) = match self.iface {
None => { None => {
return Err(Error::Stop("iface None".to_string())); return Err(Error::Stop("iface None".to_string()));
} }
Some(res) => { Some(res) => res,
res
}
}; };
let config = self.config.clone(); let config = self.config.clone();
let vnt_status_manager = VntStatusManger::new(); let vnt_status_manager = VntStatusManger::new();
let client_cipher = Cipher::new_password(config.cipher_model, config.password.clone(), config.token.clone()); let finger = if config.finger {
Some(config.token.clone())
} else {
None
};
let client_cipher =
Cipher::new_password(config.cipher_model, config.password.clone(), finger);
let virtual_ip = response.virtual_ip; let virtual_ip = response.virtual_ip;
let virtual_gateway = response.virtual_gateway; let virtual_gateway = response.virtual_gateway;
let virtual_netmask = response.virtual_netmask; let virtual_netmask = response.virtual_netmask;
let current_device = Arc::new(AtomicCell::new(CurrentDeviceInfo::new(virtual_ip, virtual_gateway, virtual_netmask, config.server_address))); let current_device = Arc::new(AtomicCell::new(CurrentDeviceInfo::new(
virtual_ip,
virtual_gateway,
virtual_netmask,
config.server_address,
)));
let (cone_sender, cone_receiver) = channel(3); let (cone_sender, cone_receiver) = channel(3);
let (symmetric_sender, symmetric_receiver) = channel(2); let (symmetric_sender, symmetric_receiver) = channel(2);
@@ -199,23 +244,45 @@ impl VntUtil {
} else { } else {
(None, None) (None, None)
}; };
let context = Context::new(Arc::new(self.main_channel), tcp_sender, current_device.clone(), 1); let context = Context::new(
let punch = Punch::new(context.clone()); Arc::new(self.main_channel),
self.main_channel_ipv6.map(|v| Arc::new(v)),
tcp_sender,
current_device.clone(),
1,
);
let punch = Punch::new(context.clone(), config.punch_model);
let idle = Idle::new(Duration::from_secs(16), context.clone()); let idle = Idle::new(Duration::from_secs(16), context.clone());
let channel_sender = ChannelSender::new(context.clone()); let channel_sender = ChannelSender::new(context.clone());
let register = Arc::new(registration_handler::Register::new(self.server_cipher.clone(), channel_sender.clone(), let register = Arc::new(registration_handler::Register::new(
config.server_address, config.token.clone(), self.server_cipher.clone(),
config.device_id.clone(), config.name.clone(), config.password.is_some())); channel_sender.clone(),
let device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>> = Arc::new(Mutex::new((response.epoch, response.device_info_list))); config.server_address,
config.token.clone(),
config.device_id.clone(),
config.name.clone(),
config.password.is_some(),
));
let device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>> =
Arc::new(Mutex::new((response.epoch, response.device_info_list)));
let peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>> = Arc::new(DashMap::new()); let peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>> = Arc::new(DashMap::new());
let connect_status = Arc::new(AtomicCell::new(ConnectStatus::Connected)); let connect_status = Arc::new(AtomicCell::new(ConnectStatus::Connected));
let local_port = context.main_local_ipv4_port().unwrap_or(0);
let local_ip = crate::nat::local_ip()?; let local_ipv4_addr = crate::nat::local_ipv4_addr(local_port);
let local_port = context.main_local_port()?; let ipv6_port = context.main_local_ipv6_port().unwrap_or(0);
let ipv6_addr = crate::nat::local_ipv6_addr(ipv6_port);
// NAT检测 // NAT检测
let nat_test = NatTest::new(config.stun_server.clone(), response.public_ip, response.public_port, local_ip, local_port).await; let nat_test = NatTest::new(
config.stun_server.clone(),
response.public_ip,
response.public_port,
local_ipv4_addr,
ipv6_addr,
)
.await;
let in_external_route = if config.in_ips.is_empty() { let in_external_route = if config.in_ips.is_empty() {
None None
} else { } else {
@@ -224,7 +291,12 @@ impl VntUtil {
let (tcp_proxy, udp_proxy, ip_proxy_map) = if config.out_ips.is_empty() { let (tcp_proxy, udp_proxy, ip_proxy_map) = if config.out_ips.is_empty() {
(None, None, None) (None, None, None)
} else { } else {
let (tcp_proxy, udp_proxy, ip_proxy_map) = crate::ip_proxy::init_proxy(channel_sender.clone(), current_device.clone(), client_cipher.clone()).await?; let (tcp_proxy, udp_proxy, ip_proxy_map) = crate::ip_proxy::init_proxy(
channel_sender.clone(),
current_device.clone(),
client_cipher.clone(),
)
.await?;
(Some(tcp_proxy), Some(udp_proxy), Some(ip_proxy_map)) (Some(tcp_proxy), Some(udp_proxy), Some(ip_proxy_map))
}; };
let out_external_route = AllowExternalRoute::new(config.out_ips); let out_external_route = AllowExternalRoute::new(config.out_ips);
@@ -236,58 +308,143 @@ impl VntUtil {
}; };
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
if config.tap { if config.tap {
tap_handler::start(vnt_status_manager.worker("tap_handler"), channel_sender.clone(), device_reader, device_writer.clone(), tap_handler::start(
igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), vnt_status_manager.worker("tap_handler"),
client_cipher.clone(), self.server_cipher.clone(), config.parallel); channel_sender.clone(),
device_reader,
device_writer.clone(),
igmp_server.clone(),
current_device.clone(),
in_external_route,
ip_proxy_map.clone(),
client_cipher.clone(),
self.server_cipher.clone(),
config.parallel,
);
} else { } else {
tun_handler::start(vnt_status_manager.worker("tun_handler"), channel_sender.clone(), device_reader, device_writer.clone(), tun_handler::start(
igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), vnt_status_manager.worker("tun_handler"),
client_cipher.clone(), self.server_cipher.clone(), config.parallel).await; channel_sender.clone(),
device_reader,
device_writer.clone(),
igmp_server.clone(),
current_device.clone(),
in_external_route,
ip_proxy_map.clone(),
client_cipher.clone(),
self.server_cipher.clone(),
config.parallel,
)
.await;
} }
#[cfg(any(target_os = "android"))] #[cfg(any(target_os = "android"))]
tun_handler::start(vnt_status_manager.worker("android tun_handler"), channel_sender.clone(), device_reader, device_writer.clone(), tun_handler::start(
igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), cipher.clone(), config.parallel).await; vnt_status_manager.worker("android tun_handler"),
channel_sender.clone(),
device_reader,
device_writer.clone(),
igmp_server.clone(),
current_device.clone(),
in_external_route,
ip_proxy_map.clone(),
client_cipher.clone(),
self.server_cipher.clone(),
config.parallel,
)
.await;
//外部数据接收处理 //外部数据接收处理
let channel_recv_handler = ChannelDataHandler::new(current_device.clone(), device_list.clone(), let channel_recv_handler = ChannelDataHandler::new(
register.clone(), nat_test.clone(), igmp_server, current_device.clone(),
device_writer.clone(), connect_status.clone(), device_list.clone(),
peer_nat_info_map.clone(), ip_proxy_map, out_external_route, register.clone(),
cone_sender, symmetric_sender, client_cipher.clone(), nat_test.clone(),
self.server_cipher.clone(), self.rsa_cipher.clone(), config.relay, config.token.clone()); igmp_server,
device_writer.clone(),
connect_status.clone(),
peer_nat_info_map.clone(),
ip_proxy_map,
out_external_route,
cone_sender,
symmetric_sender,
client_cipher.clone(),
self.server_cipher.clone(),
self.rsa_cipher.clone(),
config.relay,
config.token.clone(),
);
{ {
let channel = Channel::new(context.clone(), channel_recv_handler); let channel = Channel::new(context.clone(), channel_recv_handler);
let channel_worker = vnt_status_manager.worker("channel_worker"); let channel_worker = vnt_status_manager.worker("channel_worker");
let relay = config.relay; let relay = config.relay;
tokio::spawn(async move {
channel
.start(channel_worker, tcp, 14, 65, relay, config.parallel)
.await
});
}
{
let nat_test = nat_test.clone();
let device_list = device_list.clone();
let current_device = current_device.clone();
// 定时心跳
heartbeat_handler::start_heartbeat(
vnt_status_manager.worker("heartbeat"),
channel_sender.clone(),
device_list.clone(),
current_device.clone(),
config.server_address_str,
client_cipher.clone(),
self.server_cipher.clone(),
);
// 空闲检查
heartbeat_handler::start_idle(
vnt_status_manager.worker("idle"),
idle,
channel_sender.clone(),
);
if !config.relay {
// 打洞处理
punch_handler::start(
vnt_status_manager.worker("cone_receiver"),
cone_receiver,
punch.clone(),
current_device.clone(),
client_cipher.clone(),
);
punch_handler::start(
vnt_status_manager.worker("symmetric_receiver"),
symmetric_receiver,
punch,
current_device.clone(),
client_cipher.clone(),
);
tokio::spawn(punch_handler::start_punch(
vnt_status_manager.worker("punch_handler"),
nat_test,
device_list,
channel_sender,
current_device,
client_cipher.clone(),
));
}
}
{
//代理
if let Some(tcp_proxy) = tcp_proxy { if let Some(tcp_proxy) = tcp_proxy {
tokio::spawn(tcp_proxy.start()); tokio::spawn(tcp_proxy.start());
} }
if let Some(udp_proxy) = udp_proxy { if let Some(udp_proxy) = udp_proxy {
tokio::spawn(udp_proxy.start()); tokio::spawn(udp_proxy.start());
} }
let context = context.clone();
let nat_test = nat_test.clone();
//延迟切换类型,避免无效流量
tokio::spawn(async move { tokio::spawn(async move {
channel.start(channel_worker, tcp, 14, 65, relay, config.parallel).await tokio::time::sleep(Duration::from_secs(15)).await;
context.switch(nat_test.nat_info().nat_type);
}); });
} }
{
let other_worker = vnt_status_manager.worker("punch_handler");
let nat_test = nat_test.clone();
let device_list = device_list.clone();
let current_device = current_device.clone();
// 定时心跳
heartbeat_handler::start_heartbeat(other_worker.worker("heartbeat"), channel_sender.clone(), device_list.clone(),
current_device.clone(), config.server_address_str, client_cipher.clone(), self.server_cipher.clone());
// 空闲检查
heartbeat_handler::start_idle(other_worker.worker("idle"), idle, channel_sender.clone());
if !config.relay {
// 打洞处理
punch_handler::start(other_worker.worker("cone_receiver"), cone_receiver, punch.clone(), current_device.clone(), client_cipher.clone());
punch_handler::start(other_worker.worker("symmetric_receiver"), symmetric_receiver, punch, current_device.clone(), client_cipher.clone());
tokio::spawn(punch_handler::start_punch(other_worker, nat_test,
device_list, channel_sender, current_device, client_cipher.clone()));
}
}
context.switch(nat_test.nat_info().nat_type);
Ok(Vnt { Ok(Vnt {
config: self.config, config: self.config,
current_device, current_device,
@@ -344,8 +501,10 @@ impl Vnt {
self.vnt_status_manager.stop_all(); self.vnt_status_manager.stop_all();
self.device_writer.close()?; self.device_writer.close()?;
let virtual_gateway = self.current_device.load().virtual_gateway; let virtual_gateway = self.current_device.load().virtual_gateway;
let _ = std::net::UdpSocket::bind("0.0.0.0:0")?.send_to(&[0], let _ = std::net::UdpSocket::bind("0.0.0.0:0")?.send_to(
SocketAddr::V4(SocketAddrV4::new(virtual_gateway, 10000))); &[0],
SocketAddr::V4(SocketAddrV4::new(virtual_gateway, 10000)),
);
Ok(()) Ok(())
} }
pub async fn wait_stop(&mut self) { pub async fn wait_stop(&mut self) {
@@ -391,20 +550,33 @@ pub struct Config {
pub server_encrypt: bool, pub server_encrypt: bool,
pub parallel: usize, pub parallel: usize,
pub cipher_model: CipherModel, pub cipher_model: CipherModel,
pub finger: bool,
pub punch_model: PunchModel,
} }
impl Config { impl Config {
pub fn new(tap: bool, token: String, pub fn new(
device_id: String, tap: bool,
name: String, token: String,
server_address: SocketAddr, device_id: String,
server_address_str: String, name: String,
mut stun_server: Vec<String>, server_address: SocketAddr,
in_ips: Vec<(u32, u32, Ipv4Addr)>, out_ips: Vec<(u32, u32)>, server_address_str: String,
password: Option<String>, simulate_multicast: bool, mtu: Option<u16>, tcp: bool, mut stun_server: Vec<String>,
ip: Option<Ipv4Addr>, in_ips: Vec<(u32, u32, Ipv4Addr)>,
relay: bool, server_encrypt: bool, parallel: usize, cipher_model: CipherModel) -> Self { out_ips: Vec<(u32, u32)>,
password: Option<String>,
simulate_multicast: bool,
mtu: Option<u16>,
tcp: bool,
ip: Option<Ipv4Addr>,
relay: bool,
server_encrypt: bool,
parallel: usize,
cipher_model: CipherModel,
finger: bool,
punch_model: PunchModel,
) -> Self {
for x in stun_server.iter_mut() { for x in stun_server.iter_mut() {
if !x.contains(":") { if !x.contains(":") {
x.push_str(":3478"); x.push_str(":3478");
@@ -429,6 +601,8 @@ impl Config {
server_encrypt, server_encrypt,
parallel, parallel,
cipher_model, cipher_model,
finger,
punch_model,
} }
} }
} }
+8 -5
View File
@@ -1,7 +1,7 @@
use crate::util::wait::WaitGroup;
use std::sync::Arc; use std::sync::Arc;
use tokio::sync::watch; use tokio::sync::watch;
use tokio::sync::watch::{Receiver, Sender}; use tokio::sync::watch::{Receiver, Sender};
use crate::util::wait::WaitGroup;
#[derive(Copy, Clone, Eq, PartialEq)] #[derive(Copy, Clone, Eq, PartialEq)]
pub enum VntStatus { pub enum VntStatus {
@@ -10,7 +10,7 @@ pub enum VntStatus {
} }
pub struct VntWorker { pub struct VntWorker {
_name: String, name: String,
wg: WaitGroup, wg: WaitGroup,
status_s: Arc<Sender<VntStatus>>, status_s: Arc<Sender<VntStatus>>,
status_r: Receiver<VntStatus>, status_r: Receiver<VntStatus>,
@@ -20,7 +20,7 @@ impl VntWorker {
pub fn worker(&self, name: &str) -> Self { pub fn worker(&self, name: &str) -> Self {
self.wg.add(); self.wg.add();
VntWorker { VntWorker {
_name: name.to_string(), name: name.to_string(),
wg: self.wg.clone(), wg: self.wg.clone(),
status_s: self.status_s.clone(), status_s: self.status_s.clone(),
status_r: self.status_r.clone(), status_r: self.status_r.clone(),
@@ -30,6 +30,7 @@ impl VntWorker {
impl Drop for VntWorker { impl Drop for VntWorker {
fn drop(&mut self) { fn drop(&mut self) {
log::info!("任务停止:{}", self.name);
self.wg.done(); self.wg.done();
} }
} }
@@ -49,7 +50,9 @@ impl VntWorker {
return; return;
} }
} }
Err(_) => { return; } Err(_) => {
return;
}
} }
} }
} }
@@ -80,7 +83,7 @@ impl VntStatusManger {
pub fn worker(&self, name: &str) -> VntWorker { pub fn worker(&self, name: &str) -> VntWorker {
self.wg.add(); self.wg.add();
VntWorker { VntWorker {
_name: name.to_string(), name: name.to_string(),
wg: self.wg.clone(), wg: self.wg.clone(),
status_s: self.status_s.clone(), status_s: self.status_s.clone(),
status_r: self.status_r.clone(), status_r: self.status_r.clone(),
+16 -15
View File
@@ -1,11 +1,11 @@
use std::io;
use std::ops::Deref;
use std::time::Duration;
use tokio::runtime::Runtime;
use crate::cipher::RsaCipher; use crate::cipher::RsaCipher;
use crate::core::{Config, Vnt, VntUtil}; use crate::core::{Config, Vnt, VntUtil};
use crate::handle::handshake_handler::HandshakeEnum; use crate::handle::handshake_handler::HandshakeEnum;
use crate::handle::registration_handler::{RegResponse, ReqEnum}; use crate::handle::registration_handler::{RegResponse, ReqEnum};
use std::io;
use std::ops::Deref;
use std::time::Duration;
use tokio::runtime::Runtime;
pub struct VntUtilSync { pub struct VntUtilSync {
vnt_util: VntUtil, vnt_util: VntUtil,
@@ -19,12 +19,11 @@ pub struct VntSync {
impl VntUtilSync { impl VntUtilSync {
pub fn new(config: Config) -> io::Result<VntUtilSync> { pub fn new(config: Config) -> io::Result<VntUtilSync> {
let runtime = tokio::runtime::Builder::new_multi_thread().enable_all().build().unwrap(); let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()?;
let vnt_util = runtime.block_on(VntUtil::new(config))?; let vnt_util = runtime.block_on(VntUtil::new(config))?;
Ok(VntUtilSync { Ok(VntUtilSync { vnt_util, runtime })
vnt_util,
runtime,
})
} }
pub fn connect(&mut self) -> io::Result<()> { pub fn connect(&mut self) -> io::Result<()> {
self.runtime.block_on(self.vnt_util.connect()) self.runtime.block_on(self.vnt_util.connect())
@@ -51,13 +50,14 @@ impl VntUtilSync {
let vnt = runtime.block_on(self.vnt_util.build())?; let vnt = runtime.block_on(self.vnt_util.build())?;
{ {
let mut vnt = vnt.clone(); let mut vnt = vnt.clone();
std::thread::spawn(move || { std::thread::spawn(move || runtime.block_on(vnt.wait_stop()));
runtime.block_on(vnt.wait_stop())
});
} }
Ok(VntSync { Ok(VntSync {
vnt, vnt,
runtime: tokio::runtime::Builder::new_current_thread().enable_all().build().unwrap(), runtime: tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap(),
}) })
} }
} }
@@ -67,7 +67,8 @@ impl VntSync {
self.runtime.block_on(self.vnt.wait_stop()) self.runtime.block_on(self.vnt.wait_stop())
} }
pub fn wait_stop_ms(&mut self, ms: u64) -> bool { pub fn wait_stop_ms(&mut self, ms: u64) -> bool {
self.runtime.block_on(self.vnt.wait_stop_ms(Duration::from_millis(ms))) self.runtime
.block_on(self.vnt.wait_stop_ms(Duration::from_millis(ms)))
} }
} }
@@ -77,4 +78,4 @@ impl Deref for VntSync {
fn deref(&self) -> &Self::Target { fn deref(&self) -> &Self::Target {
&self.vnt &self.vnt
} }
} }
+3 -3
View File
@@ -11,7 +11,7 @@ pub struct ExternalRoute {
impl ExternalRoute { impl ExternalRoute {
pub fn new(route_table: Vec<(u32, u32, Ipv4Addr)>) -> Self { pub fn new(route_table: Vec<(u32, u32, Ipv4Addr)>) -> Self {
Self { Self {
route_table: Arc::new(route_table) route_table: Arc::new(route_table),
} }
} }
pub fn route(&self, ip: &Ipv4Addr) -> Option<Ipv4Addr> { pub fn route(&self, ip: &Ipv4Addr) -> Option<Ipv4Addr> {
@@ -33,7 +33,7 @@ pub struct AllowExternalRoute {
impl AllowExternalRoute { impl AllowExternalRoute {
pub fn new(route_table: Vec<(u32, u32)>) -> Self { pub fn new(route_table: Vec<(u32, u32)>) -> Self {
Self { Self {
route_table: Arc::new(route_table) route_table: Arc::new(route_table),
} }
} }
pub fn allow(&self, ip: &Ipv4Addr) -> bool { pub fn allow(&self, ip: &Ipv4Addr) -> bool {
@@ -45,4 +45,4 @@ impl AllowExternalRoute {
} }
false false
} }
} }
+102 -54
View File
@@ -8,10 +8,8 @@ use tokio::net::{TcpStream, UdpSocket};
use crate::channel::channel::Context; use crate::channel::channel::Context;
use crate::cipher::{Cipher, RsaCipher}; use crate::cipher::{Cipher, RsaCipher};
use crate::proto::message::{HandshakeRequest, HandshakeResponse, SecretHandshakeRequest}; use crate::proto::message::{HandshakeRequest, HandshakeResponse, SecretHandshakeRequest};
use crate::protocol::{MAX_TTL, NetPacket, Protocol, service_packet, Version};
use crate::protocol::body::ENCRYPTION_RESERVED; use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
const VERSION: &'static str = "1.2.0";
pub enum HandshakeEnum { pub enum HandshakeEnum {
NotSecret, NotSecret,
@@ -24,7 +22,7 @@ pub enum HandshakeEnum {
fn handshake_request_packet(secret: bool) -> crate::Result<NetPacket<Vec<u8>>> { fn handshake_request_packet(secret: bool) -> crate::Result<NetPacket<Vec<u8>>> {
let mut request = HandshakeRequest::new(); let mut request = HandshakeRequest::new();
request.secret = secret; request.secret = secret;
request.version = VERSION.to_string(); request.version = crate::VNT_VERSION.to_string();
let bytes = request.write_to_bytes()?; let bytes = request.write_to_bytes()?;
let buf = vec![0u8; 12 + bytes.len()]; let buf = vec![0u8; 12 + bytes.len()];
let mut net_packet = NetPacket::new(buf)?; let mut net_packet = NetPacket::new(buf)?;
@@ -37,7 +35,11 @@ fn handshake_request_packet(secret: bool) -> crate::Result<NetPacket<Vec<u8>>> {
Ok(net_packet) Ok(net_packet)
} }
fn secret_handshake_request_packet(rsa_cipher: &RsaCipher, token: String, key: &[u8]) -> crate::Result<NetPacket<Vec<u8>>> { fn secret_handshake_request_packet(
rsa_cipher: &RsaCipher,
token: String,
key: &[u8],
) -> crate::Result<NetPacket<Vec<u8>>> {
let mut request = SecretHandshakeRequest::new(); let mut request = SecretHandshakeRequest::new();
request.token = token; request.token = token;
request.key = key.to_vec(); request.key = key.to_vec();
@@ -53,16 +55,25 @@ fn secret_handshake_request_packet(rsa_cipher: &RsaCipher, token: String, key: &
} }
/// 第一次握手,拿到公钥 /// 第一次握手,拿到公钥
pub async fn handshake(main_channel: &UdpSocket, main_tcp_channel: Option<&mut TcpStream>, pub async fn handshake(
server_address: SocketAddr, secret: bool) -> Result<Option<RsaCipher>, HandshakeEnum> { main_channel: &UdpSocket,
main_tcp_channel: Option<&mut TcpStream>,
server_address: SocketAddr,
secret: bool,
) -> Result<Option<RsaCipher>, HandshakeEnum> {
let request_packet = handshake_request_packet(secret).unwrap(); let request_packet = handshake_request_packet(secret).unwrap();
let send_buf = request_packet.buffer(); let send_buf = request_packet.buffer();
let mut recv_buf = [0u8; 10240]; let mut recv_buf = [0u8; 10240];
let len = send_recv(main_channel, main_tcp_channel, server_address, send_buf, &mut recv_buf).await?; let len = send_recv(
main_channel,
main_tcp_channel,
server_address,
send_buf,
&mut recv_buf,
)
.await?;
let net_packet = match NetPacket::new(&recv_buf[..len]) { let net_packet = match NetPacket::new(&recv_buf[..len]) {
Ok(net_packet) => { Ok(net_packet) => net_packet,
net_packet
}
Err(e) => { Err(e) => {
return Err(HandshakeEnum::Other(format!("net_packet {}", e))); return Err(HandshakeEnum::Other(format!("net_packet {}", e)));
} }
@@ -78,23 +89,31 @@ pub async fn handshake(main_channel: &UdpSocket, main_tcp_channel: Option<&mut T
return Err(HandshakeEnum::NotSecret); return Err(HandshakeEnum::NotSecret);
} }
if secret { if secret {
//转换公钥 //转换公钥
match RsaCipher::new(&response.public_key) { match RsaCipher::new(&response.public_key) {
Ok(rsa) => { Ok(rsa) => {
match rsa.finger() { match rsa.finger() {
Ok(finger) => { Ok(finger) => {
if finger != response.key_finger { if finger != response.key_finger {
return Err(HandshakeEnum::Other("finger error".to_string())); return Err(HandshakeEnum::Other(
"finger error".to_string(),
));
} }
} }
Err(e) => { Err(e) => {
return Err(HandshakeEnum::Other(format!("finger {}", e))); return Err(HandshakeEnum::Other(format!(
"finger {}",
e
)));
} }
} }
Ok(Some(rsa)) Ok(Some(rsa))
} }
Err(e) => { Err(e) => {
return Err(HandshakeEnum::Other(format!("RsaCipher {}", e))); return Err(HandshakeEnum::Other(format!(
"RsaCipher {}",
e
)));
} }
} }
} else { } else {
@@ -117,8 +136,13 @@ pub async fn handshake(main_channel: &UdpSocket, main_tcp_channel: Option<&mut T
} }
} }
async fn send_recv(main_channel: &UdpSocket, main_tcp_channel: Option<&mut TcpStream>, async fn send_recv(
server_address: SocketAddr, send_buf: &[u8], recv_buf: &mut [u8]) -> Result<usize, HandshakeEnum> { main_channel: &UdpSocket,
main_tcp_channel: Option<&mut TcpStream>,
server_address: SocketAddr,
send_buf: &[u8],
recv_buf: &mut [u8],
) -> Result<usize, HandshakeEnum> {
if let Some(main_tcp_channel) = main_tcp_channel { if let Some(main_tcp_channel) = main_tcp_channel {
let mut head = [0; 4]; let mut head = [0; 4];
let len = send_buf.len(); let len = send_buf.len();
@@ -145,20 +169,20 @@ async fn send_recv(main_channel: &UdpSocket, main_tcp_channel: Option<&mut TcpSt
if let Err(e) = main_channel.send_to(send_buf, server_address).await { if let Err(e) = main_channel.send_to(send_buf, server_address).await {
return Err(HandshakeEnum::Other(format!("send error:{}", e))); return Err(HandshakeEnum::Other(format!("send error:{}", e)));
} }
match tokio::time::timeout(Duration::from_millis(300), main_channel.recv_from(recv_buf)).await { match tokio::time::timeout(Duration::from_millis(300), main_channel.recv_from(recv_buf))
Ok(rs) => { .await
match rs { {
Ok((len, addr)) => { Ok(rs) => match rs {
if server_address != addr { Ok((len, addr)) => {
return Err(HandshakeEnum::Other(format!("invalid data,from {}", addr))); if server_address != addr {
} return Err(HandshakeEnum::Other(format!("invalid data,from {}", addr)));
Ok(len)
}
Err(e) => {
return Err(HandshakeEnum::Other(format!("receiver error:{}", e)));
} }
Ok(len)
} }
} Err(e) => {
return Err(HandshakeEnum::Other(format!("receiver error:{}", e)));
}
},
Err(_) => { Err(_) => {
return Err(HandshakeEnum::Timeout); return Err(HandshakeEnum::Timeout);
} }
@@ -167,48 +191,72 @@ async fn send_recv(main_channel: &UdpSocket, main_tcp_channel: Option<&mut TcpSt
} }
/// 第二次握手,同步对称密钥,后续将使用对称加密 /// 第二次握手,同步对称密钥,后续将使用对称加密
pub async fn secret_handshake(main_channel: &UdpSocket, main_tcp_channel: Option<&mut TcpStream>, pub async fn secret_handshake(
server_address: SocketAddr, rsa_cipher: &RsaCipher, server_cipher: &Cipher, token: String) main_channel: &UdpSocket,
-> Result<(), HandshakeEnum> { main_tcp_channel: Option<&mut TcpStream>,
let secret_packet = match secret_handshake_request_packet(rsa_cipher, token, server_cipher.key().unwrap()) { server_address: SocketAddr,
Ok(secret_packet) => { rsa_cipher: &RsaCipher,
secret_packet server_cipher: &Cipher,
} token: String,
Err(e) => { ) -> Result<(), HandshakeEnum> {
return Err(HandshakeEnum::Other(format!("secret_handshake_request_packet {}", e))); let secret_packet =
} match secret_handshake_request_packet(rsa_cipher, token, server_cipher.key().unwrap()) {
}; Ok(secret_packet) => secret_packet,
Err(e) => {
return Err(HandshakeEnum::Other(format!(
"secret_handshake_request_packet {}",
e
)));
}
};
let send_buf = secret_packet.buffer(); let send_buf = secret_packet.buffer();
let mut recv_buf = [0u8; 10240]; let mut recv_buf = [0u8; 10240];
let len = send_recv(main_channel, main_tcp_channel, server_address, send_buf, &mut recv_buf).await?; let len = send_recv(
main_channel,
main_tcp_channel,
server_address,
send_buf,
&mut recv_buf,
)
.await?;
let mut net_packet = match NetPacket::new(&mut recv_buf[..len]) { let mut net_packet = match NetPacket::new(&mut recv_buf[..len]) {
Ok(net_packet) => { net_packet } Ok(net_packet) => net_packet,
Err(e) => { Err(e) => {
return Err(HandshakeEnum::Other(format!("secret_net_packet {}", e))); return Err(HandshakeEnum::Other(format!("secret_net_packet {}", e)));
} }
}; };
match server_cipher.decrypt_ipv4(&mut net_packet) { match server_cipher.decrypt_ipv4(&mut net_packet) {
Ok(_) => { Ok(_) => {
if net_packet.is_gateway() && net_packet.protocol() == Protocol::Service if net_packet.is_gateway()
&& service_packet::Protocol::from(net_packet.transport_protocol()) == && net_packet.protocol() == Protocol::Service
service_packet::Protocol::SecretHandshakeResponse { && service_packet::Protocol::from(net_packet.transport_protocol())
== service_packet::Protocol::SecretHandshakeResponse
{
Ok(()) Ok(())
} else { } else {
Err(HandshakeEnum::Other("not match".to_string())) Err(HandshakeEnum::Other("not match".to_string()))
} }
} }
Err(e) => { Err(e) => Err(HandshakeEnum::Other(format!("decrypt_ipv4 {}", e))),
Err(HandshakeEnum::Other(format!("decrypt_ipv4 {}", e)))
}
} }
} }
pub async fn secret_handshake_req(context: &Context, pub async fn secret_handshake_req(
server_address: SocketAddr, rsa_cipher: &RsaCipher, server_cipher: &Cipher, token: String, ) -> crate::Result<()> { context: &Context,
let secret_packet = secret_handshake_request_packet(rsa_cipher, token, server_cipher.key().unwrap())?; server_address: SocketAddr,
context.send_main(secret_packet.buffer(), server_address).await?; rsa_cipher: &RsaCipher,
if context.is_main_tcp(){ server_cipher: &Cipher,
context.send_main_udp(secret_packet.buffer(),server_address).await?; token: String,
) -> crate::Result<()> {
let secret_packet =
secret_handshake_request_packet(rsa_cipher, token, server_cipher.key().unwrap())?;
context
.send_main(secret_packet.buffer(), server_address)
.await?;
if context.is_main_tcp() {
context
.send_main_udp(secret_packet.buffer(), server_address)
.await?;
} }
Ok(()) Ok(())
} }
+92 -35
View File
@@ -1,22 +1,21 @@
use std::io;
use std::net::{Ipv4Addr, ToSocketAddrs}; use std::net::{Ipv4Addr, ToSocketAddrs};
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use std::io;
use crate::channel::idle::Idle;
use crate::channel::sender::ChannelSender;
use crate::channel::Route;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex; use parking_lot::Mutex;
use rand::prelude::SliceRandom; use rand::prelude::SliceRandom;
use crate::channel::idle::Idle;
use crate::channel::Route;
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::protocol::control_packet::PingPacket;
use crate::protocol::{control_packet, MAX_TTL, NetPacket, Protocol, Version};
use crate::protocol::body::ENCRYPTION_RESERVED; use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::PingPacket;
use crate::protocol::{control_packet, NetPacket, Protocol, Version, MAX_TTL};
pub fn start_idle(mut worker: VntWorker, idle: Idle, sender: ChannelSender) { pub fn start_idle(mut worker: VntWorker, idle: Idle, sender: ChannelSender) {
tokio::spawn(async move { tokio::spawn(async move {
@@ -35,13 +34,10 @@ pub fn start_idle(mut worker: VntWorker, idle: Idle, sender: ChannelSender) {
} }
async fn start_idle_(idle: Idle, sender: ChannelSender) -> io::Result<()> { async fn start_idle_(idle: Idle, sender: ChannelSender) -> io::Result<()> {
log::info!("启动空闲检查任务");
loop { loop {
let (peer_ip, route) = idle.next_idle().await?; let (peer_ip, route) = idle.next_idle().await?;
log::info!( log::info!("路由空闲 peer_ip:{:?},route:{:?}", peer_ip, route);
"peer_ip:{:?},route:{:?}",
peer_ip,
route
);
sender.remove_route(&peer_ip, route); sender.remove_route(&peer_ip, route);
} }
} }
@@ -70,14 +66,20 @@ pub fn start_heartbeat(
}); });
} }
fn heartbeat_packet(
fn heartbeat_packet(device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>, client_cipher: &Cipher, server_cipher: &Cipher, gateway: bool, src: Ipv4Addr, dest: Ipv4Addr) -> NetPacket<[u8; 48]> { ttl: u8,
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
client_cipher: &Cipher,
server_cipher: &Cipher,
gateway: bool,
src: Ipv4Addr,
dest: Ipv4Addr,
) -> NetPacket<[u8; 48]> {
let mut net_packet = NetPacket::new_encrypt([0u8; 12 + 4 + ENCRYPTION_RESERVED]).unwrap(); let mut net_packet = NetPacket::new_encrypt([0u8; 12 + 4 + ENCRYPTION_RESERVED]).unwrap();
net_packet.set_version(Version::V1); net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::Control); net_packet.set_protocol(Protocol::Control);
net_packet.set_transport_protocol(control_packet::Protocol::Ping.into()); net_packet.set_transport_protocol(control_packet::Protocol::Ping.into());
//只寻找两跳以内能到的目标 net_packet.first_set_ttl(ttl);
net_packet.first_set_ttl(2);
net_packet.set_source(src); net_packet.set_source(src);
net_packet.set_destination(dest); net_packet.set_destination(dest);
{ {
@@ -104,6 +106,7 @@ async fn start_heartbeat_(
server_cipher: Cipher, server_cipher: Cipher,
) -> io::Result<()> { ) -> io::Result<()> {
let mut count = 0; let mut count = 0;
log::info!("启动心跳任务");
loop { loop {
if sender.is_close() { if sender.is_close() {
return Ok(()); return Ok(());
@@ -115,14 +118,14 @@ async fn start_heartbeat_(
packet.set_version(Version::V1); packet.set_version(Version::V1);
packet.set_gateway_flag(true); packet.set_gateway_flag(true);
packet.set_protocol(Protocol::Control); packet.set_protocol(Protocol::Control);
packet.set_transport_protocol( packet.set_transport_protocol(control_packet::Protocol::AddrRequest.into());
control_packet::Protocol::AddrRequest.into(),
);
packet.first_set_ttl(MAX_TTL); packet.first_set_ttl(MAX_TTL);
packet.set_source(current_dev.virtual_ip()); packet.set_source(current_dev.virtual_ip());
packet.set_destination(current_dev.virtual_gateway); packet.set_destination(current_dev.virtual_gateway);
server_cipher.encrypt_ipv4(&mut packet)?; server_cipher.encrypt_ipv4(&mut packet)?;
let _ = sender.send_main_udp(packet.buffer(), current_dev.connect_server).await; let _ = sender
.send_main_udp(packet.buffer(), current_dev.connect_server)
.await;
} }
if count % 20 == 19 { if count % 20 == 19 {
if let Ok(mut addr) = server_address_str.to_socket_addrs() { if let Ok(mut addr) = server_address_str.to_socket_addrs() {
@@ -130,6 +133,11 @@ async fn start_heartbeat_(
if addr != current_dev.connect_server { if addr != current_dev.connect_server {
let mut tmp = current_dev.clone(); let mut tmp = current_dev.clone();
tmp.connect_server = addr; tmp.connect_server = addr;
log::info!(
"服务端地址变化,旧地址:{},新地址:{}",
current_dev.connect_server,
addr
);
if current_device.compare_exchange(current_dev, tmp).is_ok() { if current_device.compare_exchange(current_dev, tmp).is_ok() {
current_dev.connect_server = addr; current_dev.connect_server = addr;
} }
@@ -138,14 +146,20 @@ async fn start_heartbeat_(
} }
} }
let src = current_dev.virtual_ip(); let src = current_dev.virtual_ip();
let server_packet = heartbeat_packet(&device_list, &client_cipher, &server_cipher, true, src, current_dev.virtual_gateway); let server_packet = heartbeat_packet(
if let Err(e) = sender.send_main(server_packet.buffer(), current_dev.connect_server).await MAX_TTL,
&device_list,
&client_cipher,
&server_cipher,
true,
src,
current_dev.virtual_gateway,
);
if let Err(e) = sender
.send_main(server_packet.buffer(), current_dev.connect_server)
.await
{ {
log::warn!( log::warn!("connect_server:{:?},e:{:?}", current_dev.connect_server, e);
"connect_server:{:?},e:{:?}",
current_dev.connect_server,
e
);
} }
if count < 7 || count % 7 == 0 { if count < 7 || count % 7 == 0 {
let mut route_list: Option<Vec<(Ipv4Addr, Vec<Route>)>> = None; let mut route_list: Option<Vec<(Ipv4Addr, Vec<Route>)>> = None;
@@ -154,15 +168,38 @@ async fn start_heartbeat_(
if peer.virtual_ip == current_dev.virtual_ip { if peer.virtual_ip == current_dev.virtual_ip {
continue; continue;
} }
let client_packet = heartbeat_packet(&device_list, &client_cipher, &server_cipher, false, src, peer.virtual_ip); let client_packet = heartbeat_packet(
MAX_TTL,
&device_list,
&client_cipher,
&server_cipher,
false,
src,
peer.virtual_ip,
);
if let Some(route) = sender.route_one(&peer.virtual_ip) { if let Some(route) = sender.route_one(&peer.virtual_ip) {
let _ = sender.send_by_key(client_packet.buffer(), &route.route_key()).await; if let Err(e) = sender
.send_by_key(client_packet.buffer(), &route.route_key())
.await
{
log::warn!("virtual_ip:{},route:{:?},e:{:?}", peer.virtual_ip, route, e);
}
if route.is_p2p() { if route.is_p2p() {
continue; continue;
} }
} else { } else {
//没有直连路由则发送到网关 //没有直连路由则发送到网关
let _ = sender.send_main(client_packet.buffer(), current_dev.connect_server).await; if let Err(e) = sender
.send_main(client_packet.buffer(), current_dev.connect_server)
.await
{
log::warn!(
"virtual_ip:{},connect_server:{:?},e:{:?}",
peer.virtual_ip,
current_dev.connect_server,
e
);
}
} }
//再随机发送到其他地址,看有没有客户端符合转发条件 //再随机发送到其他地址,看有没有客户端符合转发条件
@@ -175,7 +212,16 @@ async fn start_heartbeat_(
'a: for (peer_ip, route_list) in route_list.iter() { 'a: for (peer_ip, route_list) in route_list.iter() {
for route in route_list { for route in route_list {
if peer_ip != &peer.virtual_ip && route.is_p2p() { if peer_ip != &peer.virtual_ip && route.is_p2p() {
let _ = sender.try_send_by_key(client_packet.buffer(), &route.route_key()); if let Err(e) =
sender.try_send_by_key(client_packet.buffer(), &route.route_key())
{
log::warn!(
"virtual_ip:{},route:{:?},e:{:?}",
peer.virtual_ip,
route,
e
);
}
num += 1; num += 1;
break; break;
} }
@@ -191,9 +237,20 @@ async fn start_heartbeat_(
if peer_ip == &current_dev.virtual_gateway { if peer_ip == &current_dev.virtual_gateway {
continue; continue;
} }
let client_packet = heartbeat_packet(&device_list, &client_cipher, &server_cipher, false, src, *peer_ip); let client_packet = heartbeat_packet(
MAX_TTL,
&device_list,
&client_cipher,
&server_cipher,
false,
src,
*peer_ip,
);
for route in route_list { for route in route_list {
if let Err(e) = sender.send_by_key(client_packet.buffer(), &route.route_key()).await { if let Err(e) = sender
.send_by_key(client_packet.buffer(), &route.route_key())
.await
{
log::warn!("peer_ip:{:?},route:{:?},e:{:?}", peer_ip, route, e); log::warn!("peer_ip:{:?},route:{:?},e:{:?}", peer_ip, route, e);
} }
tokio::time::sleep(Duration::from_millis(2)).await; tokio::time::sleep(Duration::from_millis(2)).await;
+3 -4
View File
@@ -31,17 +31,17 @@ pub struct PeerDeviceInfo {
} }
impl PeerDeviceInfo { impl PeerDeviceInfo {
pub fn new(virtual_ip: Ipv4Addr, name: String, status: u8,client_secret: bool) -> Self { pub fn new(virtual_ip: Ipv4Addr, name: String, status: u8, client_secret: bool) -> Self {
Self { Self {
virtual_ip, virtual_ip,
name, name,
status: PeerDeviceStatus::from(status), status: PeerDeviceStatus::from(status),
client_secret client_secret,
} }
} }
} }
#[derive(Copy, Clone, Debug, Eq, PartialEq,Ord, PartialOrd)] #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd)]
pub enum PeerDeviceStatus { pub enum PeerDeviceStatus {
Online, Online,
Offline, Offline,
@@ -82,7 +82,6 @@ pub struct CurrentDeviceInfo {
pub broadcast_address: Ipv4Addr, pub broadcast_address: Ipv4Addr,
//链接的服务器地址 //链接的服务器地址
pub connect_server: SocketAddr, pub connect_server: SocketAddr,
} }
impl CurrentDeviceInfo { impl CurrentDeviceInfo {
+46 -17
View File
@@ -1,25 +1,29 @@
use crate::channel::punch::{NatInfo, Punch};
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::nat::NatTest; use crate::nat::NatTest;
use crate::proto::message::{PunchInfo, PunchNatType}; use crate::proto::message::{PunchInfo, PunchNatType};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{control_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL}; use crate::protocol::{control_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL};
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex; use parking_lot::Mutex;
use protobuf::Message; use protobuf::Message;
use rand::prelude::SliceRandom; use rand::prelude::SliceRandom;
use std::io;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use std::io;
use tokio::sync::mpsc::Receiver; use tokio::sync::mpsc::Receiver;
use crate::channel::punch::{NatInfo, Punch};
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use crate::protocol::body::ENCRYPTION_RESERVED;
pub fn start(mut worker: VntWorker, receiver: Receiver<(Ipv4Addr, NatInfo)>, pub fn start(
punch: Punch, current_device: Arc<AtomicCell<CurrentDeviceInfo>>, mut worker: VntWorker,
client_cipher: Cipher, ) { receiver: Receiver<(Ipv4Addr, NatInfo)>,
punch: Punch,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) {
tokio::spawn(async move { tokio::spawn(async move {
tokio::select! { tokio::select! {
_=start0(receiver, punch, current_device,client_cipher)=>{} _=start0(receiver, punch, current_device,client_cipher)=>{}
@@ -31,11 +35,23 @@ pub fn start(mut worker: VntWorker, receiver: Receiver<(Ipv4Addr, NatInfo)>,
}); });
} }
pub async fn start0(mut receiver: Receiver<(Ipv4Addr, NatInfo)>, pub async fn start0(
mut punch: Punch, current_device: Arc<AtomicCell<CurrentDeviceInfo>>, mut receiver: Receiver<(Ipv4Addr, NatInfo)>,
client_cipher: Cipher, ) { mut punch: Punch,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) {
log::info!("启动打洞任务");
while let Some((peer_ip, nat_info)) = receiver.recv().await { while let Some((peer_ip, nat_info)) = receiver.recv().await {
if let Err(e) = start_(&client_cipher, &mut punch, &current_device, peer_ip, nat_info).await { if let Err(e) = start_(
&client_cipher,
&mut punch,
&current_device,
peer_ip,
nat_info,
)
.await
{
log::warn!("网络打洞异常 {:?}", e); log::warn!("网络打洞异常 {:?}", e);
} }
} }
@@ -70,6 +86,7 @@ pub async fn start_punch(
) { ) {
let mut num = 0; let mut num = 0;
let sleep_time = [3, 5, 7, 11, 13, 17, 19, 23, 29]; let sleep_time = [3, 5, 7, 11, 13, 17, 19, 23, 29];
log::info!("启动发起打洞请求任务");
loop { loop {
if sender.is_close() { if sender.is_close() {
break; break;
@@ -113,8 +130,16 @@ async fn start_punch_(
if count > 2 { if count > 2 {
break; break;
} }
let packet = punch_packet(client_cipher, current_device.virtual_ip(), &nat_info, info.virtual_ip)?; let packet = punch_packet(
let _ = sender.send_main(packet.buffer(), current_device.connect_server).await; client_cipher,
current_device.virtual_ip(),
&nat_info,
info.virtual_ip,
)
.unwrap();
let _ = sender
.send_main(packet.buffer(), current_device.connect_server)
.await;
} }
tokio::time::sleep(sleep_time).await; tokio::time::sleep(sleep_time).await;
Ok(()) Ok(())
@@ -135,8 +160,12 @@ pub fn punch_packet(
.collect(); .collect();
punch_reply.public_port = nat_info.public_port as u32; punch_reply.public_port = nat_info.public_port as u32;
punch_reply.public_port_range = nat_info.public_port_range as u32; punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.local_ip = u32::from_be_bytes(nat_info.local_ip.octets()); punch_reply.local_ip = u32::from_be_bytes(nat_info.local_ipv4_addr.ip().octets());
punch_reply.local_port = nat_info.local_port as u32; punch_reply.local_port = nat_info.local_ipv4_addr.port() as u32;
if !nat_info.ipv6_addr.ip().is_unspecified() {
punch_reply.ipv6_port = nat_info.ipv6_addr.port() as u32;
punch_reply.ipv6 = nat_info.ipv6_addr.ip().octets().to_vec();
}
punch_reply.nat_type = protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type)); punch_reply.nat_type = protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
let bytes = punch_reply.write_to_bytes()?; let bytes = punch_reply.write_to_bytes()?;
let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?; let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
+333 -140
View File
@@ -1,4 +1,4 @@
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; use std::net::{Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6};
use std::sync::Arc; use std::sync::Arc;
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
@@ -7,29 +7,32 @@ use parking_lot::Mutex;
use protobuf::Message; use protobuf::Message;
use tokio::sync::mpsc::Sender; use tokio::sync::mpsc::Sender;
use packet::icmp::{icmp, Kind};
use packet::icmp::icmp::HeaderOther; use packet::icmp::icmp::HeaderOther;
use packet::icmp::{icmp, Kind};
use packet::ip::ipv4; use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet; use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::{Route, RouteKey};
use crate::channel::channel::Context; use crate::channel::channel::Context;
use crate::channel::punch::{NatInfo, NatType}; use crate::channel::punch::{NatInfo, NatType};
use crate::channel::{Route, RouteKey};
use crate::cipher::{Cipher, RsaCipher}; use crate::cipher::{Cipher, RsaCipher};
use crate::error::Error; use crate::error::Error;
use crate::external_route::AllowExternalRoute; use crate::external_route::AllowExternalRoute;
use crate::handle::{ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo, PeerDeviceStatus};
use crate::handle::handshake_handler::secret_handshake_req; use crate::handle::handshake_handler::secret_handshake_req;
use crate::handle::registration_handler::Register; use crate::handle::registration_handler::Register;
use crate::handle::{ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo, PeerDeviceStatus};
use crate::igmp_server::IgmpServer; use crate::igmp_server::IgmpServer;
use crate::ip_proxy::IpProxyMap; use crate::ip_proxy::IpProxyMap;
use crate::nat; use crate::nat;
use crate::nat::NatTest; use crate::nat::NatTest;
use crate::proto::message::{DeviceList, PunchInfo, PunchNatType, RegistrationResponse}; use crate::proto::message::{DeviceList, PunchInfo, PunchNatType, RegistrationResponse};
use crate::protocol::{control_packet, ip_turn_packet, MAX_TTL, NetPacket, other_turn_packet, Protocol, service_packet, Version};
use crate::protocol::body::ENCRYPTION_RESERVED; use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::ControlPacket; use crate::protocol::control_packet::ControlPacket;
use crate::protocol::error_packet::InErrorPacket; use crate::protocol::error_packet::InErrorPacket;
use crate::protocol::{
control_packet, ip_turn_packet, other_turn_packet, service_packet, NetPacket, Protocol,
Version, MAX_TTL,
};
use crate::tun_tap_device::DeviceWriter; use crate::tun_tap_device::DeviceWriter;
#[derive(Clone)] #[derive(Clone)]
@@ -54,22 +57,25 @@ pub struct ChannelDataHandler {
} }
impl ChannelDataHandler { impl ChannelDataHandler {
pub fn new(current_device: Arc<AtomicCell<CurrentDeviceInfo>>, pub fn new(
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>, current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
register: Arc<Register>, device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
nat_test: NatTest, register: Arc<Register>,
igmp_server: Option<IgmpServer>, nat_test: NatTest,
device_writer: DeviceWriter, igmp_server: Option<IgmpServer>,
connect_status: Arc<AtomicCell<ConnectStatus>>, device_writer: DeviceWriter,
peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>>, connect_status: Arc<AtomicCell<ConnectStatus>>,
ip_proxy_map: Option<IpProxyMap>, peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>>,
out_external_route: AllowExternalRoute, ip_proxy_map: Option<IpProxyMap>,
cone_sender: Sender<(Ipv4Addr, NatInfo)>, out_external_route: AllowExternalRoute,
symmetric_sender: Sender<(Ipv4Addr, NatInfo)>, cone_sender: Sender<(Ipv4Addr, NatInfo)>,
client_cipher: Cipher, symmetric_sender: Sender<(Ipv4Addr, NatInfo)>,
server_cipher: Cipher, client_cipher: Cipher,
rsa_cipher: Option<RsaCipher>, server_cipher: Cipher,
relay: bool, token: String, ) -> Self { rsa_cipher: Option<RsaCipher>,
relay: bool,
token: String,
) -> Self {
Self { Self {
current_device, current_device,
device_list, device_list,
@@ -92,18 +98,29 @@ impl ChannelDataHandler {
} }
} }
impl ChannelDataHandler { impl ChannelDataHandler {
pub async fn handle(&self, buf: &mut [u8], start: usize, end: usize, route_key: RouteKey, context: &Context) { pub async fn handle(
&self,
buf: &mut [u8],
start: usize,
end: usize,
route_key: RouteKey,
context: &Context,
) {
assert_eq!(start, 14); assert_eq!(start, 14);
match self.handle0(&mut buf[..end], &route_key, context).await { match self.handle0(&mut buf[..end], &route_key, context).await {
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
log::warn!("{:?}",e); log::warn!("{:?}", e);
} }
} }
} }
async fn handle0(&self, buf: &mut [u8], route_key: &RouteKey, context: &Context) -> crate::Result<()> { async fn handle0(
&self,
buf: &mut [u8],
route_key: &RouteKey,
context: &Context,
) -> crate::Result<()> {
let mut net_packet = NetPacket::new(&mut buf[14..])?; let mut net_packet = NetPacket::new(&mut buf[14..])?;
if net_packet.ttl() == 0 || net_packet.source_ttl() < net_packet.ttl() { if net_packet.ttl() == 0 || net_packet.source_ttl() < net_packet.ttl() {
return Ok(()); return Ok(());
@@ -112,9 +129,13 @@ impl ChannelDataHandler {
context.update_read_time(&source, route_key); context.update_read_time(&source, route_key);
let current_device = self.current_device.load(); let current_device = self.current_device.load();
let destination = net_packet.destination(); let destination = net_packet.destination();
let not_broadcast = !destination.is_broadcast() && !destination.is_multicast() && destination != current_device.broadcast_address; let not_broadcast = !destination.is_broadcast()
&& !destination.is_multicast()
&& destination != current_device.broadcast_address;
if current_device.virtual_ip() != destination if current_device.virtual_ip() != destination
&& not_broadcast && !destination.is_unspecified() { && not_broadcast
&& !destination.is_unspecified()
{
//校验指纹,不需要解密 //校验指纹,不需要解密
self.client_cipher.check_finger(&net_packet)?; self.client_cipher.check_finger(&net_packet)?;
net_packet.set_ttl(net_packet.ttl() - 1); net_packet.set_ttl(net_packet.ttl() - 1);
@@ -123,26 +144,42 @@ impl ChannelDataHandler {
// 转发 // 转发
if let Some(route) = context.route_one(&destination) { if let Some(route) = context.route_one(&destination) {
if route.metric <= net_packet.ttl() { if route.metric <= net_packet.ttl() {
context.send_by_key(net_packet.buffer(), &route.route_key()).await?; context
.send_by_key(net_packet.buffer(), &route.route_key())
.await?;
} }
} else if (ttl > 1 || destination == current_device.virtual_gateway()) } else if (ttl > 1 || destination == current_device.virtual_gateway())
&& source != current_device.virtual_gateway() { && source != current_device.virtual_gateway()
{
//网关默认要转发一次,生存时间不够的发到网关也会被丢弃 //网关默认要转发一次,生存时间不够的发到网关也会被丢弃
context.send_main(net_packet.buffer(), current_device.connect_server).await?; context
.send_main(net_packet.buffer(), current_device.connect_server)
.await?;
} }
} }
return Ok(()); return Ok(());
} }
if net_packet.is_gateway() { if net_packet.is_gateway() {
if net_packet.protocol() == Protocol::Error && net_packet.transport_protocol() == crate::protocol::error_packet::Protocol::NoKey.into() { if net_packet.protocol() == Protocol::Error
&& net_packet.transport_protocol()
== crate::protocol::error_packet::Protocol::NoKey.into()
{
if let Some(rsa_cipher) = &self.rsa_cipher { if let Some(rsa_cipher) = &self.rsa_cipher {
secret_handshake_req(context, current_device.connect_server, rsa_cipher, &self.server_cipher, self.token.clone()).await?; secret_handshake_req(
context,
current_device.connect_server,
rsa_cipher,
&self.server_cipher,
self.token.clone(),
)
.await?;
} }
} else { } else {
//服务端解密 //服务端解密
self.server_cipher.decrypt_ipv4(&mut net_packet)?; self.server_cipher.decrypt_ipv4(&mut net_packet)?;
let data_len = net_packet.data_len(); let data_len = net_packet.data_len();
self.server_packet_handle(context, current_device, buf, data_len, route_key).await?; self.server_packet_handle(context, current_device, buf, data_len, route_key)
.await?;
} }
return Ok(()); return Ok(());
} }
@@ -161,7 +198,8 @@ impl ChannelDataHandler {
} }
ipv4::protocol::Protocol::Icmp => { ipv4::protocol::Protocol::Icmp => {
if ipv4.destination_ip() == destination { if ipv4.destination_ip() == destination {
let mut icmp_packet = icmp::IcmpPacket::new(ipv4.payload_mut())?; let mut icmp_packet =
icmp::IcmpPacket::new(ipv4.payload_mut())?;
if icmp_packet.kind() == Kind::EchoRequest { if icmp_packet.kind() == Kind::EchoRequest {
//开启ping //开启ping
icmp_packet.set_kind(Kind::EchoReply); icmp_packet.set_kind(Kind::EchoReply);
@@ -187,56 +225,81 @@ impl ChannelDataHandler {
ipv4::protocol::Protocol::Tcp => { ipv4::protocol::Protocol::Tcp => {
let dest_ip = ipv4.destination_ip(); let dest_ip = ipv4.destination_ip();
//转发到代理目标地址 //转发到代理目标地址
let mut tcp_packet = packet::tcp::tcp::TcpPacket::new(source, destination, ipv4.payload_mut())?; let mut tcp_packet = packet::tcp::tcp::TcpPacket::new(
source,
destination,
ipv4.payload_mut(),
)?;
let source_port = tcp_packet.source_port(); let source_port = tcp_packet.source_port();
let dest_port = tcp_packet.destination_port(); let dest_port = tcp_packet.destination_port();
tcp_packet.set_destination_port(ip_proxy_map.tcp_proxy_port); tcp_packet
.set_destination_port(ip_proxy_map.tcp_proxy_port);
tcp_packet.update_checksum(); tcp_packet.update_checksum();
ipv4.set_destination_ip(destination); ipv4.set_destination_ip(destination);
ipv4.update_checksum(); ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port); let key = SocketAddrV4::new(source, source_port);
//https://github.com/crossbeam-rs/crossbeam/issues/1023 //https://github.com/crossbeam-rs/crossbeam/issues/1023
ip_proxy_map.tcp_proxy_map.insert(key, SocketAddrV4::new(dest_ip, dest_port)); ip_proxy_map
.tcp_proxy_map
.insert(key, SocketAddrV4::new(dest_ip, dest_port));
} }
ipv4::protocol::Protocol::Udp => { ipv4::protocol::Protocol::Udp => {
let dest_ip = ipv4.destination_ip(); let dest_ip = ipv4.destination_ip();
//转发到代理目标地址 //转发到代理目标地址
let mut udp_packet = packet::udp::udp::UdpPacket::new(source, destination, ipv4.payload_mut())?; let mut udp_packet = packet::udp::udp::UdpPacket::new(
source,
destination,
ipv4.payload_mut(),
)?;
let source_port = udp_packet.source_port(); let source_port = udp_packet.source_port();
let dest_port = udp_packet.destination_port(); let dest_port = udp_packet.destination_port();
udp_packet.set_destination_port(ip_proxy_map.udp_proxy_port); udp_packet
.set_destination_port(ip_proxy_map.udp_proxy_port);
udp_packet.update_checksum(); udp_packet.update_checksum();
ipv4.set_destination_ip(destination); ipv4.set_destination_ip(destination);
ipv4.update_checksum(); ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port); let key = SocketAddrV4::new(source, source_port);
ip_proxy_map.udp_proxy_map.insert(key, SocketAddrV4::new(dest_ip, dest_port)); ip_proxy_map
.udp_proxy_map
.insert(key, SocketAddrV4::new(dest_ip, dest_port));
} }
ipv4::protocol::Protocol::Icmp => { ipv4::protocol::Protocol::Icmp => {
let dest_ip = ipv4.destination_ip(); let dest_ip = ipv4.destination_ip();
//转发到代理目标地址 //转发到代理目标地址
let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?; let icmp_packet =
icmp::IcmpPacket::new(ipv4.payload())?;
match icmp_packet.header_other() { match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => { HeaderOther::Identifier(id, seq) => {
ip_proxy_map.icmp_proxy_map.insert((dest_ip, id, seq), source); ip_proxy_map
ip_proxy_map.send_icmp(ipv4.payload(), &dest_ip)?; .icmp_proxy_map
.insert((dest_ip, id, seq), source);
ip_proxy_map
.send_icmp(ipv4.payload(), &dest_ip)?;
} }
_ => { _ => {
log::warn!("不支持的ip代理Icmp协议:{}",destination); log::warn!(
return Err(Error::Warn("不支持的ip代理Icmp协议".to_string())); "不支持的ip代理Icmp协议:{}",
destination
);
return Err(Error::Warn(
"不支持的ip代理Icmp协议".to_string(),
));
} }
} }
} }
_ => { _ => {
log::warn!("不支持的ip代理ipv4协议:{}",destination); log::warn!("不支持的ip代理ipv4协议:{}", destination);
return Err(Error::Warn("不支持的ip代理ipv4协议".to_string())); return Err(Error::Warn(
"不支持的ip代理ipv4协议".to_string(),
));
} }
} }
} else { } else {
log::warn!("没有ip代理规则:{}",destination); log::warn!("没有ip代理规则:{}", destination);
return Err(Error::Warn("没有ip代理规则".to_string())); return Err(Error::Warn("没有ip代理规则".to_string()));
} }
} else { } else {
log::warn!("不支持ip代理:{}",destination); log::warn!("不支持ip代理:{}", destination);
return Err(Error::Warn("不支持ip代理".to_string())); return Err(Error::Warn("不支持ip代理".to_string()));
} }
} }
@@ -254,19 +317,30 @@ impl ChannelDataHandler {
Protocol::Service => {} Protocol::Service => {}
Protocol::Error => {} Protocol::Error => {}
Protocol::Control => { Protocol::Control => {
self.control(context, current_device, source, net_packet, route_key).await?; self.control(context, current_device, source, net_packet, route_key)
.await?;
} }
Protocol::OtherTurn => { Protocol::OtherTurn => {
self.other_turn(context, current_device, source, net_packet, route_key).await?; self.other_turn(context, current_device, source, net_packet, route_key)
.await?;
} }
Protocol::UnKnow(e) => { Protocol::UnKnow(e) => {
log::info!("不支持的协议:{}",e); log::info!("不支持的协议:{}", e);
} }
} }
Ok(()) Ok(())
} }
async fn pong_packet(&self, gateway: bool, metric: u8, context: &Context, current_device: CurrentDeviceInfo, source: Ipv4Addr, pong_packet: control_packet::PongPacket<&[u8]>, route_key: &RouteKey) -> crate::Result<()> { async fn pong_packet(
&self,
gateway: bool,
metric: u8,
context: &Context,
current_device: CurrentDeviceInfo,
source: Ipv4Addr,
pong_packet: control_packet::PongPacket<&[u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
let current_time = crate::handle::now_time() as u16; let current_time = crate::handle::now_time() as u16;
if current_time < pong_packet.time() { if current_time < pong_packet.time() {
return Ok(()); return Ok(());
@@ -286,12 +360,21 @@ impl ChannelDataHandler {
poll_device.set_protocol(Protocol::Service); poll_device.set_protocol(Protocol::Service);
poll_device.set_transport_protocol(service_packet::Protocol::PollDeviceList.into()); poll_device.set_transport_protocol(service_packet::Protocol::PollDeviceList.into());
self.server_cipher.encrypt_ipv4(&mut poll_device)?; self.server_cipher.encrypt_ipv4(&mut poll_device)?;
context.send_main(poll_device.buffer(), current_device.connect_server).await?; context
.send_main(poll_device.buffer(), current_device.connect_server)
.await?;
} }
} }
Ok(()) Ok(())
} }
async fn control(&self, context: &Context, current_device: CurrentDeviceInfo, source: Ipv4Addr, mut net_packet: NetPacket<&mut [u8]>, route_key: &RouteKey) -> crate::Result<()> { async fn control(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
source: Ipv4Addr,
mut net_packet: NetPacket<&mut [u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
let metric = net_packet.source_ttl() - net_packet.ttl() + 1; let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? { match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PingPacket(_) => { ControlPacket::PingPacket(_) => {
@@ -305,7 +388,16 @@ impl ChannelDataHandler {
context.add_route_if_absent(source, route); context.add_route_if_absent(source, route);
} }
ControlPacket::PongPacket(pong_packet) => { ControlPacket::PongPacket(pong_packet) => {
self.pong_packet(false, metric, context, current_device, source, pong_packet, route_key).await?; self.pong_packet(
false,
metric,
context,
current_device,
source,
pong_packet,
route_key,
)
.await?;
} }
ControlPacket::PunchRequest => { ControlPacket::PunchRequest => {
if self.relay { if self.relay {
@@ -328,129 +420,170 @@ impl ChannelDataHandler {
let route = Route::from(*route_key, 1, 199); let route = Route::from(*route_key, 1, 199);
context.add_route_if_absent(source, route); context.add_route_if_absent(source, route);
} }
ControlPacket::AddrRequest => { ControlPacket::AddrRequest => match route_key.addr.ip() {
match route_key.addr.ip() { std::net::IpAddr::V4(ipv4) => {
std::net::IpAddr::V4(ipv4) => { let mut packet = NetPacket::new_encrypt([0; 12 + 6 + ENCRYPTION_RESERVED])?;
let mut packet = NetPacket::new_encrypt([0; 12 + 6 + ENCRYPTION_RESERVED])?; packet.set_version(Version::V1);
packet.set_version(Version::V1); packet.set_protocol(Protocol::Control);
packet.set_protocol(Protocol::Control); packet.set_transport_protocol(control_packet::Protocol::AddrResponse.into());
packet.set_transport_protocol( packet.first_set_ttl(MAX_TTL);
control_packet::Protocol::AddrResponse.into(), packet.set_source(current_device.virtual_ip());
); packet.set_destination(source);
packet.first_set_ttl(MAX_TTL); let mut addr_packet = control_packet::AddrPacket::new(packet.payload_mut())?;
packet.set_source(current_device.virtual_ip()); addr_packet.set_ipv4(ipv4);
packet.set_destination(source); addr_packet.set_port(route_key.addr.port());
let mut addr_packet = control_packet::AddrPacket::new(packet.payload_mut())?; self.client_cipher.encrypt_ipv4(&mut packet)?;
addr_packet.set_ipv4(ipv4); context.send_by_key(packet.buffer(), route_key).await?;
addr_packet.set_port(route_key.addr.port());
self.client_cipher.encrypt_ipv4(&mut packet)?;
context.send_by_key(packet.buffer(), route_key).await?;
}
std::net::IpAddr::V6(_) => {}
} }
} std::net::IpAddr::V6(_) => {}
},
ControlPacket::AddrResponse(addr_packet) => { ControlPacket::AddrResponse(addr_packet) => {
if !addr_packet.ipv4().is_multicast() if !addr_packet.ipv4().is_multicast()
&& !addr_packet.ipv4().is_broadcast() && !addr_packet.ipv4().is_broadcast()
&& !addr_packet.ipv4().is_unspecified() && !addr_packet.ipv4().is_unspecified()
&& !addr_packet.ipv4().is_loopback() && !addr_packet.ipv4().is_loopback()
&& !addr_packet.ipv4().is_private() && addr_packet.port() != 0 { && !addr_packet.ipv4().is_private()
self.nat_test.update_addr(addr_packet.ipv4(), addr_packet.port()) && addr_packet.port() != 0
{
self.nat_test
.update_addr(addr_packet.ipv4(), addr_packet.port())
} }
} }
} }
Ok(()) Ok(())
} }
async fn other_turn(&self, context: &Context, current_device: CurrentDeviceInfo, source: Ipv4Addr, net_packet: NetPacket<&mut [u8]>, route_key: &RouteKey) -> crate::Result<()> { async fn other_turn(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
source: Ipv4Addr,
net_packet: NetPacket<&mut [u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
if self.relay { if self.relay {
return Ok(()); return Ok(());
} }
match other_turn_packet::Protocol::from(net_packet.transport_protocol()) { match other_turn_packet::Protocol::from(net_packet.transport_protocol()) {
other_turn_packet::Protocol::Punch => { other_turn_packet::Protocol::Punch => {
let punch_info = PunchInfo::parse_from_bytes(net_packet.payload())?; let punch_info = PunchInfo::parse_from_bytes(net_packet.payload())?;
let public_ips = punch_info.public_ip_list. let public_ips = punch_info
iter().map(|v| { Ipv4Addr::from(v.to_be_bytes()) }).collect(); .public_ip_list
let peer_nat_info = NatInfo::new(public_ips, .iter()
punch_info.public_port as u16, .map(|v| Ipv4Addr::from(v.to_be_bytes()))
punch_info.public_port_range as u16, .collect();
Ipv4Addr::from(punch_info.local_ip.to_be_bytes()), let local_ipv4_addr = SocketAddrV4::new(
punch_info.local_port as u16, Ipv4Addr::from(punch_info.local_ip.to_be_bytes()),
punch_info.nat_type.enum_value_or_default().into()); punch_info.local_port as u16,
);
let ipv6_addr = if punch_info.ipv6.len() == 16 {
let ipv6: [u8; 16] = punch_info.ipv6.try_into().unwrap();
SocketAddrV6::new(Ipv6Addr::from(ipv6), punch_info.ipv6_port as u16, 0, 0)
} else {
SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, 0, 0, 0)
};
let peer_nat_info = NatInfo::new(
public_ips,
punch_info.public_port as u16,
punch_info.public_port_range as u16,
local_ipv4_addr,
ipv6_addr,
punch_info.nat_type.enum_value_or_default().into(),
);
self.peer_nat_info_map.insert(source, peer_nat_info.clone()); self.peer_nat_info_map.insert(source, peer_nat_info.clone());
if !punch_info.reply { if !punch_info.reply {
let mut punch_reply = PunchInfo::new(); let mut punch_reply = PunchInfo::new();
punch_reply.reply = true; punch_reply.reply = true;
let nat_info = self.nat_test.nat_info(); let nat_info = self.nat_test.nat_info();
punch_reply.public_ip_list = nat_info.public_ips.iter().map(|ip| u32::from_be_bytes(ip.octets())).collect(); punch_reply.public_ip_list = nat_info
.public_ips
.iter()
.map(|ip| u32::from_be_bytes(ip.octets()))
.collect();
punch_reply.public_port = nat_info.public_port as u32; punch_reply.public_port = nat_info.public_port as u32;
punch_reply.public_port_range = nat_info.public_port_range as u32; punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.nat_type = punch_reply.nat_type =
protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type)); protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
punch_reply.local_ip = u32::from_be_bytes(nat_info.local_ip.octets()); punch_reply.local_ip =
punch_reply.local_port = nat_info.local_port as u32; u32::from_be_bytes(nat_info.local_ipv4_addr.ip().octets());
punch_reply.local_port = nat_info.local_ipv4_addr.port() as u32;
if !nat_info.ipv6_addr.ip().is_unspecified() {
punch_reply.ipv6 = nat_info.ipv6_addr.ip().octets().to_vec();
punch_reply.ipv6_port = nat_info.ipv6_addr.port() as u32;
}
let bytes = punch_reply.write_to_bytes()?; let bytes = punch_reply.write_to_bytes()?;
let mut punch_packet = let mut punch_packet =
NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?; NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
punch_packet.set_version(Version::V1); punch_packet.set_version(Version::V1);
punch_packet.set_protocol(Protocol::OtherTurn); punch_packet.set_protocol(Protocol::OtherTurn);
punch_packet.set_transport_protocol( punch_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into());
other_turn_packet::Protocol::Punch.into(),
);
punch_packet.first_set_ttl(MAX_TTL); punch_packet.first_set_ttl(MAX_TTL);
punch_packet.set_source(current_device.virtual_ip()); punch_packet.set_source(current_device.virtual_ip());
punch_packet.set_destination(source); punch_packet.set_destination(source);
punch_packet.set_payload(&bytes)?; punch_packet.set_payload(&bytes)?;
if !peer_nat_info.local_ip.is_unspecified() && peer_nat_info.local_port != 0 { // if !peer_nat_info.local_ip.is_unspecified() && peer_nat_info.local_port != 0 {
let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED])?; // let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED])?;
packet.set_version(Version::V1); // packet.set_version(Version::V1);
packet.first_set_ttl(1); // packet.first_set_ttl(1);
packet.set_protocol(Protocol::Control); // packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::PunchRequest.into()); // packet.set_transport_protocol(control_packet::Protocol::PunchRequest.into());
packet.set_source(current_device.virtual_ip()); // packet.set_source(current_device.virtual_ip());
packet.set_destination(source); // packet.set_destination(source);
self.client_cipher.encrypt_ipv4(&mut packet)?; // self.client_cipher.encrypt_ipv4(&mut packet)?;
let _ = context.send_main(packet.buffer(), SocketAddr::V4(SocketAddrV4::new(peer_nat_info.local_ip, peer_nat_info.local_port))).await; // let _ = context.try_send_main_udp(packet.buffer(),
} // SocketAddr::V4(SocketAddrV4::new(peer_nat_info.local_ip, peer_nat_info.local_port)));
// }
if self.punch(source, peer_nat_info).await { if self.punch(source, peer_nat_info).await {
self.client_cipher.encrypt_ipv4(&mut punch_packet)?; self.client_cipher.encrypt_ipv4(&mut punch_packet)?;
context.send_by_key(punch_packet.buffer(), route_key).await?; context
.send_by_key(punch_packet.buffer(), route_key)
.await?;
} }
} else { } else {
self.punch(source, peer_nat_info).await; self.punch(source, peer_nat_info).await;
} }
} }
other_turn_packet::Protocol::Unknown(e) => { other_turn_packet::Protocol::Unknown(e) => {
log::warn!("不支持的转发协议 {:?},source:{:?}",e,source); log::warn!("不支持的转发协议 {:?},source:{:?}", e, source);
} }
} }
Ok(()) Ok(())
} }
async fn punch(&self, peer_ip: Ipv4Addr, peer_nat_info: NatInfo) -> bool { async fn punch(&self, peer_ip: Ipv4Addr, peer_nat_info: NatInfo) -> bool {
match peer_nat_info.nat_type { match peer_nat_info.nat_type {
NatType::Symmetric => { NatType::Symmetric => self
self.symmetric_sender.try_send((peer_ip, peer_nat_info)).is_ok() .symmetric_sender
} .try_send((peer_ip, peer_nat_info))
NatType::Cone => { .is_ok(),
self.cone_sender.try_send((peer_ip, peer_nat_info)).is_ok() NatType::Cone => self.cone_sender.try_send((peer_ip, peer_nat_info)).is_ok(),
}
} }
} }
} }
/// 处理服务端数据 /// 处理服务端数据
impl ChannelDataHandler { impl ChannelDataHandler {
async fn server_packet_handle(&self, context: &Context, current_device: CurrentDeviceInfo, buf: &mut [u8], data_len: usize, route_key: &RouteKey) -> crate::Result<()> { async fn server_packet_handle(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
buf: &mut [u8],
data_len: usize,
route_key: &RouteKey,
) -> crate::Result<()> {
let net_packet = NetPacket::new0(data_len, &buf[14..])?; let net_packet = NetPacket::new0(data_len, &buf[14..])?;
let source = net_packet.source(); let source = net_packet.source();
match net_packet.protocol() { match net_packet.protocol() {
Protocol::Service => { Protocol::Service => {
self.service(context, current_device, net_packet, route_key).await?; self.service(context, current_device, net_packet, route_key)
.await?;
} }
Protocol::Error => { Protocol::Error => {
self.error(context, current_device, source, net_packet, route_key).await?; self.error(context, current_device, source, net_packet, route_key)
.await?;
} }
Protocol::Control => { Protocol::Control => {
self.control_gateway(context, current_device, net_packet, route_key).await?; self.control_gateway(context, current_device, net_packet, route_key)
.await?;
} }
Protocol::IpTurn => { Protocol::IpTurn => {
match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) { match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) {
@@ -484,14 +617,29 @@ impl ChannelDataHandler {
} }
return Ok(()); return Ok(());
} }
async fn control_gateway(&self, context: &Context, current_device: CurrentDeviceInfo, net_packet: NetPacket<&[u8]>, route_key: &RouteKey) -> crate::Result<()> { async fn control_gateway(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
net_packet: NetPacket<&[u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
if net_packet.source() != current_device.virtual_gateway { if net_packet.source() != current_device.virtual_gateway {
return Ok(()); return Ok(());
} }
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? { match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PongPacket(pong_packet) => { ControlPacket::PongPacket(pong_packet) => {
let metric = net_packet.source_ttl() - net_packet.ttl() + 1; let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
self.pong_packet(true, metric, context, current_device, net_packet.source(), pong_packet, route_key).await?; self.pong_packet(
true,
metric,
context,
current_device,
net_packet.source(),
pong_packet,
route_key,
)
.await?;
} }
ControlPacket::AddrResponse(addr_packet) => { ControlPacket::AddrResponse(addr_packet) => {
if addr_packet.port() != 0 if addr_packet.port() != 0
@@ -499,42 +647,78 @@ impl ChannelDataHandler {
&& !addr_packet.ipv4().is_broadcast() && !addr_packet.ipv4().is_broadcast()
&& !addr_packet.ipv4().is_unspecified() && !addr_packet.ipv4().is_unspecified()
&& !addr_packet.ipv4().is_loopback() && !addr_packet.ipv4().is_loopback()
&& !addr_packet.ipv4().is_private() { && !addr_packet.ipv4().is_private()
self.nat_test.update_addr(addr_packet.ipv4(), addr_packet.port()) {
self.nat_test
.update_addr(addr_packet.ipv4(), addr_packet.port())
} }
} }
_ => {} _ => {}
} }
Ok(()) Ok(())
} }
async fn service(&self, context: &Context, current_device: CurrentDeviceInfo, net_packet: NetPacket<&[u8]>, route_key: &RouteKey) -> crate::Result<()> { async fn service(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
net_packet: NetPacket<&[u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
match service_packet::Protocol::from(net_packet.transport_protocol()) { match service_packet::Protocol::from(net_packet.transport_protocol()) {
service_packet::Protocol::RegistrationRequest => {} service_packet::Protocol::RegistrationRequest => {}
service_packet::Protocol::RegistrationResponse => { service_packet::Protocol::RegistrationResponse => {
let response = RegistrationResponse::parse_from_bytes(net_packet.payload())?; let response = RegistrationResponse::parse_from_bytes(net_packet.payload())?;
let local_port = context.main_local_port()?;
let local_ip = nat::local_ip()?; {
let nat_info = self.nat_test.re_test(Ipv4Addr::from(response.public_ip), let context = context.clone();
response.public_port as u16, local_ip, local_port).await; let nat_test = self.nat_test.clone();
context.switch(nat_info.nat_type); tokio::spawn(async move {
let local_port = context.main_local_ipv4_port().unwrap_or(0);
let local_ipv4_addr = nat::local_ipv4_addr(local_port);
let local_port = context.main_local_ipv6_port().unwrap_or(0);
let ipv6_addr = nat::local_ipv6_addr(local_port);
let nat_info = nat_test
.re_test(
Ipv4Addr::from(response.public_ip),
response.public_port as u16,
local_ipv4_addr,
ipv6_addr,
)
.await;
context.switch(nat_info.nat_type);
});
}
let new_ip = Ipv4Addr::from(response.virtual_ip); let new_ip = Ipv4Addr::from(response.virtual_ip);
let current_ip = current_device.virtual_ip(); let current_ip = current_device.virtual_ip();
if current_ip != new_ip { if current_ip != new_ip {
// ip发生变化 // ip发生变化
log::info!("ip发生变化,old_ip:{:?},new_ip:{:?}",current_ip,new_ip); log::info!("ip发生变化,old_ip:{:?},new_ip:{:?}", current_ip, new_ip);
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
let old_netmask = current_device.virtual_netmask; let old_netmask = current_device.virtual_netmask;
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
let old_gateway = current_device.virtual_gateway(); let old_gateway = current_device.virtual_gateway();
let virtual_ip = Ipv4Addr::from(response.virtual_ip); let virtual_ip = Ipv4Addr::from(response.virtual_ip);
let virtual_gateway = Ipv4Addr::from(response.virtual_gateway); let virtual_gateway = Ipv4Addr::from(response.virtual_gateway);
let virtual_netmask = Ipv4Addr::from(response.virtual_netmask); let virtual_netmask = Ipv4Addr::from(response.virtual_netmask);
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
self.device_writer.change_ip(virtual_ip, virtual_netmask, virtual_gateway, old_netmask, old_gateway)?; self.device_writer.change_ip(
let new_current_device = CurrentDeviceInfo::new(virtual_ip, virtual_gateway, virtual_ip,
virtual_netmask, current_device.connect_server); virtual_netmask,
if let Err(e) = self.current_device.compare_exchange(current_device, new_current_device) { virtual_gateway,
log::warn!("替换失败:{:?}",e); old_netmask,
old_gateway,
)?;
let new_current_device = CurrentDeviceInfo::new(
virtual_ip,
virtual_gateway,
virtual_netmask,
current_device.connect_server,
);
if let Err(e) = self
.current_device
.compare_exchange(current_device, new_current_device)
{
log::warn!("替换失败:{:?}", e);
} }
} }
self.connect_status.store(ConnectStatus::Connected); self.connect_status.store(ConnectStatus::Connected);
@@ -554,7 +738,7 @@ impl ChannelDataHandler {
) )
}) })
.collect(); .collect();
let route = Route::from(*route_key, 2, 99); let route = Route::from(*route_key, 2, 199);
for x in &ip_list { for x in &ip_list {
if x.status == PeerDeviceStatus::Online { if x.status == PeerDeviceStatus::Online {
context.add_route_if_absent(x.virtual_ip, route); context.add_route_if_absent(x.virtual_ip, route);
@@ -567,14 +751,21 @@ impl ChannelDataHandler {
} }
} }
service_packet::Protocol::Unknown(u) => { service_packet::Protocol::Unknown(u) => {
log::warn!("未知服务协议:{}",u); log::warn!("未知服务协议:{}", u);
} }
_ => {} _ => {}
} }
Ok(()) Ok(())
} }
async fn error(&self, _context: &Context, current_device: CurrentDeviceInfo, _source: Ipv4Addr, net_packet: NetPacket<&[u8]>, _route_key: &RouteKey) -> crate::Result<()> { async fn error(
log::info!("current_device:{:?}",current_device); &self,
_context: &Context,
current_device: CurrentDeviceInfo,
_source: Ipv4Addr,
net_packet: NetPacket<&[u8]>,
_route_key: &RouteKey,
) -> crate::Result<()> {
log::info!("current_device:{:?}", current_device);
match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload())? { match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
InErrorPacket::TokenError => { InErrorPacket::TokenError => {
return Err(Error::Stop("Token error".to_string())); return Err(Error::Stop("Token error".to_string()));
@@ -587,7 +778,9 @@ impl ChannelDataHandler {
} }
self.connect_status.store(ConnectStatus::Connecting); self.connect_status.store(ConnectStatus::Connecting);
self.register.fast_register(current_device.virtual_ip).await?; self.register
.fast_register(current_device.virtual_ip)
.await?;
} }
InErrorPacket::AddressExhausted => { InErrorPacket::AddressExhausted => {
//地址用尽 //地址用尽
@@ -606,4 +799,4 @@ impl ChannelDataHandler {
} }
Ok(()) Ok(())
} }
} }
+41 -48
View File
@@ -1,18 +1,18 @@
use crossbeam_utils::atomic::AtomicCell;
use std::net::{Ipv4Addr, SocketAddr}; use std::net::{Ipv4Addr, SocketAddr};
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use crossbeam_utils::atomic::AtomicCell;
use protobuf::Message;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpStream, UdpSocket};
use crate::channel::sender::ChannelSender; use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher; use crate::cipher::Cipher;
use crate::handle::PeerDeviceInfo; use crate::handle::PeerDeviceInfo;
use protobuf::Message;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpStream, UdpSocket};
use crate::proto::message::{RegistrationRequest, RegistrationResponse}; use crate::proto::message::{RegistrationRequest, RegistrationResponse};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::error_packet::InErrorPacket; use crate::protocol::error_packet::InErrorPacket;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL}; use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
use crate::protocol::body::ENCRYPTION_RESERVED;
pub enum ReqEnum { pub enum ReqEnum {
TokenError, TokenError,
@@ -47,8 +47,17 @@ pub async fn registration(
ip: Ipv4Addr, ip: Ipv4Addr,
client_secret: bool, client_secret: bool,
) -> Result<RegResponse, ReqEnum> { ) -> Result<RegResponse, ReqEnum> {
let request_packet = let request_packet = registration_request_packet(
registration_request_packet(server_cipher, token.clone(), device_id.clone(), name.clone(), ip, false, false, client_secret).unwrap(); server_cipher,
token.clone(),
device_id.clone(),
name.clone(),
ip,
false,
false,
client_secret,
)
.unwrap();
let buf = request_packet.buffer(); let buf = request_packet.buffer();
let mut recv_buf = [0u8; 10240]; let mut recv_buf = [0u8; 10240];
let recv_buf = if let Some(main_tcp_channel) = main_tcp_channel { let recv_buf = if let Some(main_tcp_channel) = main_tcp_channel {
@@ -75,29 +84,30 @@ pub async fn registration(
if let Err(e) = main_channel.send_to(buf, server_address).await { if let Err(e) = main_channel.send_to(buf, server_address).await {
return Err(ReqEnum::Other(format!("send error:{}", e))); return Err(ReqEnum::Other(format!("send error:{}", e)));
} }
match tokio::time::timeout(Duration::from_millis(300), main_channel.recv_from(&mut recv_buf)).await { match tokio::time::timeout(
Ok(rs) => { Duration::from_millis(300),
match rs { main_channel.recv_from(&mut recv_buf),
Ok((len, addr)) => { )
if server_address != addr { .await
return Err(ReqEnum::Other(format!("invalid data,from {}", addr))); {
} Ok(rs) => match rs {
&mut recv_buf[..len] Ok((len, addr)) => {
} if server_address != addr {
Err(e) => { return Err(ReqEnum::Other(format!("invalid data,from {}", addr)));
return Err(ReqEnum::Other(format!("receiver error:{}", e)));
} }
&mut recv_buf[..len]
} }
} Err(e) => {
return Err(ReqEnum::Other(format!("receiver error:{}", e)));
}
},
Err(_) => { Err(_) => {
return Err(ReqEnum::Timeout); return Err(ReqEnum::Timeout);
} }
} }
}; };
let mut net_packet = match NetPacket::new(recv_buf) { let mut net_packet = match NetPacket::new(recv_buf) {
Ok(net_packet) => { Ok(net_packet) => net_packet,
net_packet
}
Err(e) => { Err(e) => {
return Err(ReqEnum::ServerError(format!("{}", e))); return Err(ReqEnum::ServerError(format!("{}", e)));
} }
@@ -133,14 +143,10 @@ pub async fn registration(
public_port: response.public_port as u16, public_port: response.public_port as u16,
}) })
} }
Err(_) => { Err(_) => Err(ReqEnum::ServerError("invalid data".to_string())),
Err(ReqEnum::ServerError("invalid data".to_string()))
}
} }
} }
_ => { _ => Err(ReqEnum::ServerError("invalid data".to_string())),
Err(ReqEnum::ServerError("invalid data".to_string()))
}
} }
} }
Protocol::Error => { Protocol::Error => {
@@ -150,24 +156,14 @@ pub async fn registration(
InErrorPacket::Disconnect => { InErrorPacket::Disconnect => {
Err(ReqEnum::ServerError("disconnect".to_string())) Err(ReqEnum::ServerError("disconnect".to_string()))
} }
InErrorPacket::AddressExhausted => { InErrorPacket::AddressExhausted => Err(ReqEnum::AddressExhausted),
Err(ReqEnum::AddressExhausted)
}
InErrorPacket::OtherError(e) => match e.message() { InErrorPacket::OtherError(e) => match e.message() {
Ok(str) => { Ok(str) => Err(ReqEnum::ServerError(str)),
Err(ReqEnum::ServerError(str))
}
Err(e) => Err(ReqEnum::Other(format!("{}", e))), Err(e) => Err(ReqEnum::Other(format!("{}", e))),
}, },
InErrorPacket::IpAlreadyExists => { InErrorPacket::IpAlreadyExists => Err(ReqEnum::IpAlreadyExists),
Err(ReqEnum::IpAlreadyExists) InErrorPacket::InvalidIp => Err(ReqEnum::InvalidIp),
} InErrorPacket::NoKey => Err(ReqEnum::ServerError("no key".to_string())),
InErrorPacket::InvalidIp => {
Err(ReqEnum::InvalidIp)
}
InErrorPacket::NoKey => {
Err(ReqEnum::ServerError("no key".to_string()))
}
}, },
Err(e) => Err(ReqEnum::Other(format!("{}", e))), Err(e) => Err(ReqEnum::Other(format!("{}", e))),
} }
@@ -193,7 +189,7 @@ fn registration_request_packet(
request.virtual_ip = ip.into(); request.virtual_ip = ip.into();
request.allow_ip_change = allow_ip_change; request.allow_ip_change = allow_ip_change;
request.is_fast = is_fast; request.is_fast = is_fast;
request.version = "1.2.0".to_string(); request.version = crate::VNT_VERSION.to_string();
request.client_secret = client_secret; request.client_secret = client_secret;
let bytes = request.write_to_bytes()?; let bytes = request.write_to_bytes()?;
let buf = vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED]; let buf = vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED];
@@ -243,10 +239,7 @@ impl Register {
pub async fn fast_register(&self, ip: Ipv4Addr) -> crate::Result<()> { pub async fn fast_register(&self, ip: Ipv4Addr) -> crate::Result<()> {
let last = self.time.load(); let last = self.time.load();
if last.elapsed() < Duration::from_secs(2) if last.elapsed() < Duration::from_secs(2)
|| self || self.time.compare_exchange(last, Instant::now()).is_err()
.time
.compare_exchange(last, Instant::now())
.is_err()
{ {
//短时间不重复注册 //短时间不重复注册
return Ok(()); return Ok(());
+11 -4
View File
@@ -1,7 +1,10 @@
use byte_pool::Block; use byte_pool::Block;
#[derive(Clone)] #[derive(Clone)]
pub struct BufSenderGroup(usize, Vec<tokio::sync::mpsc::Sender<(Block<'static>, usize, usize)>>); pub struct BufSenderGroup(
usize,
Vec<tokio::sync::mpsc::Sender<(Block<'static>, usize, usize)>>,
);
pub struct BufReceiverGroup(pub Vec<tokio::sync::mpsc::Receiver<(Block<'static>, usize, usize)>>); pub struct BufReceiverGroup(pub Vec<tokio::sync::mpsc::Receiver<(Block<'static>, usize, usize)>>);
@@ -17,9 +20,13 @@ pub fn buf_channel_group(size: usize) -> (BufSenderGroup, BufReceiverGroup) {
let mut buf_sender_group = Vec::with_capacity(size); let mut buf_sender_group = Vec::with_capacity(size);
let mut buf_receiver_group = Vec::with_capacity(size); let mut buf_receiver_group = Vec::with_capacity(size);
for _ in 0..size { for _ in 0..size {
let (buf_sender, buf_receiver) = tokio::sync::mpsc::channel::<(Block<'static>, usize, usize)>(10); let (buf_sender, buf_receiver) =
tokio::sync::mpsc::channel::<(Block<'static>, usize, usize)>(10);
buf_sender_group.push(buf_sender); buf_sender_group.push(buf_sender);
buf_receiver_group.push(buf_receiver); buf_receiver_group.push(buf_receiver);
} }
(BufSenderGroup(0, buf_sender_group), BufReceiverGroup(buf_receiver_group)) (
} BufSenderGroup(0, buf_sender_group),
BufReceiverGroup(buf_receiver_group),
)
}
+109 -48
View File
@@ -1,28 +1,34 @@
use std::io; use crate::channel::sender::ChannelSender;
use std::net::SocketAddrV4; use crate::cipher::Cipher;
use std::sync::Arc; use crate::error::*;
use parking_lot::RwLock; use crate::external_route::ExternalRoute;
use crate::handle::{check_dest, CurrentDeviceInfo};
use crate::igmp_server::{IgmpServer, Multicast};
use crate::ip_proxy::IpProxyMap;
use crate::protocol;
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::ip_turn_packet::BroadcastPacket;
use crate::protocol::{ip_turn_packet, NetPacket, Version, MAX_TTL};
use packet::ip::ipv4::packet::IpV4Packet; use packet::ip::ipv4::packet::IpV4Packet;
use packet::ip::ipv4::protocol::Protocol; use packet::ip::ipv4::protocol::Protocol;
use packet::tcp::tcp::TcpPacket; use packet::tcp::tcp::TcpPacket;
use packet::udp::udp::UdpPacket; use packet::udp::udp::UdpPacket;
use crate::channel::sender::ChannelSender; use parking_lot::RwLock;
use crate::cipher::Cipher; use std::io;
use crate::external_route::ExternalRoute; use std::net::{Ipv4Addr, SocketAddrV4};
use crate::handle::{check_dest, CurrentDeviceInfo}; use std::sync::Arc;
use crate::ip_proxy::IpProxyMap;
use crate::protocol::{ip_turn_packet, MAX_TTL, NetPacket, Version};
use crate::error::*;
use crate::igmp_server::{IgmpServer, Multicast};
use crate::protocol;
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::ip_turn_packet::BroadcastPacket;
pub mod channel_group; pub mod channel_group;
pub mod tun_handler;
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
pub mod tap_handler; pub mod tap_handler;
pub mod tun_handler;
async fn broadcast(server_cipher: &Cipher, multicast_members: Option<Arc<RwLock<Multicast>>>, sender: &ChannelSender, net_packet: &mut NetPacket<&mut [u8]>, current_device: &CurrentDeviceInfo) -> Result<()> { async fn broadcast(
server_cipher: &Cipher,
multicast_members: Option<Arc<RwLock<Multicast>>>,
sender: &ChannelSender,
net_packet: &mut NetPacket<&mut [u8]>,
current_device: &CurrentDeviceInfo,
) -> Result<()> {
let mut peer_ips = Vec::with_capacity(8); let mut peer_ips = Vec::with_capacity(8);
let vec = sender.route_table_one(); let vec = sender.route_table_one();
let mut relay_count = 0; let mut relay_count = 0;
@@ -40,7 +46,11 @@ async fn broadcast(server_cipher: &Cipher, multicast_members: Option<Arc<RwLock<
} }
} }
if route.is_p2p() if route.is_p2p()
&& sender.send_by_key(net_packet.buffer(), &route.route_key()).await.is_ok() { && sender
.send_by_key(net_packet.buffer(), &route.route_key())
.await
.is_ok()
{
peer_ips.push(peer_ip); peer_ips.push(peer_ip);
} else { } else {
relay_count += 1; relay_count += 1;
@@ -52,16 +62,22 @@ async fn broadcast(server_cipher: &Cipher, multicast_members: Option<Arc<RwLock<
} }
//转发到服务端的可选择广播,还要进行服务端加密 //转发到服务端的可选择广播,还要进行服务端加密
if peer_ips.is_empty() { if peer_ips.is_empty() {
sender.send_main(net_packet.buffer(), current_device.connect_server).await?; sender
.send_main(net_packet.buffer(), current_device.connect_server)
.await?;
} else { } else {
let buf = vec![0 as u8; 12 + 1 + peer_ips.len() * 4 + net_packet.data_len() + ENCRYPTION_RESERVED]; let buf = vec![
0 as u8;
12 + 1 + peer_ips.len() * 4 + net_packet.data_len() + ENCRYPTION_RESERVED
];
//剩余的发送到服务端,需要告知哪些已发送过 //剩余的发送到服务端,需要告知哪些已发送过
let mut server_packet = NetPacket::new_encrypt(buf)?; let mut server_packet = NetPacket::new_encrypt(buf)?;
server_packet.set_version(Version::V1); server_packet.set_version(Version::V1);
server_packet.set_gateway_flag(true); server_packet.set_gateway_flag(true);
server_packet.first_set_ttl(MAX_TTL); server_packet.first_set_ttl(MAX_TTL);
server_packet.set_source(net_packet.source()); server_packet.set_source(net_packet.source());
server_packet.set_destination(current_device.virtual_gateway); //使用对应的目的地址
server_packet.set_destination(net_packet.destination());
server_packet.set_protocol(protocol::Protocol::IpTurn); server_packet.set_protocol(protocol::Protocol::IpTurn);
server_packet.set_transport_protocol(ip_turn_packet::Protocol::Ipv4Broadcast.into()); server_packet.set_transport_protocol(ip_turn_packet::Protocol::Ipv4Broadcast.into());
@@ -69,7 +85,9 @@ async fn broadcast(server_cipher: &Cipher, multicast_members: Option<Arc<RwLock<
broadcast.set_address(&peer_ips)?; broadcast.set_address(&peer_ips)?;
broadcast.set_data(net_packet.buffer())?; broadcast.set_data(net_packet.buffer())?;
server_cipher.encrypt_ipv4(&mut server_packet)?; server_cipher.encrypt_ipv4(&mut server_packet)?;
sender.send_main(server_packet.buffer(), current_device.connect_server).await?; sender
.send_main(server_packet.buffer(), current_device.connect_server)
.await?;
} }
Ok(()) Ok(())
} }
@@ -78,14 +96,17 @@ async fn broadcast(server_cipher: &Cipher, multicast_members: Option<Arc<RwLock<
/// |12字节开头|ip报文|至少1024字节结尾| /// |12字节开头|ip报文|至少1024字节结尾|
/// ///
#[inline] #[inline]
pub async fn base_handle(sender: &ChannelSender, buf: &mut [u8], pub async fn base_handle(
data_len: usize,//数据总长度=12+ip包长度 sender: &ChannelSender,
igmp_server: &Option<IgmpServer>, buf: &mut [u8],
current_device: CurrentDeviceInfo, data_len: usize, //数据总长度=12+ip包长度
ip_route: &Option<ExternalRoute>, igmp_server: &Option<IgmpServer>,
proxy_map: &Option<IpProxyMap>, current_device: CurrentDeviceInfo,
client_cipher: &Cipher, ip_route: &Option<ExternalRoute>,
server_cipher: &Cipher) -> Result<()> { proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> Result<()> {
let ipv4_packet = IpV4Packet::new(&buf[12..data_len])?; let ipv4_packet = IpV4Packet::new(&buf[12..data_len])?;
let protocol = ipv4_packet.protocol(); let protocol = ipv4_packet.protocol();
let ip_head_len = ipv4_packet.header_len() as usize * 4; let ip_head_len = ipv4_packet.header_len() as usize * 4;
@@ -105,7 +126,9 @@ pub async fn base_handle(sender: &ChannelSender, buf: &mut [u8],
if protocol == Protocol::Icmp { if protocol == Protocol::Icmp {
net_packet.set_gateway_flag(true); net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?; server_cipher.encrypt_ipv4(&mut net_packet)?;
sender.send_main(net_packet.buffer(), current_device.connect_server).await?; sender
.send_main(net_packet.buffer(), current_device.connect_server)
.await?;
} }
return Ok(()); return Ok(());
} }
@@ -117,17 +140,28 @@ pub async fn base_handle(sender: &ChannelSender, buf: &mut [u8],
net_packet.set_destination(current_device.virtual_gateway); net_packet.set_destination(current_device.virtual_gateway);
net_packet.set_gateway_flag(true); net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?; server_cipher.encrypt_ipv4(&mut net_packet)?;
sender.send_main(net_packet.buffer(), current_device.connect_server).await?; sender
.send_main(net_packet.buffer(), current_device.connect_server)
.await?;
} }
} }
Protocol::Udp => { Protocol::Udp => {
client_cipher.encrypt_ipv4(&mut net_packet)?;
let multicast_members = if let Some(igmp_server) = igmp_server { let multicast_members = if let Some(igmp_server) = igmp_server {
igmp_server.load(&dest_ip) igmp_server.load(&dest_ip)
} else { } else {
//当作广播处理
net_packet.set_destination(Ipv4Addr::BROADCAST);
None None
}; };
broadcast(server_cipher, multicast_members, sender, &mut net_packet, &current_device).await?; client_cipher.encrypt_ipv4(&mut net_packet)?;
broadcast(
server_cipher,
multicast_members,
sender,
&mut net_packet,
&current_device,
)
.await?;
} }
_ => {} _ => {}
} }
@@ -136,10 +170,21 @@ pub async fn base_handle(sender: &ChannelSender, buf: &mut [u8],
if dest_ip.is_broadcast() || current_device.broadcast_address == dest_ip { if dest_ip.is_broadcast() || current_device.broadcast_address == dest_ip {
// 广播 发送到直连目标 // 广播 发送到直连目标
client_cipher.encrypt_ipv4(&mut net_packet)?; client_cipher.encrypt_ipv4(&mut net_packet)?;
broadcast(server_cipher, None, sender, &mut net_packet, &current_device).await?; broadcast(
server_cipher,
None,
sender,
&mut net_packet,
&current_device,
)
.await?;
return Ok(()); return Ok(());
} }
if !check_dest(dest_ip, current_device.virtual_netmask, current_device.virtual_network) { if !check_dest(
dest_ip,
current_device.virtual_netmask,
current_device.virtual_network,
) {
if let Some(ip_route) = ip_route { if let Some(ip_route) = ip_route {
if let Some(r_dest_ip) = ip_route.route(&dest_ip) { if let Some(r_dest_ip) = ip_route.route(&dest_ip) {
//路由的目标不能是自己 //路由的目标不能是自己
@@ -159,35 +204,45 @@ pub async fn base_handle(sender: &ChannelSender, buf: &mut [u8],
match protocol { match protocol {
Protocol::Tcp => { Protocol::Tcp => {
let dest_addr = { let dest_addr = {
let tcp_packet = TcpPacket::new(src_ip, dest_ip, let tcp_packet = TcpPacket::new(
&mut net_packet.payload_mut()[ip_head_len..])?; src_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
SocketAddrV4::new(dest_ip, tcp_packet.destination_port()) SocketAddrV4::new(dest_ip, tcp_packet.destination_port())
}; };
if let Some(entry) = proxy_map.tcp_proxy_map.get(&dest_addr) { if let Some(entry) = proxy_map.tcp_proxy_map.get(&dest_addr) {
let source_addr = entry.value(); let source_addr = entry.value();
let source_ip = *source_addr.ip(); let source_ip = *source_addr.ip();
let mut tcp_packet = TcpPacket::new(source_ip, dest_ip, let mut tcp_packet = TcpPacket::new(
&mut net_packet.payload_mut()[ip_head_len..])?; source_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
tcp_packet.set_source_port(source_addr.port()); tcp_packet.set_source_port(source_addr.port());
tcp_packet.update_checksum(); tcp_packet.update_checksum();
let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?; let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?;
ipv4_packet.set_source_ip(source_ip); ipv4_packet.set_source_ip(source_ip);
ipv4_packet.update_checksum(); ipv4_packet.update_checksum();
}else{
log::warn!("不存在接口 {:?}",dest_addr);
} }
} }
Protocol::Udp => { Protocol::Udp => {
let dest_addr = { let dest_addr = {
let udp_packet = UdpPacket::new(src_ip, dest_ip, let udp_packet = UdpPacket::new(
&mut net_packet.payload_mut()[ip_head_len..])?; src_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
SocketAddrV4::new(dest_ip, udp_packet.destination_port()) SocketAddrV4::new(dest_ip, udp_packet.destination_port())
}; };
if let Some(entry) = proxy_map.udp_proxy_map.get(&dest_addr) { if let Some(entry) = proxy_map.udp_proxy_map.get(&dest_addr) {
let source_addr = entry.value(); let source_addr = entry.value();
let source_ip = *source_addr.ip(); let source_ip = *source_addr.ip();
let mut udp_packet = UdpPacket::new(source_ip, dest_ip, let mut udp_packet = UdpPacket::new(
&mut net_packet.payload_mut()[ip_head_len..])?; source_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
udp_packet.set_source_port(source_addr.port()); udp_packet.set_source_port(source_addr.port());
udp_packet.update_checksum(); udp_packet.update_checksum();
let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?; let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?;
@@ -200,8 +255,14 @@ pub async fn base_handle(sender: &ChannelSender, buf: &mut [u8],
} }
client_cipher.encrypt_ipv4(&mut net_packet)?; client_cipher.encrypt_ipv4(&mut net_packet)?;
//优先发到直连到地址 //优先发到直连到地址
if sender.send_by_id(net_packet.buffer(), &dest_ip).await.is_err() { if sender
sender.send_main(net_packet.buffer(), current_device.connect_server).await?; .send_by_id(net_packet.buffer(), &dest_ip)
.await
.is_err()
{
sender
.send_main(net_packet.buffer(), current_device.connect_server)
.await?;
} }
return Ok(()); return Ok(());
} }
+150 -59
View File
@@ -1,6 +1,6 @@
use std::{io, thread};
use std::sync::Arc;
use byte_pool::BytePool; use byte_pool::BytePool;
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
use lazy_static::lazy_static; use lazy_static::lazy_static;
@@ -17,36 +17,56 @@ use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher; use crate::cipher::Cipher;
use crate::core::status::VntWorker; use crate::core::status::VntWorker;
use crate::external_route::ExternalRoute; use crate::external_route::ExternalRoute;
use crate::handle::CurrentDeviceInfo;
use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup}; use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup};
use crate::handle::CurrentDeviceInfo;
use crate::igmp_server::IgmpServer; use crate::igmp_server::IgmpServer;
use crate::ip_proxy::IpProxyMap; use crate::ip_proxy::IpProxyMap;
use crate::tun_tap_device::{DeviceReader, DeviceWriter}; use crate::tun_tap_device::{DeviceReader, DeviceWriter};
lazy_static! { lazy_static! {
static ref POOL:BytePool<Vec<u8>> = BytePool::<Vec<u8>>::new(); static ref POOL: BytePool<Vec<u8>> = BytePool::<Vec<u8>>::new();
} }
pub fn start(worker: VntWorker, sender: ChannelSender, pub fn start(
device_reader: DeviceReader, worker: VntWorker,
device_writer: DeviceWriter, sender: ChannelSender,
igmp_server: Option<IgmpServer>, device_reader: DeviceReader,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>, device_writer: DeviceWriter,
ip_route: Option<ExternalRoute>, igmp_server: Option<IgmpServer>,
ip_proxy_map: Option<IpProxyMap>, current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher, server_cipher: Cipher, parallel: usize) { ip_route: Option<ExternalRoute>,
ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
parallel: usize,
) {
if parallel == 1 { if parallel == 1 {
thread::Builder::new().name("tap_handler".into()).spawn(move || { thread::Builder::new()
tokio::runtime::Builder::new_current_thread() .name("tap_handler".into())
.enable_all().build().unwrap() .spawn(move || {
.block_on(async move { tokio::runtime::Builder::new_current_thread()
if let Err(e) = start_simple(sender, device_reader, .enable_all()
device_writer, igmp_server, .build()
current_device, ip_route, ip_proxy_map, client_cipher, server_cipher).await { .unwrap()
log::warn!("tap:{:?}",e); .block_on(async move {
} if let Err(e) = start_simple(
worker.stop_all(); sender,
}); device_reader,
}).unwrap(); device_writer,
igmp_server,
current_device,
ip_route,
ip_proxy_map,
client_cipher,
server_cipher,
)
.await
{
log::warn!("tap:{:?}", e);
}
worker.stop_all();
});
})
.unwrap();
} else { } else {
let (buf_sender, buf_receiver) = buf_channel_group(parallel); let (buf_sender, buf_receiver) = buf_channel_group(parallel);
for mut buf_receiver in buf_receiver.0 { for mut buf_receiver in buf_receiver.0 {
@@ -60,8 +80,20 @@ pub fn start(worker: VntWorker, sender: ChannelSender,
let server_cipher = server_cipher.clone(); let server_cipher = server_cipher.clone();
tokio::spawn(async move { tokio::spawn(async move {
while let Some((mut buf, _, len)) = buf_receiver.recv().await { while let Some((mut buf, _, len)) = buf_receiver.recv().await {
match handle(&mut buf, len, &igmp_server, &current_device, &device_writer, &sender, match handle(
&ip_route, &ip_proxy_map, &client_cipher, &server_cipher).await { &mut buf,
len,
&igmp_server,
&current_device,
&device_writer,
&sender,
&ip_route,
&ip_proxy_map,
&client_cipher,
&server_cipher,
)
.await
{
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
log::warn!("{:?}", e) log::warn!("{:?}", e)
@@ -70,22 +102,29 @@ pub fn start(worker: VntWorker, sender: ChannelSender,
} }
}); });
} }
thread::Builder::new().name("tap_handler".into()).spawn(move || { thread::Builder::new()
tokio::runtime::Builder::new_current_thread() .name("tap_handler".into())
.enable_all().build().unwrap() .spawn(move || {
.block_on(async move { tokio::runtime::Builder::new_current_thread()
if let Err(e) = start_(sender, device_reader, buf_sender).await { .enable_all()
log::warn!("tap:{:?}",e); .build()
} .unwrap()
worker.stop_all(); .block_on(async move {
}); if let Err(e) = start_(sender, device_reader, buf_sender).await {
}).unwrap(); log::warn!("tap:{:?}", e);
}
worker.stop_all();
});
})
.unwrap();
} }
} }
async fn start_(sender: ChannelSender, async fn start_(
device_reader: DeviceReader, sender: ChannelSender,
mut buf_sender: BufSenderGroup) -> io::Result<()> { device_reader: DeviceReader,
mut buf_sender: BufSenderGroup,
) -> io::Result<()> {
loop { loop {
let mut buf = POOL.alloc(4096); let mut buf = POOL.alloc(4096);
if sender.is_close() { if sender.is_close() {
@@ -94,36 +133,65 @@ async fn start_(sender: ChannelSender,
let start = 0; let start = 0;
let len = device_reader.read(&mut buf)?; let len = device_reader.read(&mut buf)?;
if !buf_sender.send((buf, start, len)).await { if !buf_sender.send((buf, start, len)).await {
return Err(io::Error::new(io::ErrorKind::Other, "tap buf_sender发送失败")); return Err(io::Error::new(
io::ErrorKind::Other,
"tap buf_sender发送失败",
));
} }
} }
} }
async fn start_simple(sender: ChannelSender, async fn start_simple(
device_reader: DeviceReader, sender: ChannelSender,
device_writer: DeviceWriter, device_reader: DeviceReader,
igmp_server: Option<IgmpServer>, device_writer: DeviceWriter,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>, igmp_server: Option<IgmpServer>,
ip_route: Option<ExternalRoute>, current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_proxy_map: Option<IpProxyMap>, ip_route: Option<ExternalRoute>,
client_cipher: Cipher, server_cipher: Cipher) -> io::Result<()> { ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
) -> io::Result<()> {
let mut buf = [0; 4096]; let mut buf = [0; 4096];
loop { loop {
let len = device_reader.read(&mut buf)?; let len = device_reader.read(&mut buf)?;
if let Err(e) = handle(&mut buf, len, &igmp_server, &current_device, &device_writer, &sender, &ip_route, &ip_proxy_map, &client_cipher, &server_cipher).await { if let Err(e) = handle(
log::warn!("tap handle{:?}",e); &mut buf,
len,
&igmp_server,
&current_device,
&device_writer,
&sender,
&ip_route,
&ip_proxy_map,
&client_cipher,
&server_cipher,
)
.await
{
log::warn!("tap handle{:?}", e);
} }
} }
} }
async fn handle(buf: &mut [u8], len: usize, igmp_server: &Option<IgmpServer>, current_device: &AtomicCell<CurrentDeviceInfo>, async fn handle(
device_writer: &DeviceWriter, sender: &ChannelSender, ip_route: &Option<ExternalRoute>, buf: &mut [u8],
proxy_map: &Option<IpProxyMap>, client_cipher: &Cipher, server_cipher: &Cipher) -> crate::Result<()> { len: usize,
igmp_server: &Option<IgmpServer>,
current_device: &AtomicCell<CurrentDeviceInfo>,
device_writer: &DeviceWriter,
sender: &ChannelSender,
ip_route: &Option<ExternalRoute>,
proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> crate::Result<()> {
let mut ethernet_packet = EthernetPacket::new(&mut buf[..len])?; let mut ethernet_packet = EthernetPacket::new(&mut buf[..len])?;
let current_device = current_device.load(); let current_device = current_device.load();
match ethernet_packet.protocol() { match ethernet_packet.protocol() {
ethernet::protocol::Protocol::Arp => { ethernet::protocol::Protocol::Arp => {
let mut out_ethernet_packet = EthernetPacket::unchecked(ethernet_packet.buffer.to_vec()); let mut out_ethernet_packet =
EthernetPacket::unchecked(ethernet_packet.buffer.to_vec());
let arp_packet = ArpPacket::unchecked(ethernet_packet.payload()); let arp_packet = ArpPacket::unchecked(ethernet_packet.payload());
let mut out_arp_packet = ArpPacket::unchecked(out_ethernet_packet.payload_mut()); let mut out_arp_packet = ArpPacket::unchecked(out_ethernet_packet.payload_mut());
let sender_h = arp_packet.sender_hardware_addr(); let sender_h = arp_packet.sender_hardware_addr();
@@ -133,12 +201,26 @@ async fn handle(buf: &mut [u8], len: usize, igmp_server: &Option<IgmpServer>, cu
return Ok(()); return Ok(());
} }
//回复一个虚假的MAC地址 //回复一个虚假的MAC地址
out_arp_packet.set_sender_hardware_addr(&[target_p[0], target_p[1], target_p[2], target_p[3], !sender_h[5], 234]); out_arp_packet.set_sender_hardware_addr(&[
target_p[0],
target_p[1],
target_p[2],
target_p[3],
!sender_h[5],
234,
]);
out_arp_packet.set_sender_protocol_addr(target_p); out_arp_packet.set_sender_protocol_addr(target_p);
out_arp_packet.set_target_hardware_addr(sender_h); out_arp_packet.set_target_hardware_addr(sender_h);
out_arp_packet.set_target_protocol_addr(sender_p); out_arp_packet.set_target_protocol_addr(sender_p);
out_arp_packet.set_op_code(2); out_arp_packet.set_op_code(2);
out_ethernet_packet.set_source(&[target_p[0], target_p[1], target_p[2], target_p[3], !sender_h[5], 234]); out_ethernet_packet.set_source(&[
target_p[0],
target_p[1],
target_p[2],
target_p[3],
!sender_h[5],
234,
]);
out_ethernet_packet.set_destination(sender_h); out_ethernet_packet.set_destination(sender_h);
device_writer.write_ethernet_tap(&out_ethernet_packet.buffer)?; device_writer.write_ethernet_tap(&out_ethernet_packet.buffer)?;
} }
@@ -169,8 +251,18 @@ async fn handle(buf: &mut [u8], len: usize, igmp_server: &Option<IgmpServer>, cu
return Ok(()); return Ok(());
} }
// 以太网帧头部14字节,预留12字节 // 以太网帧头部14字节,预留12字节
return crate::handle::tun_tap::base_handle(sender, &mut buf[2..], len - 2, igmp_server, current_device, return crate::handle::tun_tap::base_handle(
ip_route, proxy_map, client_cipher, server_cipher).await; sender,
&mut buf[2..],
len - 2,
igmp_server,
current_device,
ip_route,
proxy_map,
client_cipher,
server_cipher,
)
.await;
} }
_ => { _ => {
// log::warn!("不支持的二层协议:{:?}",p) // log::warn!("不支持的二层协议:{:?}",p)
@@ -178,4 +270,3 @@ async fn handle(buf: &mut [u8], len: usize, igmp_server: &Option<IgmpServer>, cu
} }
Ok(()) Ok(())
} }
+138 -58
View File
@@ -1,21 +1,21 @@
use std::{io, thread};
use std::sync::Arc;
use byte_pool::BytePool; use byte_pool::BytePool;
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
use packet::icmp::Kind;
use packet::icmp::icmp::IcmpPacket;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::sender::ChannelSender; use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher; use crate::cipher::Cipher;
use crate::core::status::VntWorker; use crate::core::status::VntWorker;
use packet::icmp::icmp::IcmpPacket;
use packet::icmp::Kind;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::error::*; use crate::error::*;
use crate::external_route::ExternalRoute; use crate::external_route::ExternalRoute;
use crate::handle::CurrentDeviceInfo;
use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup}; use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup};
use crate::handle::CurrentDeviceInfo;
use crate::igmp_server::IgmpServer; use crate::igmp_server::IgmpServer;
use crate::ip_proxy::IpProxyMap; use crate::ip_proxy::IpProxyMap;
use crate::tun_tap_device::{DeviceReader, DeviceWriter}; use crate::tun_tap_device::{DeviceReader, DeviceWriter};
@@ -40,10 +40,18 @@ fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) ->
/// 接收tun数据,并且转发到udp上 /// 接收tun数据,并且转发到udp上
#[inline] #[inline]
async fn handle(sender: &ChannelSender, data: &mut [u8], len: usize, device_writer: &DeviceWriter, async fn handle(
igmp_server: &Option<IgmpServer>, current_device: CurrentDeviceInfo, sender: &ChannelSender,
ip_route: &Option<ExternalRoute>, proxy_map: &Option<IpProxyMap>, data: &mut [u8],
client_cipher: &Cipher, server_cipher: &Cipher) -> Result<()> { len: usize,
device_writer: &DeviceWriter,
igmp_server: &Option<IgmpServer>,
current_device: CurrentDeviceInfo,
ip_route: &Option<ExternalRoute>,
proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> Result<()> {
let ipv4_packet = if let Ok(ipv4_packet) = IpV4Packet::new(&mut data[12..len]) { let ipv4_packet = if let Ok(ipv4_packet) = IpV4Packet::new(&mut data[12..len]) {
ipv4_packet ipv4_packet
} else { } else {
@@ -57,30 +65,62 @@ async fn handle(sender: &ChannelSender, data: &mut [u8], len: usize, device_writ
if src_ip == dest_ip { if src_ip == dest_ip {
return icmp(&device_writer, ipv4_packet); return icmp(&device_writer, ipv4_packet);
} }
return crate::handle::tun_tap::base_handle(sender, data, len, igmp_server, return crate::handle::tun_tap::base_handle(
current_device, ip_route, proxy_map, client_cipher, server_cipher).await; sender,
data,
len,
igmp_server,
current_device,
ip_route,
proxy_map,
client_cipher,
server_cipher,
)
.await;
} }
pub async fn start(worker: VntWorker, sender: ChannelSender, pub async fn start(
device_reader: DeviceReader, worker: VntWorker,
device_writer: DeviceWriter, sender: ChannelSender,
igmp_server: Option<IgmpServer>, device_reader: DeviceReader,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>, device_writer: DeviceWriter,
ip_route: Option<ExternalRoute>, igmp_server: Option<IgmpServer>,
ip_proxy_map: Option<IpProxyMap>, current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher, server_cipher: Cipher, parallel: usize) { ip_route: Option<ExternalRoute>,
ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
parallel: usize,
) {
if parallel == 1 { if parallel == 1 {
thread::Builder::new().name("tun_handler".into()).spawn(move || { thread::Builder::new()
tokio::runtime::Builder::new_current_thread() .name("tun_handler".into())
.enable_all().build().unwrap() .spawn(move || {
.block_on(async move { tokio::runtime::Builder::new_current_thread()
if let Err(e) = start_simple(sender, device_reader, &device_writer, igmp_server, current_device, ip_route, ip_proxy_map, client_cipher, server_cipher).await { .enable_all()
log::warn!("stop:{}",e); .build()
} .unwrap()
let _ = device_writer.close(); .block_on(async move {
worker.stop_all(); if let Err(e) = start_simple(
}) sender,
}).unwrap(); device_reader,
&device_writer,
igmp_server,
current_device,
ip_route,
ip_proxy_map,
client_cipher,
server_cipher,
)
.await
{
log::warn!("stop:{}", e);
}
let _ = device_writer.close();
worker.stop_all();
})
})
.unwrap();
} else { } else {
let (buf_sender, buf_receiver) = buf_channel_group(parallel); let (buf_sender, buf_receiver) = buf_channel_group(parallel);
for mut buf_receiver in buf_receiver.0 { for mut buf_receiver in buf_receiver.0 {
@@ -94,8 +134,20 @@ pub async fn start(worker: VntWorker, sender: ChannelSender,
let server_cipher = server_cipher.clone(); let server_cipher = server_cipher.clone();
tokio::spawn(async move { tokio::spawn(async move {
while let Some((mut buf, start, len)) = buf_receiver.recv().await { while let Some((mut buf, start, len)) = buf_receiver.recv().await {
match handle(&sender, &mut buf[start..], len, &device_writer, &igmp_server, current_device.load(), match handle(
&ip_route, &ip_proxy_map, &client_cipher, &server_cipher).await { &sender,
&mut buf[start..],
len,
&device_writer,
&igmp_server,
current_device.load(),
&ip_route,
&ip_proxy_map,
&client_cipher,
&server_cipher,
)
.await
{
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
log::warn!("{:?}", e) log::warn!("{:?}", e)
@@ -105,21 +157,30 @@ pub async fn start(worker: VntWorker, sender: ChannelSender,
}); });
} }
thread::Builder::new().name("tun_handler".into()).spawn(move || { thread::Builder::new()
tokio::runtime::Builder::new_current_thread() .name("tun_handler".into())
.enable_all().build().unwrap() .spawn(move || {
.block_on(async move { tokio::runtime::Builder::new_current_thread()
if let Err(e) = start_(sender, device_reader, buf_sender).await { .enable_all()
log::warn!("stop:{}",e); .build()
} .unwrap()
let _ = device_writer.close(); .block_on(async move {
worker.stop_all(); if let Err(e) = start_(sender, device_reader, buf_sender).await {
}) log::warn!("stop:{}", e);
}).unwrap(); }
let _ = device_writer.close();
worker.stop_all();
})
})
.unwrap();
} }
} }
async fn start_(sender: ChannelSender, device_reader: DeviceReader, mut buf_sender: BufSenderGroup) -> io::Result<()> { async fn start_(
sender: ChannelSender,
device_reader: DeviceReader,
mut buf_sender: BufSenderGroup,
) -> io::Result<()> {
loop { loop {
let mut buf = POOL.alloc(4096); let mut buf = POOL.alloc(4096);
buf[..12].fill(0); buf[..12].fill(0);
@@ -129,21 +190,27 @@ async fn start_(sender: ChannelSender, device_reader: DeviceReader, mut buf_send
let start = 0; let start = 0;
let len = device_reader.read(&mut buf[12..])? + 12; let len = device_reader.read(&mut buf[12..])? + 12;
#[cfg(any(target_os = "macos"))] #[cfg(any(target_os = "macos"))]
let start = 4; let start = 4;
if !buf_sender.send((buf, start, len)).await { if !buf_sender.send((buf, start, len)).await {
return Err(io::Error::new(io::ErrorKind::Other, "tun buf_sender发送失败")); return Err(io::Error::new(
io::ErrorKind::Other,
"tun buf_sender发送失败",
));
} }
} }
} }
async fn start_simple(sender: ChannelSender, async fn start_simple(
device_reader: DeviceReader, sender: ChannelSender,
device_writer: &DeviceWriter, device_reader: DeviceReader,
igmp_server: Option<IgmpServer>, device_writer: &DeviceWriter,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>, igmp_server: Option<IgmpServer>,
ip_route: Option<ExternalRoute>, current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_proxy_map: Option<IpProxyMap>, ip_route: Option<ExternalRoute>,
client_cipher: Cipher, server_cipher: Cipher) -> io::Result<()> { ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
) -> io::Result<()> {
let mut buf = [0; 4096]; let mut buf = [0; 4096];
loop { loop {
if sender.is_close() { if sender.is_close() {
@@ -152,8 +219,21 @@ async fn start_simple(sender: ChannelSender,
buf[..12].fill(0); buf[..12].fill(0);
let len = device_reader.read(&mut buf[12..])? + 12; let len = device_reader.read(&mut buf[12..])? + 12;
#[cfg(any(target_os = "macos"))] #[cfg(any(target_os = "macos"))]
let mut buf = &mut buf[4..]; let mut buf = &mut buf[4..];
match handle(&sender, &mut buf, len, device_writer, &igmp_server, current_device.load(), &ip_route, &ip_proxy_map, &client_cipher, &server_cipher).await { match handle(
&sender,
&mut buf,
len,
device_writer,
&igmp_server,
current_device.load(),
&ip_route,
&ip_proxy_map,
&client_cipher,
&server_cipher,
)
.await
{
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
log::warn!("{:?}", e) log::warn!("{:?}", e)
+41 -40
View File
@@ -1,14 +1,14 @@
use std::collections::{HashMap, HashSet}; use crate::tun_tap_device::DeviceWriter;
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use dashmap::DashMap; use dashmap::DashMap;
use parking_lot::RwLock;
use packet::igmp::igmp_v2::IgmpV2Packet; use packet::igmp::igmp_v2::IgmpV2Packet;
use packet::igmp::igmp_v3::{IgmpV3QueryPacket, IgmpV3RecordType, IgmpV3ReportPacket}; use packet::igmp::igmp_v3::{IgmpV3QueryPacket, IgmpV3RecordType, IgmpV3ReportPacket};
use packet::igmp::IgmpType; use packet::igmp::IgmpType;
use packet::ip::ipv4::protocol::Protocol; use packet::ip::ipv4::protocol::Protocol;
use crate::tun_tap_device::DeviceWriter; use parking_lot::RwLock;
use std::collections::{HashMap, HashSet};
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::time::{Duration, Instant};
//1. 定时发送query,启动时20秒一次,连发3次,之后8分钟一次 //1. 定时发送query,启动时20秒一次,连发3次,之后8分钟一次
//2. 接收网关的igmp report 维护组播源信息 //2. 接收网关的igmp report 维护组播源信息
@@ -90,9 +90,7 @@ impl IgmpServer {
std::thread::sleep(Duration::from_secs(20)) std::thread::sleep(Duration::from_secs(20))
} }
}); });
Self { Self { multicast }
multicast,
}
} }
pub fn load(&self, multicast_addr: &Ipv4Addr) -> Option<Arc<RwLock<Multicast>>> { pub fn load(&self, multicast_addr: &Ipv4Addr) -> Option<Arc<RwLock<Multicast>>> {
if let Some(entry) = self.multicast.get(multicast_addr) { if let Some(entry) = self.multicast.get(multicast_addr) {
@@ -125,9 +123,11 @@ impl IgmpServer {
return Ok(()); return Ok(());
} }
let multi = { let multi = {
self.multicast.entry(multicast_addr).or_insert_with(|| { self.multicast
Arc::new(RwLock::new(Multicast::new())) .entry(multicast_addr)
}).value().clone() .or_insert_with(|| Arc::new(RwLock::new(Multicast::new())))
.value()
.clone()
}; };
let mut guard = multi.write(); let mut guard = multi.write();
guard.members.insert(source, Instant::now()); guard.members.insert(source, Instant::now());
@@ -153,13 +153,17 @@ impl IgmpServer {
if !multicast_addr.is_multicast() { if !multicast_addr.is_multicast() {
return Ok(()); return Ok(());
} }
let multi = self.multicast.entry(multicast_addr).or_insert_with(|| { let multi = self
Arc::new(RwLock::new(Multicast::new())) .multicast
}).value().clone(); .entry(multicast_addr)
.or_insert_with(|| Arc::new(RwLock::new(Multicast::new())))
.value()
.clone();
let mut guard = multi.write(); let mut guard = multi.write();
match group_record.record_type() { match group_record.record_type() {
IgmpV3RecordType::ModeIsInclude | IgmpV3RecordType::ChangeToIncludeMode => { IgmpV3RecordType::ModeIsInclude
| IgmpV3RecordType::ChangeToIncludeMode => {
match group_record.source_addresses() { match group_record.source_addresses() {
None => { None => {
//不接收所有 //不接收所有
@@ -173,7 +177,8 @@ impl IgmpServer {
} }
} }
IgmpV3RecordType::ModeIsExclude | IgmpV3RecordType::ChangeToExcludeMode => { IgmpV3RecordType::ModeIsExclude
| IgmpV3RecordType::ChangeToExcludeMode => {
match group_record.source_addresses() { match group_record.source_addresses() {
None => { None => {
//接收所有 //接收所有
@@ -190,40 +195,36 @@ impl IgmpServer {
//在已有源的基础上,接收目标源,如果是排除模式,则删除;是包含模式则添加 //在已有源的基础上,接收目标源,如果是排除模式,则删除;是包含模式则添加
match group_record.source_addresses() { match group_record.source_addresses() {
None => {} None => {}
Some(src) => { Some(src) => match guard.map.get_mut(&source) {
match guard.map.get_mut(&source) { None => {}
None => {} Some((is_include, set)) => {
Some((is_include, set)) => { for ip in src {
for ip in src { if *is_include {
if *is_include { set.insert(ip);
set.insert(ip); } else {
} else { set.remove(&ip);
set.remove(&ip);
}
} }
} }
} }
} },
} }
} }
IgmpV3RecordType::BlockOldSources => { IgmpV3RecordType::BlockOldSources => {
//在已有源的基础上,不接收目标源 //在已有源的基础上,不接收目标源
match group_record.source_addresses() { match group_record.source_addresses() {
None => {} None => {}
Some(src) => { Some(src) => match guard.map.get_mut(&source) {
match guard.map.get_mut(&source) { None => {}
None => {} Some((is_include, set)) => {
Some((is_include, set)) => { for ip in src {
for ip in src { if *is_include {
if *is_include { set.remove(&ip);
set.remove(&ip); } else {
} else { set.insert(ip);
set.insert(ip);
}
} }
} }
} }
} },
} }
} }
IgmpV3RecordType::Unknown(_) => {} IgmpV3RecordType::Unknown(_) => {}
@@ -235,4 +236,4 @@ impl IgmpServer {
} }
Ok(()) Ok(())
} }
} }
+58 -29
View File
@@ -1,20 +1,20 @@
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use std::io; use std::io;
use std::mem::MaybeUninit; use std::mem::MaybeUninit;
use std::net::{IpAddr, Ipv4Addr, SocketAddrV4}; use std::net::{IpAddr, Ipv4Addr, SocketAddrV4};
use std::sync::Arc; use std::sync::Arc;
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use socket2::{Domain, SockAddr, Socket, Type}; use socket2::{Domain, SockAddr, Socket, Type};
use packet::icmp::icmp;
use packet::icmp::icmp::HeaderOther;
use packet::ip::ipv4;
use crate::channel::sender::ChannelSender; use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher; use crate::cipher::Cipher;
use crate::handle::CurrentDeviceInfo; use crate::handle::CurrentDeviceInfo;
use crate::protocol::{MAX_TTL, NetPacket, Protocol, Version};
use crate::protocol::body::ENCRYPTION_RESERVED; use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{NetPacket, Protocol, Version, MAX_TTL};
use packet::icmp::icmp;
use packet::icmp::icmp::HeaderOther;
use packet::ip::ipv4;
pub struct IcmpProxy { pub struct IcmpProxy {
icmp_socket: Arc<Socket>, icmp_socket: Arc<Socket>,
@@ -26,9 +26,18 @@ pub struct IcmpProxy {
} }
impl IcmpProxy { impl IcmpProxy {
pub fn new(addr: SocketAddrV4, icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>, pub fn new(
sender: ChannelSender, current_device: Arc<AtomicCell<CurrentDeviceInfo>>, client_cipher: Cipher) -> io::Result<IcmpProxy> { addr: SocketAddrV4,
let icmp_socket = Arc::new(Socket::new(Domain::IPV4, Type::RAW, Some(socket2::Protocol::ICMPV4))?); icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
sender: ChannelSender,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<IcmpProxy> {
let icmp_socket = Arc::new(Socket::new(
Domain::IPV4,
Type::RAW,
Some(socket2::Protocol::ICMPV4),
)?);
icmp_socket.bind(&SockAddr::from(addr))?; icmp_socket.bind(&SockAddr::from(addr))?;
Ok(IcmpProxy { Ok(IcmpProxy {
icmp_socket, icmp_socket,
@@ -43,8 +52,7 @@ impl IcmpProxy {
} }
pub fn start(self) { pub fn start(self) {
let mut buf = [0 as u8; 1500]; let mut buf = [0 as u8; 1500];
let data: &mut [MaybeUninit<u8>] = let data: &mut [MaybeUninit<u8>] = unsafe { std::mem::transmute(&mut buf[..]) };
unsafe { std::mem::transmute(&mut buf[..]) };
loop { loop {
match self.recv(data) { match self.recv(data) {
@@ -57,29 +65,54 @@ impl IcmpProxy {
Ok(icmp_packet) => { Ok(icmp_packet) => {
match icmp_packet.header_other() { match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => { HeaderOther::Identifier(id, seq) => {
if let Some(entry) = self.icmp_proxy_map.get(&(peer_ip, id, seq)) { if let Some(entry) =
self.icmp_proxy_map.get(&(peer_ip, id, seq))
{
//将数据发送到真实的来源 //将数据发送到真实的来源
let dest_ip = *entry.value(); let dest_ip = *entry.value();
drop(entry); drop(entry);
ipv4_packet.set_destination_ip(dest_ip); ipv4_packet.set_destination_ip(dest_ip);
ipv4_packet.update_checksum(); ipv4_packet.update_checksum();
let current_device = self.current_device.load(); let current_device =
let virtual_ip = current_device.virtual_ip(); self.current_device.load();
let connect_server = current_device.connect_server; let virtual_ip =
let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + len + ENCRYPTION_RESERVED]).unwrap(); current_device.virtual_ip();
let connect_server =
current_device.connect_server;
let mut net_packet =
NetPacket::new_encrypt(vec![
0u8;
12 + len + ENCRYPTION_RESERVED
])
.unwrap();
net_packet.set_version(Version::V1); net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::IpTurn); net_packet.set_protocol(Protocol::IpTurn);
net_packet.set_transport_protocol(crate::protocol::ip_turn_packet::Protocol::Ipv4.into()); net_packet.set_transport_protocol(crate::protocol::ip_turn_packet::Protocol::Ipv4.into());
net_packet.first_set_ttl(MAX_TTL); net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip); net_packet.set_source(virtual_ip);
net_packet.set_destination(dest_ip); net_packet.set_destination(dest_ip);
net_packet.set_payload(ipv4_packet.buffer).unwrap(); net_packet
if let Err(e) = self.client_cipher.encrypt_ipv4(&mut net_packet) { .set_payload(ipv4_packet.buffer)
log::warn!("加密失败:{}",e); .unwrap();
if let Err(e) = self
.client_cipher
.encrypt_ipv4(&mut net_packet)
{
log::warn!("加密失败:{}", e);
continue; continue;
} }
if self.sender.try_send_by_id(net_packet.buffer(), &dest_ip).is_err() { if self
let _ = self.sender.try_send_main(net_packet.buffer(), connect_server); .sender
.try_send_by_id(
net_packet.buffer(),
&dest_ip,
)
.is_err()
{
let _ = self.sender.try_send_main(
net_packet.buffer(),
connect_server,
);
} }
} }
} }
@@ -98,7 +131,7 @@ impl IcmpProxy {
} }
} }
Err(e) => { Err(e) => {
log::warn!("icmp代理异常:{:?}",e); log::warn!("icmp代理异常:{:?}", e);
} }
} }
} }
@@ -106,16 +139,12 @@ impl IcmpProxy {
fn recv(&self, buf: &mut [MaybeUninit<u8>]) -> io::Result<(usize, IpAddr)> { fn recv(&self, buf: &mut [MaybeUninit<u8>]) -> io::Result<(usize, IpAddr)> {
let (size, addr) = self.icmp_socket.recv_from(buf)?; let (size, addr) = self.icmp_socket.recv_from(buf)?;
let addr = match addr.as_socket() { let addr = match addr.as_socket() {
None => { None => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
IpAddr::V4(Ipv4Addr::UNSPECIFIED) Some(add) => add.ip(),
}
Some(add) => {
add.ip()
}
}; };
Ok((size, addr)) Ok((size, addr))
} }
// fn send_to(&self, buf: &[u8], addr: SocketAddrV4) -> io::Result<usize> { // fn send_to(&self, buf: &[u8], addr: SocketAddrV4) -> io::Result<usize> {
// self.icmp_socket.send_to(buf, &SockAddr::from(addr)) // self.icmp_socket.send_to(buf, &SockAddr::from(addr))
// } // }
} }
+36 -22
View File
@@ -1,16 +1,16 @@
use std::{io, thread};
use std::net::{Ipv4Addr, SocketAddrV4};
use std::sync::Arc;
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use socket2::{SockAddr, Socket};
use tokio::net::{TcpListener, UdpSocket};
use crate::channel::sender::ChannelSender; use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher; use crate::cipher::Cipher;
use crate::handle::CurrentDeviceInfo; use crate::handle::CurrentDeviceInfo;
use crate::ip_proxy::icmp_proxy::IcmpProxy; use crate::ip_proxy::icmp_proxy::IcmpProxy;
use crate::ip_proxy::tcp_proxy::TcpProxy; use crate::ip_proxy::tcp_proxy::TcpProxy;
use crate::ip_proxy::udp_proxy::UdpProxy; use crate::ip_proxy::udp_proxy::UdpProxy;
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use socket2::{SockAddr, Socket};
use std::net::{Ipv4Addr, SocketAddrV4};
use std::sync::Arc;
use std::{io, thread};
use tokio::net::{TcpListener, UdpSocket};
pub mod icmp_proxy; pub mod icmp_proxy;
pub mod tcp_proxy; pub mod tcp_proxy;
@@ -37,13 +37,18 @@ pub struct IpProxyMap {
impl IpProxyMap { impl IpProxyMap {
pub fn send_icmp(&self, buf: &[u8], dest: &Ipv4Addr) -> io::Result<usize> { pub fn send_icmp(&self, buf: &[u8], dest: &Ipv4Addr) -> io::Result<usize> {
self.icmp_socket.send_to(buf, &SockAddr::from(SocketAddrV4::new(*dest, 0))) self.icmp_socket
.send_to(buf, &SockAddr::from(SocketAddrV4::new(*dest, 0)))
} }
} }
pub async fn init_proxy(sender: ChannelSender, current_device: Arc<AtomicCell<CurrentDeviceInfo>>, client_cipher: Cipher,) -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> { pub async fn init_proxy(
let tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new()); sender: ChannelSender,
let udp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new()); current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> {
let tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new());
let udp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new());
let icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>> = Arc::new(DashMap::new()); let icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>> = Arc::new(DashMap::new());
let tcp_listener = TcpListener::bind("0.0.0.0:0").await?; let tcp_listener = TcpListener::bind("0.0.0.0:0").await?;
let udp_socket = UdpSocket::bind("0.0.0.0:0").await?; let udp_socket = UdpSocket::bind("0.0.0.0:0").await?;
@@ -52,19 +57,28 @@ pub async fn init_proxy(sender: ChannelSender, current_device: Arc<AtomicCell<Cu
let tcp_proxy = TcpProxy::new(tcp_listener, tcp_proxy_map.clone()); let tcp_proxy = TcpProxy::new(tcp_listener, tcp_proxy_map.clone());
let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone()); let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone());
let addr = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); let addr = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
let icmp_proxy = IcmpProxy::new(addr, icmp_proxy_map.clone(), let icmp_proxy = IcmpProxy::new(
sender.clone(), current_device.clone(),client_cipher)?; addr,
icmp_proxy_map.clone(),
sender.clone(),
current_device.clone(),
client_cipher,
)?;
let icmp_socket = icmp_proxy.icmp_socket(); let icmp_socket = icmp_proxy.icmp_socket();
thread::spawn(move || { thread::spawn(move || {
icmp_proxy.start(); icmp_proxy.start();
}); });
Ok((tcp_proxy, udp_proxy, IpProxyMap { Ok((
tcp_proxy_port, tcp_proxy,
udp_proxy_port, udp_proxy,
tcp_proxy_map, IpProxyMap {
udp_proxy_map, tcp_proxy_port,
icmp_proxy_map, udp_proxy_port,
icmp_socket, tcp_proxy_map,
})) udp_proxy_map,
} icmp_proxy_map,
icmp_socket,
},
))
}
+32 -24
View File
@@ -1,7 +1,7 @@
use dashmap::DashMap;
use std::io; use std::io;
use std::net::{SocketAddr, SocketAddrV4}; use std::net::{SocketAddr, SocketAddrV4};
use std::sync::Arc; use std::sync::Arc;
use dashmap::DashMap;
use tokio::net::{TcpListener, TcpStream}; use tokio::net::{TcpListener, TcpStream};
@@ -11,7 +11,10 @@ pub struct TcpProxy {
} }
impl TcpProxy { impl TcpProxy {
pub fn new(tcp_listener: TcpListener, tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>) -> Self { pub fn new(
tcp_listener: TcpListener,
tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
) -> Self {
Self { Self {
tcp_listener, tcp_listener,
tcp_proxy_map, tcp_proxy_map,
@@ -22,31 +25,36 @@ impl TcpProxy {
let tcp_proxy_map = self.tcp_proxy_map; let tcp_proxy_map = self.tcp_proxy_map;
loop { loop {
match tcp_listener.accept().await { match tcp_listener.accept().await {
Ok((tcp_stream, sender_addr)) => { Ok((tcp_stream, sender_addr)) => match sender_addr {
match sender_addr { SocketAddr::V4(sender_addr) => {
SocketAddr::V4(sender_addr) => { if let Some(entry) = tcp_proxy_map.get(&sender_addr) {
if let Some(entry) = tcp_proxy_map.get(&sender_addr) { let dest_addr = *entry.value();
let dest_addr = *entry.value(); drop(entry);
drop(entry); let peer_tcp_stream = match TcpStream::connect(dest_addr).await {
let peer_tcp_stream = match TcpStream::connect(dest_addr).await { Ok(peer_tcp_stream) => peer_tcp_stream,
Ok(peer_tcp_stream) => { peer_tcp_stream } Err(e) => {
Err(e) => { log::warn!(
log::warn!("tcp代理异常:{:?},来源:{},目标:{}",e,sender_addr,dest_addr); "tcp代理异常:{:?},来源:{},目标:{}",
continue; e,
} sender_addr,
}; dest_addr
tokio::spawn(async move { );
if let Err(e) = proxy(tcp_stream, peer_tcp_stream).await { continue;
log::warn!("{}->{},{}",sender_addr,dest_addr,e); }
} };
}); tokio::spawn(async move {
} if let Err(e) = proxy(tcp_stream, peer_tcp_stream).await {
log::warn!("{}->{},{}", sender_addr, dest_addr, e);
}
});
} else {
log::warn!("tcp代理异常: 来源:{},未找到目标", sender_addr);
} }
SocketAddr::V6(_) => {}
} }
} SocketAddr::V6(_) => {}
},
Err(e) => { Err(e) => {
log::warn!("tcp代理监听:{:?}",e); log::warn!("tcp代理监听:{:?}", e);
} }
} }
} }
+43 -33
View File
@@ -1,8 +1,8 @@
use dashmap::DashMap;
use std::io; use std::io;
use std::net::{SocketAddr, SocketAddrV4}; use std::net::{SocketAddr, SocketAddrV4};
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use dashmap::DashMap;
use tokio::net::UdpSocket; use tokio::net::UdpSocket;
/// 一个udp代理,作用是利用系统协议栈,将udp数据报解析出来再转发到目的地址 /// 一个udp代理,作用是利用系统协议栈,将udp数据报解析出来再转发到目的地址
@@ -14,10 +14,7 @@ pub struct UdpProxy {
impl UdpProxy { impl UdpProxy {
pub fn new(udp_socket: UdpSocket, map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>) -> Self { pub fn new(udp_socket: UdpSocket, map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>) -> Self {
let udp_socket = Arc::new(udp_socket); let udp_socket = Arc::new(udp_socket);
Self { Self { udp_socket, map }
udp_socket,
map,
}
} }
pub async fn start(self) { pub async fn start(self) {
let map = self.map; let map = self.map;
@@ -27,28 +24,33 @@ impl UdpProxy {
loop { loop {
match udp_socket.recv_from(&mut buf).await { match udp_socket.recv_from(&mut buf).await {
Ok((len, sender_addr)) => { Ok((len, sender_addr)) => match sender_addr {
match sender_addr { SocketAddr::V4(sender_addr) => {
SocketAddr::V4(sender_addr) => { match start0(&buf[..len], sender_addr, &inner_map, &map, &udp_socket).await
match start0(&buf[..len], sender_addr, &inner_map, &map, &udp_socket).await { {
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
log::warn!("udp代理异常:{:?},来源:{}",e,sender_addr); log::warn!("udp代理异常:{:?},来源:{}", e, sender_addr);
}
} }
} }
SocketAddr::V6(_) => {}
} }
} SocketAddr::V6(_) => {}
},
Err(e) => { Err(e) => {
log::warn!("udp代理异常:{:?}",e); log::warn!("udp代理异常:{:?}", e);
} }
}; };
} }
} }
} }
async fn start0(buf: &[u8], sender_addr: SocketAddrV4, inner_map: &Arc<DashMap<SocketAddrV4, Arc<UdpSocket>>>, map: &Arc<DashMap<SocketAddrV4, SocketAddrV4>>, udp_socket: &Arc<UdpSocket>) -> io::Result<()> { async fn start0(
buf: &[u8],
sender_addr: SocketAddrV4,
inner_map: &Arc<DashMap<SocketAddrV4, Arc<UdpSocket>>>,
map: &Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
udp_socket: &Arc<UdpSocket>,
) -> io::Result<()> {
if let Some(entry) = inner_map.get(&sender_addr) { if let Some(entry) = inner_map.get(&sender_addr) {
let udp = entry.value().clone(); let udp = entry.value().clone();
drop(entry); drop(entry);
@@ -67,27 +69,35 @@ async fn start0(buf: &[u8], sender_addr: SocketAddrV4, inner_map: &Arc<DashMap<S
tokio::spawn(async move { tokio::spawn(async move {
let mut buf = [0u8; 65536]; let mut buf = [0u8; 65536];
loop { loop {
match tokio::time::timeout(Duration::from_secs(300), peer_udp_socket.recv(&mut buf)).await { match tokio::time::timeout(Duration::from_secs(300), peer_udp_socket.recv(&mut buf))
Ok(rs) => { .await
match rs { {
Ok(len) => { Ok(rs) => match rs {
match udp_socket.send_to(&buf[..len], sender_addr).await { Ok(len) => match udp_socket.send_to(&buf[..len], sender_addr).await {
Ok(_) => {} Ok(_) => {}
Err(e) => {
log::warn!("udp代理异常:{:?},来源:{},目标:{}",e,sender_addr,dest_addr);
break;
}
}
}
Err(e) => { Err(e) => {
log::warn!("udp代理异常:{:?},来源:{},目标:{}",e,sender_addr,dest_addr); log::warn!(
"udp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
break; break;
} }
},
Err(e) => {
log::warn!(
"udp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
break;
} }
} },
Err(_) => { Err(_) => {
//超时关闭 //超时关闭
log::warn!("udp代理超时关闭,来源:{},目标:{}",sender_addr,dest_addr); log::warn!("udp代理超时关闭,来源:{},目标:{}", sender_addr, dest_addr);
break; break;
} }
} }
@@ -97,4 +107,4 @@ async fn start0(buf: &[u8], sender_addr: SocketAddrV4, inner_map: &Arc<DashMap<S
}); });
} }
Ok(()) Ok(())
} }
+7 -7
View File
@@ -1,17 +1,17 @@
use crate::error::Error; use crate::error::Error;
pub const VNT_VERSION: &'static str = "1.2.2";
pub type Result<T> = std::result::Result<T, Error>; pub type Result<T> = std::result::Result<T, Error>;
pub mod channel;
pub mod cipher;
pub mod core;
pub mod error; pub mod error;
pub mod external_route;
pub mod handle; pub mod handle;
pub mod igmp_server;
pub mod ip_proxy;
pub mod nat; pub mod nat;
pub mod proto; pub mod proto;
pub mod protocol; pub mod protocol;
pub mod ip_proxy;
pub mod external_route;
pub mod igmp_server;
pub mod tun_tap_device; pub mod tun_tap_device;
pub mod core;
pub mod channel;
pub mod util; pub mod util;
pub mod cipher;
+53 -27
View File
@@ -1,6 +1,6 @@
use std::io; use std::io;
use std::net::{IpAddr, Ipv4Addr};
use std::net::UdpSocket; use std::net::UdpSocket;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6};
use std::sync::Arc; use std::sync::Arc;
use parking_lot::Mutex; use parking_lot::Mutex;
@@ -10,16 +10,42 @@ use crate::proto::message::PunchNatType;
mod stun_test; mod stun_test;
pub fn local_ip() -> io::Result<Ipv4Addr> { pub fn local_ipv4() -> io::Result<Ipv4Addr> {
let socket = UdpSocket::bind("0.0.0.0:0")?; let socket = UdpSocket::bind("0.0.0.0:0")?;
socket.connect("8.8.8.8:80")?; socket.connect("8.8.8.8:80")?;
let addr = socket.local_addr()?; let addr = socket.local_addr()?;
match addr.ip() { match addr.ip() {
IpAddr::V4(ip) => { IpAddr::V4(ip) => Ok(ip),
Ok(ip) IpAddr::V6(_) => Ok(Ipv4Addr::UNSPECIFIED),
}
}
pub fn local_ipv6() -> io::Result<Ipv6Addr> {
let socket = UdpSocket::bind("[::]:0")?;
socket.connect("[2001:4860:4860::8888]:80")?;
let addr = socket.local_addr()?;
match addr.ip() {
IpAddr::V4(_) => Ok(Ipv6Addr::UNSPECIFIED),
IpAddr::V6(ip) => Ok(ip),
}
}
pub fn local_ipv4_addr(port: u16) -> SocketAddrV4 {
match local_ipv4() {
Ok(ipv4) => SocketAddrV4::new(ipv4, port),
Err(e) => {
log::warn!("获取本地ipv4地址失败:{}", e);
SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0)
} }
IpAddr::V6(_) => { }
Ok(Ipv4Addr::UNSPECIFIED) }
pub fn local_ipv6_addr(port: u16) -> SocketAddrV6 {
match local_ipv6() {
Ok(ipv6) => SocketAddrV6::new(ipv6, port, 0, 0),
Err(e) => {
log::warn!("获取本地ipv6地址失败:{}", e);
SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, 0, 0, 0)
} }
} }
} }
@@ -53,22 +79,21 @@ impl NatTest {
mut stun_server: Vec<String>, mut stun_server: Vec<String>,
public_ip: Ipv4Addr, public_ip: Ipv4Addr,
public_port: u16, public_port: u16,
local_ip: Ipv4Addr, local_ipv4_addr: SocketAddrV4,
local_port: u16, ipv6_addr: SocketAddrV6,
) -> NatTest { ) -> NatTest {
let server = stun_server[0].clone(); let server = stun_server[0].clone();
stun_server.resize(3, server); stun_server.resize(3, server);
let info = NatTest::re_test_( let nat_info = Self::re_test_(
&stun_server, &stun_server,
public_ip, public_ip,
public_port, public_port,
local_ip, local_ipv4_addr,
local_port, ipv6_addr,
).await; )
NatTest { .await;
stun_server, let info = Arc::new(Mutex::new(nat_info));
info: Arc::new(Mutex::new(info)), NatTest { stun_server, info }
}
} }
pub fn nat_info(&self) -> NatInfo { pub fn nat_info(&self) -> NatInfo {
self.info.lock().clone() self.info.lock().clone()
@@ -84,16 +109,17 @@ impl NatTest {
&self, &self,
public_ip: Ipv4Addr, public_ip: Ipv4Addr,
public_port: u16, public_port: u16,
local_ip: Ipv4Addr, local_ipv4_addr: SocketAddrV4,
local_port: u16, ipv6_addr: SocketAddrV6,
) -> NatInfo { ) -> NatInfo {
let info = NatTest::re_test_( let info = NatTest::re_test_(
&self.stun_server, &self.stun_server,
public_ip, public_ip,
public_port, public_port,
local_ip, local_ipv4_addr,
local_port, ipv6_addr,
).await; )
.await;
*self.info.lock() = info.clone(); *self.info.lock() = info.clone();
info info
} }
@@ -101,8 +127,8 @@ impl NatTest {
stun_server: &Vec<String>, stun_server: &Vec<String>,
public_ip: Ipv4Addr, public_ip: Ipv4Addr,
public_port: u16, public_port: u16,
local_ip: Ipv4Addr, local_ipv4_addr: SocketAddrV4,
local_port: u16, ipv6_addr: SocketAddrV6,
) -> NatInfo { ) -> NatInfo {
return match stun_test::stun_test_nat(stun_server.clone()).await { return match stun_test::stun_test_nat(stun_server.clone()).await {
Ok((nat_type, ips, port_range)) => { Ok((nat_type, ips, port_range)) => {
@@ -117,8 +143,8 @@ impl NatTest {
public_ips, public_ips,
public_port, public_port,
port_range, port_range,
local_ip, local_ipv4_addr,
local_port, ipv6_addr,
nat_type, nat_type,
) )
} }
@@ -128,8 +154,8 @@ impl NatTest {
vec![public_ip], vec![public_ip],
public_port, public_port,
0, 0,
local_ip, local_ipv4_addr,
local_port, ipv6_addr,
NatType::Cone, NatType::Cone,
) )
} }
+18 -10
View File
@@ -3,9 +3,9 @@ use std::io;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}; use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6};
use std::time::Duration; use std::time::Duration;
use crate::channel::punch::NatType;
use stun_format::Attr; use stun_format::Attr;
use tokio::net::UdpSocket; use tokio::net::UdpSocket;
use crate::channel::punch::NatType;
pub async fn stun_test_nat(stun_servers: Vec<String>) -> io::Result<(NatType, Vec<Ipv4Addr>, u16)> { pub async fn stun_test_nat(stun_servers: Vec<String>) -> io::Result<(NatType, Vec<Ipv4Addr>, u16)> {
let mut h = Vec::new(); let mut h = Vec::new();
@@ -68,21 +68,30 @@ async fn test_nat(stun_server: String) -> io::Result<(NatType, Vec<Ipv4Addr>, u1
Ok((nat_type, hash_set.into_iter().collect(), port_range)) Ok((nat_type, hash_set.into_iter().collect(), port_range))
} }
async fn test_nat_(udp: &UdpSocket, change_ip: bool, change_port: bool) -> io::Result<(SocketAddr, SocketAddr)> { async fn test_nat_(
udp: &UdpSocket,
change_ip: bool,
change_port: bool,
) -> io::Result<(SocketAddr, SocketAddr)> {
for _ in 0..2 { for _ in 0..2 {
let mut buf = [0u8; 28]; let mut buf = [0u8; 28];
let mut msg = stun_format::MsgBuilder::from(buf.as_mut_slice()); let mut msg = stun_format::MsgBuilder::from(buf.as_mut_slice());
msg.typ(stun_format::MsgType::BindingRequest).unwrap(); msg.typ(stun_format::MsgType::BindingRequest).unwrap();
msg.tid(1).unwrap(); msg.tid(1).unwrap();
msg.add_attr(Attr::ChangeRequest { change_ip, change_port }).unwrap(); msg.add_attr(Attr::ChangeRequest {
change_ip,
change_port,
})
.unwrap();
udp.send(msg.as_bytes()).await?; udp.send(msg.as_bytes()).await?;
let mut buf = [0; 10240]; let mut buf = [0; 10240];
let (len, addr) = match tokio::time::timeout(Duration::from_millis(300), udp.recv_from(&mut buf)).await { let (len, addr) =
Ok(rs) => { rs? } match tokio::time::timeout(Duration::from_millis(300), udp.recv_from(&mut buf)).await {
Err(_) => { Ok(rs) => rs?,
continue; Err(_) => {
} continue;
}; }
};
let msg = stun_format::Msg::from(&buf[..len]); let msg = stun_format::Msg::from(&buf[..len]);
let mut mapped_addr = None; let mut mapped_addr = None;
let mut changed_addr = None; let mut changed_addr = None;
@@ -126,4 +135,3 @@ fn stun_addr(addr: stun_format::SocketAddr) -> SocketAddr {
} }
} }
} }
+39 -20
View File
@@ -1313,8 +1313,10 @@ pub struct PunchInfo {
pub local_ip: u32, pub local_ip: u32,
// @@protoc_insertion_point(field:PunchInfo.local_port) // @@protoc_insertion_point(field:PunchInfo.local_port)
pub local_port: u32, pub local_port: u32,
// @@protoc_insertion_point(field:PunchInfo.public_ipv6_list) // @@protoc_insertion_point(field:PunchInfo.ipv6)
pub public_ipv6_list: ::std::vec::Vec<::std::vec::Vec<u8>>, pub ipv6: ::std::vec::Vec<u8>,
// @@protoc_insertion_point(field:PunchInfo.ipv6_port)
pub ipv6_port: u32,
// special fields // special fields
// @@protoc_insertion_point(special_field:PunchInfo.special_fields) // @@protoc_insertion_point(special_field:PunchInfo.special_fields)
pub special_fields: ::protobuf::SpecialFields, pub special_fields: ::protobuf::SpecialFields,
@@ -1332,7 +1334,7 @@ impl PunchInfo {
} }
fn generated_message_descriptor_data() -> ::protobuf::reflect::GeneratedMessageDescriptorData { fn generated_message_descriptor_data() -> ::protobuf::reflect::GeneratedMessageDescriptorData {
let mut fields = ::std::vec::Vec::with_capacity(8); let mut fields = ::std::vec::Vec::with_capacity(9);
let mut oneofs = ::std::vec::Vec::with_capacity(0); let mut oneofs = ::std::vec::Vec::with_capacity(0);
fields.push(::protobuf::reflect::rt::v2::make_vec_simpler_accessor::<_, _>( fields.push(::protobuf::reflect::rt::v2::make_vec_simpler_accessor::<_, _>(
"public_ip_list", "public_ip_list",
@@ -1369,10 +1371,15 @@ impl PunchInfo {
|m: &PunchInfo| { &m.local_port }, |m: &PunchInfo| { &m.local_port },
|m: &mut PunchInfo| { &mut m.local_port }, |m: &mut PunchInfo| { &mut m.local_port },
)); ));
fields.push(::protobuf::reflect::rt::v2::make_vec_simpler_accessor::<_, _>( fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"public_ipv6_list", "ipv6",
|m: &PunchInfo| { &m.public_ipv6_list }, |m: &PunchInfo| { &m.ipv6 },
|m: &mut PunchInfo| { &mut m.public_ipv6_list }, |m: &mut PunchInfo| { &mut m.ipv6 },
));
fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"ipv6_port",
|m: &PunchInfo| { &m.ipv6_port },
|m: &mut PunchInfo| { &mut m.ipv6_port },
)); ));
::protobuf::reflect::GeneratedMessageDescriptorData::new_2::<PunchInfo>( ::protobuf::reflect::GeneratedMessageDescriptorData::new_2::<PunchInfo>(
"PunchInfo", "PunchInfo",
@@ -1417,7 +1424,10 @@ impl ::protobuf::Message for PunchInfo {
self.local_port = is.read_uint32()?; self.local_port = is.read_uint32()?;
}, },
74 => { 74 => {
self.public_ipv6_list.push(is.read_bytes()?); self.ipv6 = is.read_bytes()?;
},
80 => {
self.ipv6_port = is.read_uint32()?;
}, },
tag => { tag => {
::protobuf::rt::read_unknown_or_skip_group(tag, is, self.special_fields.mut_unknown_fields())?; ::protobuf::rt::read_unknown_or_skip_group(tag, is, self.special_fields.mut_unknown_fields())?;
@@ -1450,9 +1460,12 @@ impl ::protobuf::Message for PunchInfo {
if self.local_port != 0 { if self.local_port != 0 {
my_size += ::protobuf::rt::uint32_size(8, self.local_port); my_size += ::protobuf::rt::uint32_size(8, self.local_port);
} }
for value in &self.public_ipv6_list { if !self.ipv6.is_empty() {
my_size += ::protobuf::rt::bytes_size(9, &value); my_size += ::protobuf::rt::bytes_size(9, &self.ipv6);
}; }
if self.ipv6_port != 0 {
my_size += ::protobuf::rt::uint32_size(10, self.ipv6_port);
}
my_size += ::protobuf::rt::unknown_fields_size(self.special_fields.unknown_fields()); my_size += ::protobuf::rt::unknown_fields_size(self.special_fields.unknown_fields());
self.special_fields.cached_size().set(my_size as u32); self.special_fields.cached_size().set(my_size as u32);
my_size my_size
@@ -1480,9 +1493,12 @@ impl ::protobuf::Message for PunchInfo {
if self.local_port != 0 { if self.local_port != 0 {
os.write_uint32(8, self.local_port)?; os.write_uint32(8, self.local_port)?;
} }
for v in &self.public_ipv6_list { if !self.ipv6.is_empty() {
os.write_bytes(9, &v)?; os.write_bytes(9, &self.ipv6)?;
}; }
if self.ipv6_port != 0 {
os.write_uint32(10, self.ipv6_port)?;
}
os.write_unknown_fields(self.special_fields.unknown_fields())?; os.write_unknown_fields(self.special_fields.unknown_fields())?;
::std::result::Result::Ok(()) ::std::result::Result::Ok(())
} }
@@ -1507,7 +1523,8 @@ impl ::protobuf::Message for PunchInfo {
self.reply = false; self.reply = false;
self.local_ip = 0; self.local_ip = 0;
self.local_port = 0; self.local_port = 0;
self.public_ipv6_list.clear(); self.ipv6.clear();
self.ipv6_port = 0;
self.special_fields.clear(); self.special_fields.clear();
} }
@@ -1520,7 +1537,8 @@ impl ::protobuf::Message for PunchInfo {
reply: false, reply: false,
local_ip: 0, local_ip: 0,
local_port: 0, local_port: 0,
public_ipv6_list: ::std::vec::Vec::new(), ipv6: ::std::vec::Vec::new(),
ipv6_port: 0,
special_fields: ::protobuf::SpecialFields::new(), special_fields: ::protobuf::SpecialFields::new(),
}; };
&instance &instance
@@ -1625,15 +1643,16 @@ static file_descriptor_proto_data: &'static [u8] = b"\
\n\rdevice_status\x18\x03\x20\x01(\rR\x0cdeviceStatus\x12#\n\rclient_sec\ \n\rdevice_status\x18\x03\x20\x01(\rR\x0cdeviceStatus\x12#\n\rclient_sec\
ret\x18\x04\x20\x01(\x08R\x0cclientSecret\"Y\n\nDeviceList\x12\x14\n\x05\ ret\x18\x04\x20\x01(\x08R\x0cclientSecret\"Y\n\nDeviceList\x12\x14\n\x05\
epoch\x18\x01\x20\x01(\rR\x05epoch\x125\n\x10device_info_list\x18\x02\ epoch\x18\x01\x20\x01(\rR\x05epoch\x125\n\x10device_info_list\x18\x02\
\x20\x03(\x0b2\x0b.DeviceInfoR\x0edeviceInfoList\"\xa2\x02\n\tPunchInfo\ \x20\x03(\x0b2\x0b.DeviceInfoR\x0edeviceInfoList\"\xa9\x02\n\tPunchInfo\
\x12$\n\x0epublic_ip_list\x18\x02\x20\x03(\x07R\x0cpublicIpList\x12\x1f\ \x12$\n\x0epublic_ip_list\x18\x02\x20\x03(\x07R\x0cpublicIpList\x12\x1f\
\n\x0bpublic_port\x18\x03\x20\x01(\rR\npublicPort\x12*\n\x11public_port_\ \n\x0bpublic_port\x18\x03\x20\x01(\rR\npublicPort\x12*\n\x11public_port_\
range\x18\x04\x20\x01(\rR\x0fpublicPortRange\x12(\n\x08nat_type\x18\x05\ range\x18\x04\x20\x01(\rR\x0fpublicPortRange\x12(\n\x08nat_type\x18\x05\
\x20\x01(\x0e2\r.PunchNatTypeR\x07natType\x12\x14\n\x05reply\x18\x06\x20\ \x20\x01(\x0e2\r.PunchNatTypeR\x07natType\x12\x14\n\x05reply\x18\x06\x20\
\x01(\x08R\x05reply\x12\x19\n\x08local_ip\x18\x07\x20\x01(\x07R\x07local\ \x01(\x08R\x05reply\x12\x19\n\x08local_ip\x18\x07\x20\x01(\x07R\x07local\
Ip\x12\x1d\n\nlocal_port\x18\x08\x20\x01(\rR\tlocalPort\x12(\n\x10public\ Ip\x12\x1d\n\nlocal_port\x18\x08\x20\x01(\rR\tlocalPort\x12\x12\n\x04ipv\
_ipv6_list\x18\t\x20\x03(\x0cR\x0epublicIpv6List*'\n\x0cPunchNatType\x12\ 6\x18\t\x20\x01(\x0cR\x04ipv6\x12\x1b\n\tipv6_port\x18\n\x20\x01(\rR\x08\
\r\n\tSymmetric\x10\0\x12\x08\n\x04Cone\x10\x01b\x06proto3\ ipv6Port*'\n\x0cPunchNatType\x12\r\n\tSymmetric\x10\0\x12\x08\n\x04Cone\
\x10\x01b\x06proto3\
"; ";
/// `FileDescriptorProto` object which was a source for this generated file /// `FileDescriptorProto` object which was a source for this generated file
+169 -102
View File
@@ -2,66 +2,83 @@ use std::{fmt, io};
pub const ENCRYPTION_RESERVED: usize = 32; pub const ENCRYPTION_RESERVED: usize = 32;
/* aes_gcm加密数据体 /* aes_gcm加密数据体
0 15 31 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 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
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| | | |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| random(32) | | random(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| tag(32) | | tag(32) |
| tag(32) | | tag(32) |
| tag(32) | | tag(32) |
| tag(32) | | tag(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| finger(32) | | finger(32) |
| finger(32) | | finger(32) |
| finger(32) | | finger(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
finger用于快速校验数据是否被修改使tokenfinger finger用于快速校验数据是否被修改使tokenfinger
() ()
*/ */
pub struct SecretBody<B> { pub struct SecretBody<B> {
buffer: B, buffer: B,
exist_finger: bool,
} }
impl<B: AsRef<[u8]>> SecretBody<B> { impl<B: AsRef<[u8]>> SecretBody<B> {
pub fn new(buffer: B) -> io::Result<SecretBody<B>> { pub fn new(buffer: B, exist_finger: bool) -> io::Result<SecretBody<B>> {
let len = buffer.as_ref().len(); let len = buffer.as_ref().len();
let min_len = if exist_finger { 32 } else { 32 - 12 };
// 不能大于udp最大载荷长度 // 不能大于udp最大载荷长度
if len < 32 || len > 65535 - 20 - 8 - 12 { if len < min_len || len > 65535 - 20 - 8 - 12 {
return Err(io::Error::new( return Err(io::Error::new(
io::ErrorKind::InvalidData, io::ErrorKind::InvalidData,
"length overflow", "SecretBody length overflow",
)); ));
} }
Ok(SecretBody { buffer }) Ok(SecretBody {
} buffer,
pub fn data(&self) -> &[u8] { exist_finger,
let end = self.buffer.as_ref().len() - 32; })
&self.buffer.as_ref()[..end]
} }
pub fn random(&self) -> u32 { pub fn random(&self) -> u32 {
let end = self.buffer.as_ref().len() - 16 - 12; let mut end = self.buffer.as_ref().len() - 16;
if self.exist_finger {
end -= 12;
}
u32::from_be_bytes(self.buffer.as_ref()[end - 4..end].try_into().unwrap()) u32::from_be_bytes(self.buffer.as_ref()[end - 4..end].try_into().unwrap())
} }
pub fn body(&self) -> &[u8] { pub fn body(&self) -> &[u8] {
let end = self.buffer.as_ref().len() - 16 - 12; let mut end = self.buffer.as_ref().len() - 16;
if self.exist_finger {
end -= 12;
}
&self.buffer.as_ref()[..end] &self.buffer.as_ref()[..end]
} }
pub fn tag(&self) -> &[u8] { pub fn tag(&self) -> &[u8] {
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
&self.buffer.as_ref()[end - 16..end] &self.buffer.as_ref()[end - 16..end]
} }
/// 数据部分+tag部分 /// 数据部分+tag部分
pub fn en_body(&self) -> &[u8] { pub fn en_body(&self) -> &[u8] {
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
&self.buffer.as_ref()[..end] &self.buffer.as_ref()[..end]
} }
pub fn finger(&self) -> &[u8] { pub fn finger(&self) -> &[u8] {
let end = self.buffer.as_ref().len(); if self.exist_finger {
&self.buffer.as_ref()[end - 12..end] let end = self.buffer.as_ref().len();
&self.buffer.as_ref()[end - 12..end]
} else {
&[]
}
} }
pub fn buffer(&self) -> &[u8] { pub fn buffer(&self) -> &[u8] {
self.buffer.as_ref() self.buffer.as_ref()
@@ -69,16 +86,11 @@ impl<B: AsRef<[u8]>> SecretBody<B> {
} }
impl<B: AsRef<[u8]> + AsMut<[u8]>> SecretBody<B> { impl<B: AsRef<[u8]> + AsMut<[u8]>> SecretBody<B> {
pub fn set_data(&mut self, data: &[u8]) -> io::Result<()> {
let end = self.buffer.as_ref().len() - 32;
if end - 4 != data.len() {
return Err(io::Error::new(io::ErrorKind::InvalidData, "end-4 != data.len"));
}
self.buffer.as_mut()[..end].copy_from_slice(data);
Ok(())
}
pub fn set_random(&mut self, random: u32) { pub fn set_random(&mut self, random: u32) {
let end = self.buffer.as_ref().len() - 16 - 12; let mut end = self.buffer.as_ref().len() - 16;
if self.exist_finger {
end -= 12;
}
self.buffer.as_mut()[end - 4..end].copy_from_slice(&random.to_be_bytes()); self.buffer.as_mut()[end - 4..end].copy_from_slice(&random.to_be_bytes());
} }
@@ -86,35 +98,53 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> SecretBody<B> {
if tag.len() != 16 { if tag.len() != 16 {
return Err(io::Error::new(io::ErrorKind::InvalidData, "tag.len != 16")); return Err(io::Error::new(io::ErrorKind::InvalidData, "tag.len != 16"));
} }
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
self.buffer.as_mut()[end - 16..end].copy_from_slice(tag); self.buffer.as_mut()[end - 16..end].copy_from_slice(tag);
Ok(()) Ok(())
} }
pub fn set_finger(&mut self, finger: &[u8]) -> io::Result<()> { pub fn set_finger(&mut self, finger: &[u8]) -> io::Result<()> {
if finger.len() != 12 { if self.exist_finger {
return Err(io::Error::new(io::ErrorKind::InvalidData, "finger.len != 12")); if finger.len() != 12 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"finger.len != 12",
));
}
let end = self.buffer.as_ref().len();
self.buffer.as_mut()[end - 12..end].copy_from_slice(finger);
Ok(())
} else {
Err(io::Error::new(
io::ErrorKind::InvalidData,
"not exist finger",
))
} }
let end = self.buffer.as_ref().len();
self.buffer.as_mut()[end - 12..end].copy_from_slice(finger);
Ok(())
} }
pub fn data_mut(&mut self) -> &mut [u8] {
let end = self.buffer.as_ref().len() - 32;
&mut self.buffer.as_mut()[..end]
}
/// 数据部分 /// 数据部分
pub fn body_mut(&mut self) -> &mut [u8] { pub fn body_mut(&mut self) -> &mut [u8] {
let end = self.buffer.as_ref().len() - 12 - 16; let mut end = self.buffer.as_ref().len() - 16;
if self.exist_finger {
end -= 12;
}
&mut self.buffer.as_mut()[..end] &mut self.buffer.as_mut()[..end]
} }
pub fn tag_mut(&mut self) -> &mut [u8] { pub fn tag_mut(&mut self) -> &mut [u8] {
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
&mut self.buffer.as_mut()[end - 16..end] &mut self.buffer.as_mut()[end - 16..end]
} }
/// 数据部分+tag部分 /// 数据部分+tag部分
pub fn en_body_mut(&mut self) -> &mut [u8] { pub fn en_body_mut(&mut self) -> &mut [u8] {
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
&mut self.buffer.as_mut()[..end] &mut self.buffer.as_mut()[..end]
} }
pub fn buffer_mut(&mut self) -> &mut [u8] { pub fn buffer_mut(&mut self) -> &mut [u8] {
@@ -128,85 +158,116 @@ impl<B: AsRef<[u8]>> fmt::Debug for SecretBody<B> {
.field("random", &self.random()) .field("random", &self.random())
.field("body", &self.body()) .field("body", &self.body())
.field("tag", &self.tag()) .field("tag", &self.tag())
.field("finger", &self.finger())
.finish() .finish()
} }
} }
/* aes_cbc加密数据体 /* aes_cbc加密数据体
0 15 31 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 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
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| | | |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| random(32) | | random(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| finger(32) | | finger(32) |
| finger(32) | | finger(32) |
| finger(32) | | finger(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
finger用于快速校验数据是否被修改使tokenfinger finger用于快速校验数据是否被修改使tokenfinger
() ()
*/ */
pub struct AesCbcSecretBody<B> { pub struct AesCbcSecretBody<B> {
buffer: B, buffer: B,
exist_finger: bool,
} }
impl<B: AsRef<[u8]>> AesCbcSecretBody<B> { impl<B: AsRef<[u8]>> AesCbcSecretBody<B> {
pub fn new(buffer: B) -> io::Result<AesCbcSecretBody<B>> { pub fn new(buffer: B, exist_finger: bool) -> io::Result<AesCbcSecretBody<B>> {
let len = buffer.as_ref().len(); let len = buffer.as_ref().len();
let min_len = if exist_finger { 16 } else { 16 - 12 };
// 不能大于udp最大载荷长度 // 不能大于udp最大载荷长度
if len < 16 || len > 65535 - 20 - 8 - 12 { if len < min_len || len > 65535 - 20 - 8 - 12 {
return Err(io::Error::new( return Err(io::Error::new(
io::ErrorKind::InvalidData, io::ErrorKind::InvalidData,
"length overflow", "AesCbcSecretBody length overflow",
)); ));
} }
Ok(AesCbcSecretBody { buffer }) Ok(AesCbcSecretBody {
buffer,
exist_finger,
})
} }
pub fn en_body(&self) -> &[u8] { pub fn en_body(&self) -> &[u8] {
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
&self.buffer.as_ref()[..end] &self.buffer.as_ref()[..end]
} }
pub fn finger(&self) -> &[u8] { pub fn finger(&self) -> &[u8] {
let end = self.buffer.as_ref().len(); if self.exist_finger {
&self.buffer.as_ref()[end - 12..end] let end = self.buffer.as_ref().len();
&self.buffer.as_ref()[end - 12..end]
} else {
&[]
}
} }
} }
impl<B: AsRef<[u8]> + AsMut<[u8]>> AesCbcSecretBody<B> { impl<B: AsRef<[u8]> + AsMut<[u8]>> AesCbcSecretBody<B> {
pub fn set_random(&mut self, random: u32) { pub fn set_random(&mut self, random: u32) {
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
self.buffer.as_mut()[end - 4..end].copy_from_slice(&random.to_be_bytes()); self.buffer.as_mut()[end - 4..end].copy_from_slice(&random.to_be_bytes());
} }
pub fn set_finger(&mut self, finger: &[u8]) -> io::Result<()> { pub fn set_finger(&mut self, finger: &[u8]) -> io::Result<()> {
if finger.len() != 12 { if self.exist_finger {
return Err(io::Error::new(io::ErrorKind::InvalidData, "finger.len != 12")); if finger.len() != 12 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"finger.len != 12",
));
}
let end = self.buffer.as_ref().len();
self.buffer.as_mut()[end - 12..end].copy_from_slice(finger);
Ok(())
} else {
Err(io::Error::new(
io::ErrorKind::InvalidData,
"cbc not exist finger",
))
} }
let end = self.buffer.as_ref().len();
self.buffer.as_mut()[end - 12..end].copy_from_slice(finger);
Ok(())
} }
pub fn en_body_mut(&mut self) -> &mut [u8] { pub fn en_body_mut(&mut self) -> &mut [u8] {
let end = self.buffer.as_ref().len() - 12; let mut end = self.buffer.as_ref().len();
if self.exist_finger {
end -= 12;
}
&mut self.buffer.as_mut()[..end] &mut self.buffer.as_mut()[..end]
} }
} }
/* rsa加密数据体 /* rsa加密数据体
0 15 31 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 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
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| (n) | | (n) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| random(32) | | random(32) |
| random(32) | | random(32) |
| random(32) | | random(32) |
| random(32) | | random(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| finger(32) | | finger(32) |
| finger(32) | | finger(32) |
| finger(32) | | finger(32) |
| finger(32) | | finger(32) |
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
*/ */
pub struct RsaSecretBody<B> { pub struct RsaSecretBody<B> {
buffer: B, buffer: B,
} }
@@ -247,7 +308,10 @@ impl<B: AsRef<[u8]>> RsaSecretBody<B> {
impl<B: AsRef<[u8]> + AsMut<[u8]>> RsaSecretBody<B> { impl<B: AsRef<[u8]> + AsMut<[u8]>> RsaSecretBody<B> {
pub fn set_random(&mut self, random: &[u8]) -> io::Result<()> { pub fn set_random(&mut self, random: &[u8]) -> io::Result<()> {
if random.len() != 16 { if random.len() != 16 {
return Err(io::Error::new(io::ErrorKind::InvalidData, "random.len != 16")); return Err(io::Error::new(
io::ErrorKind::InvalidData,
"random.len != 16",
));
} }
let end = self.buffer.as_ref().len() - 16; let end = self.buffer.as_ref().len() - 16;
self.buffer.as_mut()[end - 16..end].copy_from_slice(random); self.buffer.as_mut()[end - 16..end].copy_from_slice(random);
@@ -259,10 +323,13 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> RsaSecretBody<B> {
} }
pub fn set_finger(&mut self, finger: &[u8]) -> io::Result<()> { pub fn set_finger(&mut self, finger: &[u8]) -> io::Result<()> {
if finger.len() != 16 { if finger.len() != 16 {
return Err(io::Error::new(io::ErrorKind::InvalidData, "finger.len != 16")); return Err(io::Error::new(
io::ErrorKind::InvalidData,
"finger.len != 16",
));
} }
let end = self.buffer.as_ref().len(); let end = self.buffer.as_ref().len();
self.buffer.as_mut()[end - 16..end].copy_from_slice(finger); self.buffer.as_mut()[end - 16..end].copy_from_slice(finger);
Ok(()) Ok(())
} }
} }
+1 -1
View File
@@ -1,5 +1,5 @@
use std::{fmt, io};
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::{fmt, io};
#[derive(Eq, PartialEq, Copy, Clone, Debug)] #[derive(Eq, PartialEq, Copy, Clone, Debug)]
pub enum Protocol { pub enum Protocol {
+12 -14
View File
@@ -42,7 +42,7 @@ impl<B: AsRef<[u8]>> BroadcastPacket<B> {
if len < 2 + 4 || packet.addr_num() == 0 { if len < 2 + 4 || packet.addr_num() == 0 {
Err(io::Error::new( Err(io::Error::new(
io::ErrorKind::InvalidData, io::ErrorKind::InvalidData,
"InvalidData", "BroadcastPacket InvalidData",
)) ))
} else { } else {
Ok(packet) Ok(packet)
@@ -52,7 +52,7 @@ impl<B: AsRef<[u8]>> BroadcastPacket<B> {
impl<B: AsRef<[u8]>> BroadcastPacket<B> { impl<B: AsRef<[u8]>> BroadcastPacket<B> {
pub fn addr_num(&self) -> u8 { pub fn addr_num(&self) -> u8 {
self.buffer.as_ref()[1] self.buffer.as_ref()[0]
} }
/// 已经发送给了这些地址 /// 已经发送给了这些地址
pub fn addresses(&self) -> Vec<Ipv4Addr> { pub fn addresses(&self) -> Vec<Ipv4Addr> {
@@ -61,7 +61,12 @@ impl<B: AsRef<[u8]>> BroadcastPacket<B> {
let buf = self.buffer.as_ref(); let buf = self.buffer.as_ref();
let mut offset = 1; let mut offset = 1;
for _ in 0..num { for _ in 0..num {
list.push(Ipv4Addr::new(buf[offset], buf[offset + 1], buf[offset + 2], buf[offset + 3])); list.push(Ipv4Addr::new(
buf[offset],
buf[offset + 1],
buf[offset + 2],
buf[offset + 3],
));
offset += 4; offset += 4;
} }
list list
@@ -69,10 +74,7 @@ impl<B: AsRef<[u8]>> BroadcastPacket<B> {
pub fn data(&self) -> io::Result<&[u8]> { pub fn data(&self) -> io::Result<&[u8]> {
let start = 1 + self.addr_num() as usize * 4; let start = 1 + self.addr_num() as usize * 4;
if start > self.buffer.as_ref().len() { if start > self.buffer.as_ref().len() {
Err(io::Error::new( Err(io::Error::new(io::ErrorKind::InvalidData, "InvalidData"))
io::ErrorKind::InvalidData,
"InvalidData",
))
} else { } else {
Ok(&self.buffer.as_ref()[start..]) Ok(&self.buffer.as_ref()[start..])
} }
@@ -85,7 +87,7 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> BroadcastPacket<B> {
if buf.len() < 1 + addr.len() * 4 || addr.len() > u8::MAX as usize { if buf.len() < 1 + addr.len() * 4 || addr.len() > u8::MAX as usize {
Err(io::Error::new( Err(io::Error::new(
io::ErrorKind::InvalidData, io::ErrorKind::InvalidData,
"InvalidData", "addr invalid data",
)) ))
} else { } else {
buf[0] = addr.len() as u8; buf[0] = addr.len() as u8;
@@ -101,17 +103,13 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> BroadcastPacket<B> {
let num = self.addr_num() as usize; let num = self.addr_num() as usize;
let start = 1 + 4 * num; let start = 1 + 4 * num;
let buf = self.buffer.as_mut(); let buf = self.buffer.as_mut();
if start > buf.len() || start + data.len() != buf.len() { if start >= buf.len() || start + data.len() != buf.len() {
return Err(io::Error::new( return Err(io::Error::new(
io::ErrorKind::InvalidData, io::ErrorKind::InvalidData,
"InvalidData", "data invalid data",
)); ));
} }
buf[start..].copy_from_slice(data); buf[start..].copy_from_slice(data);
Ok(()) Ok(())
} }
} }
+11 -5
View File
@@ -1,6 +1,6 @@
use std::{fmt, io};
use std::net::Ipv4Addr;
use crate::protocol::body::ENCRYPTION_RESERVED; use crate::protocol::body::ENCRYPTION_RESERVED;
use std::net::Ipv4Addr;
use std::{fmt, io};
/* /*
0 15 31 0 15 31
@@ -21,9 +21,9 @@ pub const HEAD_LEN: usize = 12;
pub mod body; pub mod body;
pub mod control_packet; pub mod control_packet;
pub mod error_packet; pub mod error_packet;
pub mod service_packet;
pub mod ip_turn_packet; pub mod ip_turn_packet;
pub mod other_turn_packet; pub mod other_turn_packet;
pub mod service_packet;
#[derive(Eq, PartialEq, Copy, Clone, Debug)] #[derive(Eq, PartialEq, Copy, Clone, Debug)]
pub enum Version { pub enum Version {
@@ -230,7 +230,10 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> NetPacket<B> {
} }
pub fn set_payload(&mut self, payload: &[u8]) -> io::Result<()> { pub fn set_payload(&mut self, payload: &[u8]) -> io::Result<()> {
if self.data_len - 12 != payload.len() { if self.data_len - 12 != payload.len() {
return Err(io::Error::new(io::ErrorKind::InvalidData, "data_len - 12 != payload.len")); return Err(io::Error::new(
io::ErrorKind::InvalidData,
"data_len - 12 != payload.len",
));
} }
self.buffer.as_mut()[12..self.data_len].copy_from_slice(payload); self.buffer.as_mut()[12..self.data_len].copy_from_slice(payload);
Ok(()) Ok(())
@@ -240,7 +243,10 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> NetPacket<B> {
} }
pub fn set_data_len(&mut self, data_len: usize) -> io::Result<()> { pub fn set_data_len(&mut self, data_len: usize) -> io::Result<()> {
if data_len > self.buffer.as_ref().len() || data_len < 12 { if data_len > self.buffer.as_ref().len() || data_len < 12 {
return Err(io::Error::new(io::ErrorKind::InvalidData, "data_len invalid")); return Err(io::Error::new(
io::ErrorKind::InvalidData,
"data_len invalid",
));
} }
self.data_len = data_len; self.data_len = data_len;
Ok(()) Ok(())
+1 -1
View File
@@ -45,4 +45,4 @@ impl DeviceReader {
pub fn create(fd: i32) -> (DeviceWriter, DeviceReader) { pub fn create(fd: i32) -> (DeviceWriter, DeviceReader) {
(DeviceWriter(fd as _), DeviceReader(fd as _)) (DeviceWriter(fd as _), DeviceReader(fd as _))
} }
+77 -34
View File
@@ -1,18 +1,27 @@
use crate::tun_tap_device::linux_mac::DeviceW;
use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo};
use parking_lot::Mutex;
use std::io; use std::io;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo};
use tun::Device;
use parking_lot::Mutex;
use std::process::Command; use std::process::Command;
use std::sync::Arc; use std::sync::Arc;
use crate::tun_tap_device::linux_mac::DeviceW; use tun::Device;
pub const TUN_INTERFACE_NAME: &str = "vnt-tun";
pub const TAP_INTERFACE_NAME: &str = "vnt-tap";
impl DeviceWriter { impl DeviceWriter {
pub fn change_ip(&self, address: Ipv4Addr, netmask: Ipv4Addr, pub fn change_ip(
gateway: Ipv4Addr, _old_netmask: Ipv4Addr, _old_gateway: Ipv4Addr) -> io::Result<()> { &self,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
_old_netmask: Ipv4Addr,
_old_gateway: Ipv4Addr,
) -> io::Result<()> {
let mut config = tun::Configuration::default(); let mut config = tun::Configuration::default();
let broadcast_address = (!u32::from_be_bytes(netmask.octets())) let broadcast_address =
| u32::from_be_bytes(gateway.octets()); (!u32::from_be_bytes(netmask.octets())) | u32::from_be_bytes(gateway.octets());
let broadcast_address = Ipv4Addr::from(broadcast_address); let broadcast_address = Ipv4Addr::from(broadcast_address);
config config
.destination(gateway) .destination(gateway)
@@ -33,37 +42,45 @@ impl DeviceWriter {
// add_route(name, address, netmask)?; // add_route(name, address, netmask)?;
// 广播和组播路由 // 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?; add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(name, Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]))?; add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
return Ok(()); return Ok(());
} }
} }
pub fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> { pub fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
let route_add_str: String = format!( let route_add_str: String = format!("ip route add {:?}/{:?} dev {}", address, netmask, name);
"ip route add {:?}/{:?} dev {}",
address, netmask, name
);
let route_add_out = Command::new("sh") let route_add_out = Command::new("sh")
.arg("-c") .arg("-c")
.arg(&route_add_str) .arg(&route_add_str)
.output() .output()
.expect("sh exec error!"); .expect("sh exec error!");
if !route_add_out.status.success() { if !route_add_out.status.success() {
return Err(io::Error::new(io::ErrorKind::Other, format!("添加路由失败: cmd:{},out:{:?}", route_add_str, route_add_out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!(
"添加路由失败: cmd:{},out:{:?}",
route_add_str, route_add_out
),
));
} }
Ok(()) Ok(())
} }
pub fn create_device(device_type: DeviceType, pub fn create_device(
address: Ipv4Addr, device_type: DeviceType,
netmask: Ipv4Addr, address: Ipv4Addr,
gateway: Ipv4Addr, netmask: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, gateway: Ipv4Addr,
mtu: u16, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
) -> io::Result<(DeviceWriter, DeviceReader,DriverInfo)> { mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
let mut config = tun::Configuration::default(); let mut config = tun::Configuration::default();
let broadcast_address = (!u32::from_be_bytes(netmask.octets())) let broadcast_address =
| u32::from_be_bytes(gateway.octets()); (!u32::from_be_bytes(netmask.octets())) | u32::from_be_bytes(gateway.octets());
let broadcast_address = Ipv4Addr::from(broadcast_address); let broadcast_address = Ipv4Addr::from(broadcast_address);
config config
.destination(gateway) .destination(gateway)
@@ -74,12 +91,15 @@ pub fn create_device(device_type: DeviceType,
// .queues(2) 用多个队列有兼容性问题 // .queues(2) 用多个队列有兼容性问题
.up(); .up();
match device_type { match device_type {
DeviceType::Tun => {} DeviceType::Tun => {
config.name(TUN_INTERFACE_NAME);
}
DeviceType::Tap => { DeviceType::Tap => {
config.name(TAP_INTERFACE_NAME);
config.layer(tun::Layer::L2); config.layer(tun::Layer::L2);
} }
} }
let dev = tun::create(&config).unwrap(); let dev = tun::create(&config).expect("tun/tap failed to create");
let packet_information = dev.has_packet_information(); let packet_information = dev.has_packet_information();
let queue = dev.queue(0).unwrap(); let queue = dev.queue(0).unwrap();
let reader = queue.reader(); let reader = queue.reader();
@@ -92,11 +112,13 @@ pub fn create_device(device_type: DeviceType,
// add_route(name, address, netmask)?; // add_route(name, address, netmask)?;
// 广播和组播路由 // 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?; add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(name, Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]))?; add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
let device_w = match device_type { let device_w = match device_type {
DeviceType::Tun => { DeviceType::Tun => DeviceW::Tun(writer),
DeviceW::Tun(writer)
}
DeviceType::Tap => { DeviceType::Tap => {
let get_mac_cmd = format!("cat /sys/class/net/{}/address", name); let get_mac_cmd = format!("cat /sys/class/net/{}/address", name);
let mac_out = Command::new("sh") let mac_out = Command::new("sh")
@@ -105,7 +127,10 @@ pub fn create_device(device_type: DeviceType,
.output() .output()
.expect("sh exec error!"); .expect("sh exec error!");
if !mac_out.status.success() { if !mac_out.status.success() {
return Err(io::Error::new(io::ErrorKind::Other, format!("获取mac地址错误: {:?}", mac_out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("获取mac地址错误: {:?}", mac_out),
));
} }
let mac_str = String::from_utf8(mac_out.stdout).unwrap(); let mac_str = String::from_utf8(mac_out.stdout).unwrap();
let mut mac = [0; 6]; let mut mac = [0; 6];
@@ -118,15 +143,33 @@ pub fn create_device(device_type: DeviceType,
}; };
let driver_info = DriverInfo { let driver_info = DriverInfo {
device_type, device_type,
name:name.to_string(), name: name.to_string(),
version:String::new(), version: String::new(),
mac: None, mac: None,
}; };
Ok(( Ok((
DeviceWriter::new(device_w, Arc::new(Mutex::new(dev)), in_ips, address, packet_information), DeviceWriter::new(
device_w,
Arc::new(Mutex::new(dev)),
in_ips,
address,
packet_information,
),
DeviceReader::new(reader), DeviceReader::new(reader),
driver_info, driver_info,
)) ))
} }
pub fn delete_device(_device_type: DeviceType) {} pub fn delete_device(_device_type: DeviceType) {
for name in [TUN_INTERFACE_NAME, TAP_INTERFACE_NAME] {
let cmd = format!("ip link delete {}", name);
let delete_tun = Command::new("sh")
.arg("-c")
.arg(&cmd)
.output()
.expect("sh exec error!");
if !delete_tun.status.success() {
log::warn!("删除网卡失败:{:?}",delete_tun);
}
}
}
+25 -26
View File
@@ -2,15 +2,15 @@ use std::io;
use std::sync::Arc; use std::sync::Arc;
use bytes::BufMut; use bytes::BufMut;
use tun::platform::posix::{Reader, Writer}; use packet::ethernet;
use parking_lot::Mutex;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::os::unix::io::AsRawFd; use std::os::unix::io::AsRawFd;
#[cfg(any(target_os = "linux"))] #[cfg(any(target_os = "linux"))]
use tun::platform::linux::Device; use tun::platform::linux::Device;
#[cfg(any(target_os = "macos"))] #[cfg(any(target_os = "macos"))]
use tun::platform::macos::Device; use tun::platform::macos::Device;
use parking_lot::Mutex; use tun::platform::posix::{Reader, Writer};
use packet::ethernet;
use packet::ethernet::packet::EthernetPacket; use packet::ethernet::packet::EthernetPacket;
#[derive(Clone)] #[derive(Clone)]
@@ -22,12 +22,8 @@ pub enum DeviceW {
impl DeviceW { impl DeviceW {
pub fn is_tun(&self) -> bool { pub fn is_tun(&self) -> bool {
match self { match self {
DeviceW::Tun(_) => { DeviceW::Tun(_) => true,
true DeviceW::Tap(_) => false,
}
DeviceW::Tap(_) => {
false
}
} }
} }
} }
@@ -41,7 +37,13 @@ pub struct DeviceWriter {
} }
impl DeviceWriter { impl DeviceWriter {
pub fn new(writer: DeviceW,lock: Arc<Mutex<Device>>, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, _ip: Ipv4Addr, packet_information: bool) -> Self { pub fn new(
writer: DeviceW,
lock: Arc<Mutex<Device>>,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
_ip: Ipv4Addr,
packet_information: bool,
) -> Self {
Self { Self {
writer, writer,
lock, lock,
@@ -69,33 +71,30 @@ impl DeviceWriter {
///tun网卡写入ipv4数据 ///tun网卡写入ipv4数据
pub fn write_ipv4_tun(&self, buf: &[u8]) -> io::Result<()> { pub fn write_ipv4_tun(&self, buf: &[u8]) -> io::Result<()> {
match &self.writer { match &self.writer {
DeviceW::Tun(writer) => { DeviceW::Tun(writer) => Self::write(self.packet_information, writer, buf),
Self::write(self.packet_information, writer, buf) DeviceW::Tap(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
}
DeviceW::Tap(_) => {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
} }
} }
/// tap网卡写入以太网帧 /// tap网卡写入以太网帧
pub fn write_ethernet_tap(&self, buf: &[u8]) -> io::Result<()> { pub fn write_ethernet_tap(&self, buf: &[u8]) -> io::Result<()> {
match &self.writer { match &self.writer {
DeviceW::Tun(_) => { DeviceW::Tun(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
Err(io::Error::from(io::ErrorKind::Unsupported)) DeviceW::Tap((writer, _)) => Self::write(self.packet_information, writer, buf),
}
DeviceW::Tap((writer, _)) => {
Self::write(self.packet_information, writer, buf)
}
} }
} }
///写入ipv4数据,头部必须留14字节,给tap写入以太网帧头 ///写入ipv4数据,头部必须留14字节,给tap写入以太网帧头
pub fn write_ipv4(&self, buf: &mut [u8]) -> io::Result<()> { pub fn write_ipv4(&self, buf: &mut [u8]) -> io::Result<()> {
match &self.writer { match &self.writer {
DeviceW::Tun(writer) => { DeviceW::Tun(writer) => Self::write(self.packet_information, writer, &buf[14..]),
Self::write(self.packet_information, writer, &buf[14..])
}
DeviceW::Tap((writer, mac)) => { DeviceW::Tap((writer, mac)) => {
let source_mac = [buf[14 + 12], buf[14 + 13], buf[14 + 14], buf[14 + 15], !mac[5], 234]; let source_mac = [
buf[14 + 12],
buf[14 + 13],
buf[14 + 14],
buf[14 + 15],
!mac[5],
234,
];
let mut ethernet_packet = EthernetPacket::unchecked(buf); let mut ethernet_packet = EthernetPacket::unchecked(buf);
ethernet_packet.set_source(&source_mac); ethernet_packet.set_source(&source_mac);
ethernet_packet.set_destination(mac); ethernet_packet.set_destination(mac);
+56 -21
View File
@@ -1,15 +1,21 @@
use crate::tun_tap_device::linux_mac::DeviceW;
use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo};
use parking_lot::Mutex;
use std::io; use std::io;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo};
use tun::Device;
use parking_lot::Mutex;
use std::process::Command; use std::process::Command;
use std::sync::Arc; use std::sync::Arc;
use crate::tun_tap_device::linux_mac::DeviceW; use tun::Device;
impl DeviceWriter { impl DeviceWriter {
pub fn change_ip(&self, address: Ipv4Addr, netmask: Ipv4Addr, pub fn change_ip(
gateway: Ipv4Addr, _old_netmask: Ipv4Addr, _old_gateway: Ipv4Addr) -> io::Result<()> { &self,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
_old_netmask: Ipv4Addr,
_old_gateway: Ipv4Addr,
) -> io::Result<()> {
let mut config = tun::Configuration::default(); let mut config = tun::Configuration::default();
config config
.destination(gateway) .destination(gateway)
@@ -21,7 +27,7 @@ impl DeviceWriter {
return Err(io::Error::new(io::ErrorKind::Other, format!("{:?}", e))); return Err(io::Error::new(io::ErrorKind::Other, format!("{:?}", e)));
} }
if let Err(e) = config_ip(dev.name(), address, netmask, gateway) { if let Err(e) = config_ip(dev.name(), address, netmask, gateway) {
log::error!("{}",e); log::error!("{}", e);
} }
let name = dev.name(); let name = dev.name();
for (address, netmask) in &self.in_ips { for (address, netmask) in &self.in_ips {
@@ -31,17 +37,22 @@ impl DeviceWriter {
add_route(name, address, netmask)?; add_route(name, address, netmask)?;
// 广播和组播路由 // 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?; add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(name, Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]))?; add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
return Ok(()); return Ok(());
} }
} }
pub fn create_device(device_type: DeviceType, pub fn create_device(
address: Ipv4Addr, device_type: DeviceType,
netmask: Ipv4Addr, address: Ipv4Addr,
gateway: Ipv4Addr, netmask: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, gateway: Ipv4Addr,
mtu: u16, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> { ) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
match device_type { match device_type {
DeviceType::Tun => {} DeviceType::Tun => {}
@@ -68,7 +79,11 @@ pub fn create_device(device_type: DeviceType,
add_route(name, address, netmask)?; add_route(name, address, netmask)?;
// 广播和组播路由 // 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?; add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(name, Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]))?; add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
let packet_information = dev.has_packet_information(); let packet_information = dev.has_packet_information();
let queue = dev.queue(0).unwrap(); let queue = dev.queue(0).unwrap();
let reader = queue.reader(); let reader = queue.reader();
@@ -80,9 +95,15 @@ pub fn create_device(device_type: DeviceType,
mac: None, mac: None,
}; };
Ok(( Ok((
DeviceWriter::new(DeviceW::Tun(writer), Arc::new(Mutex::new(dev)), in_ips, address, packet_information), DeviceWriter::new(
DeviceW::Tun(writer),
Arc::new(Mutex::new(dev)),
in_ips,
address,
packet_information,
),
DeviceReader::new(reader), DeviceReader::new(reader),
driver_info driver_info,
)) ))
} }
@@ -97,12 +118,23 @@ fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()>
.output() .output()
.expect("sh exec error!"); .expect("sh exec error!");
if !route_add_out.status.success() { if !route_add_out.status.success() {
return Err(io::Error::new(io::ErrorKind::Other, format!("添加路由失败: cmd:{},out:{:?}", route_add_str, route_add_out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!(
"添加路由失败: cmd:{},out:{:?}",
route_add_str, route_add_out
),
));
} }
Ok(()) Ok(())
} }
fn config_ip(name: &str, address: Ipv4Addr, _netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> { fn config_ip(
name: &str,
address: Ipv4Addr,
_netmask: Ipv4Addr,
gateway: Ipv4Addr,
) -> io::Result<()> {
let up_eth_str: String = format!("ifconfig {} {:?} {:?} up ", name, address, gateway); let up_eth_str: String = format!("ifconfig {} {:?} {:?} up ", name, address, gateway);
let up_eth_out = Command::new("sh") let up_eth_out = Command::new("sh")
.arg("-c") .arg("-c")
@@ -110,9 +142,12 @@ fn config_ip(name: &str, address: Ipv4Addr, _netmask: Ipv4Addr, gateway: Ipv4Add
.output() .output()
.expect("sh exec error!"); .expect("sh exec error!");
if !up_eth_out.status.success() { if !up_eth_out.status.success() {
return Err(io::Error::new(io::ErrorKind::Other, format!("设置网络地址失败: cmd:{},out:{:?}", up_eth_str, up_eth_out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置网络地址失败: cmd:{},out:{:?}", up_eth_str, up_eth_out),
));
} }
Ok(()) Ok(())
} }
pub fn delete_device(_device_type: DeviceType) {} pub fn delete_device(_device_type: DeviceType) {}
+15 -16
View File
@@ -1,25 +1,24 @@
#[cfg(target_os = "windows")]
mod windows;
#[cfg(any(target_os = "linux"))]
mod linux;
#[cfg(target_os = "macos")]
mod mac;
#[cfg(any(target_os = "linux", target_os = "macos"))]
mod linux_mac;
#[cfg(target_os = "android")] #[cfg(target_os = "android")]
mod android; mod android;
#[cfg(any(target_os = "linux"))]
mod linux;
#[cfg(any(target_os = "linux", target_os = "macos"))]
mod linux_mac;
#[cfg(target_os = "macos")]
mod mac;
#[cfg(target_os = "windows")]
mod windows;
#[cfg(target_os = "android")]
pub use android::create;
#[cfg(target_os = "android")]
pub use android::{DeviceReader, DeviceWriter};
#[cfg(any(target_os = "linux"))] #[cfg(any(target_os = "linux"))]
pub use linux::create_device; pub use linux::create_device;
#[cfg(any(target_os = "linux"))] #[cfg(any(target_os = "linux"))]
pub use linux::delete_device; pub use linux::delete_device;
#[cfg(target_os = "android")]
pub use android::create;
#[cfg(any(target_os = "linux", target_os = "macos"))] #[cfg(any(target_os = "linux", target_os = "macos"))]
pub use linux_mac::{DeviceWriter, DeviceReader}; pub use linux_mac::{DeviceReader, DeviceWriter};
#[cfg(target_os = "android")]
pub use android::{DeviceWriter, DeviceReader};
#[cfg(target_os = "macos")] #[cfg(target_os = "macos")]
pub use mac::create_device; pub use mac::create_device;
#[cfg(target_os = "macos")] #[cfg(target_os = "macos")]
@@ -30,7 +29,7 @@ pub use windows::create_device;
#[cfg(target_os = "windows")] #[cfg(target_os = "windows")]
pub use windows::delete_device; pub use windows::delete_device;
#[cfg(target_os = "windows")] #[cfg(target_os = "windows")]
pub use windows::{DeviceWriter, DeviceReader}; pub use windows::{DeviceReader, DeviceWriter};
#[derive(Copy, Clone, Debug, Eq, PartialEq)] #[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum DeviceType { pub enum DeviceType {
@@ -50,4 +49,4 @@ pub struct DriverInfo {
pub name: String, pub name: String,
pub version: String, pub version: String,
pub mac: Option<String>, pub mac: Option<String>,
} }
+59 -68
View File
@@ -1,14 +1,14 @@
use std::{io, thread}; use crate::tun_tap_device::{DeviceType, DriverInfo};
use libloading::Library;
use packet::ethernet;
use packet::ethernet::packet::EthernetPacket;
use parking_lot::Mutex;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::os::windows::process::CommandExt; use std::os::windows::process::CommandExt;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use libloading::Library; use std::{io, thread};
use parking_lot::Mutex;
use packet::ethernet;
use packet::ethernet::packet::EthernetPacket;
use win_tun_tap::{IFace, TapDevice, TunDevice}; use win_tun_tap::{IFace, TapDevice, TunDevice};
use crate::tun_tap_device::{DriverInfo, DeviceType};
pub const TUN_INTERFACE_NAME: &str = "Vnt-Tun-V1"; pub const TUN_INTERFACE_NAME: &str = "Vnt-Tun-V1";
pub const TUN_POOL_NAME: &str = "Vnt-Tun-V1"; pub const TUN_POOL_NAME: &str = "Vnt-Tun-V1";
@@ -22,12 +22,8 @@ pub enum Device {
impl Device { impl Device {
pub fn is_tun(&self) -> bool { pub fn is_tun(&self) -> bool {
match self { match self {
Device::Tun(_) => { Device::Tun(_) => true,
true Device::Tap(_) => false,
}
Device::Tap(_) => {
false
}
} }
} }
} }
@@ -59,17 +55,13 @@ impl DeviceWriter {
dev.send_packet(packet); dev.send_packet(packet);
Ok(()) Ok(())
} }
Device::Tap(_) => { Device::Tap(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
Err(io::Error::from(io::ErrorKind::Unsupported))
}
} }
} }
/// tap网卡写入以太网帧 /// tap网卡写入以太网帧
pub fn write_ethernet_tap(&self, buf: &[u8]) -> io::Result<()> { pub fn write_ethernet_tap(&self, buf: &[u8]) -> io::Result<()> {
match self.device.as_ref() { match self.device.as_ref() {
Device::Tun(_) => { Device::Tun(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
Err(io::Error::from(io::ErrorKind::Unsupported))
}
Device::Tap((dev, _)) => { Device::Tap((dev, _)) => {
dev.write(buf)?; dev.write(buf)?;
Ok(()) Ok(())
@@ -85,7 +77,14 @@ impl DeviceWriter {
dev.send_packet(packet); dev.send_packet(packet);
} }
Device::Tap((dev, mac)) => { Device::Tap((dev, mac)) => {
let source_mac = [buf[14 + 12], buf[14 + 13], buf[14 + 14], buf[14 + 15], !mac[5], 234]; let source_mac = [
buf[14 + 12],
buf[14 + 13],
buf[14 + 14],
buf[14 + 15],
!mac[5],
234,
];
let mut ethernet_packet = EthernetPacket::unchecked(buf); let mut ethernet_packet = EthernetPacket::unchecked(buf);
ethernet_packet.set_source(&source_mac); ethernet_packet.set_source(&source_mac);
ethernet_packet.set_destination(mac); ethernet_packet.set_destination(mac);
@@ -105,16 +104,10 @@ impl DeviceWriter {
) -> io::Result<()> { ) -> io::Result<()> {
let _guard = self.lock.lock(); let _guard = self.lock.lock();
let dev: &dyn IFace = match self.device.as_ref() { let dev: &dyn IFace = match self.device.as_ref() {
Device::Tun(dev) => { Device::Tun(dev) => dev as &dyn IFace,
dev as &dyn IFace Device::Tap((dev, _)) => dev as &dyn IFace,
}
Device::Tap((dev, _)) => {
dev as &dyn IFace
}
}; };
if let Err(e) = if let Err(e) = dev.delete_route(dest(old_gateway, old_gateway), old_netmask, old_gateway) {
dev.delete_route(dest(old_gateway, old_gateway), old_netmask, old_gateway)
{
log::warn!("{:?}", e); log::warn!("{:?}", e);
} }
dev.set_ip(address, netmask)?; dev.set_ip(address, netmask)?;
@@ -125,18 +118,19 @@ impl DeviceWriter {
dev.add_route(address, netmask, gateway, 1)?; dev.add_route(address, netmask, gateway, 1)?;
// 广播和组播路由 // 广播和组播路由
dev.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?; dev.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?;
dev.add_route(Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]), gateway, 1)?; dev.add_route(
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
gateway,
1,
)?;
delete_cache(); delete_cache();
Ok(()) Ok(())
} }
pub fn close(&self) -> io::Result<()> { pub fn close(&self) -> io::Result<()> {
match self.device.as_ref() { match self.device.as_ref() {
Device::Tun(dev) => { Device::Tun(dev) => dev.shutdown(),
dev.shutdown() Device::Tap((dev, _)) => dev.shutdown(),
}
Device::Tap((dev, _)) => {
dev.shutdown()
}
} }
} }
pub fn is_tun(&self) -> bool { pub fn is_tun(&self) -> bool {
@@ -161,9 +155,7 @@ pub struct DeviceReader {
impl DeviceReader { impl DeviceReader {
pub fn new(device: Arc<Device>) -> Self { pub fn new(device: Arc<Device>) -> Self {
Self { Self { device }
device,
}
} }
} }
@@ -180,9 +172,7 @@ impl DeviceReader {
buf[..len].copy_from_slice(packet); buf[..len].copy_from_slice(packet);
Ok(len) Ok(len)
} }
Device::Tap((dev, _)) => { Device::Tap((dev, _)) => dev.read(buf),
dev.read(buf)
}
} }
} }
} }
@@ -224,10 +214,7 @@ fn create_tun(
) { ) {
Ok(tun_device) => tun_device, Ok(tun_device) => tun_device,
Err(e) => { Err(e) => {
return Err(io::Error::new( return Err(io::Error::new(io::ErrorKind::Other, format!("{:?}", e)));
io::ErrorKind::Other,
format!("{:?}", e),
));
} }
} }
} }
@@ -245,7 +232,12 @@ fn create_tun(
tun_device.add_route(address, netmask, gateway, 1)?; tun_device.add_route(address, netmask, gateway, 1)?;
// 广播和组播路由 // 广播和组播路由
tun_device.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?; tun_device.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?;
tun_device.add_route(Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]), gateway, 1)?; tun_device.add_route(
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
gateway,
1,
)?;
delete_cache(); delete_cache();
let device = Arc::new(Device::Tun(tun_device)); let device = Arc::new(Device::Tun(tun_device));
let driver_info = DriverInfo { let driver_info = DriverInfo {
@@ -257,7 +249,7 @@ fn create_tun(
Ok(( Ok((
DeviceWriter::new(device.clone(), in_ips, address), DeviceWriter::new(device.clone(), in_ips, address),
DeviceReader::new(device), DeviceReader::new(device),
driver_info driver_info,
)) ))
} }
} }
@@ -272,7 +264,7 @@ fn delete_cache() {
.output() .output()
.unwrap(); .unwrap();
if !out.status.success() { if !out.status.success() {
log::warn!("删除缓存失败:{:?}",out); log::warn!("删除缓存失败:{:?}", out);
} }
} }
@@ -318,7 +310,12 @@ fn create_tap(
} }
// 广播和组播路由 // 广播和组播路由
tap_device.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?; tap_device.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?;
tap_device.add_route(Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]), gateway, 1)?; tap_device.add_route(
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
gateway,
1,
)?;
delete_cache(); delete_cache();
let tap = Arc::new(Device::Tap((tap_device, mac))); let tap = Arc::new(Device::Tap((tap_device, mac)));
let driver_info = DriverInfo { let driver_info = DriverInfo {
@@ -330,7 +327,7 @@ fn create_tap(
Ok(( Ok((
DeviceWriter::new(tap.clone(), in_ips, address), DeviceWriter::new(tap.clone(), in_ips, address),
DeviceReader::new(tap), DeviceReader::new(tap),
driver_info driver_info,
)) ))
} }
@@ -344,29 +341,23 @@ fn delete_tap() {
let _ = tap_device.delete(); let _ = tap_device.delete();
} }
pub fn create_device(device_type: DeviceType, address: Ipv4Addr, pub fn create_device(
netmask: Ipv4Addr, device_type: DeviceType,
gateway: Ipv4Addr, address: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, netmask: Ipv4Addr,
mtu: u16) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> { gateway: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
match device_type { match device_type {
DeviceType::Tun => { DeviceType::Tun => create_tun(address, netmask, gateway, in_ips, mtu),
create_tun(address, netmask, gateway, in_ips, mtu) DeviceType::Tap => create_tap(address, netmask, gateway, in_ips, mtu),
}
DeviceType::Tap => {
create_tap(address, netmask, gateway, in_ips, mtu)
}
} }
} }
pub fn delete_device(device_type: DeviceType) { pub fn delete_device(device_type: DeviceType) {
match device_type { match device_type {
DeviceType::Tun => { DeviceType::Tun => delete_tun(),
delete_tun() DeviceType::Tap => delete_tap(),
}
DeviceType::Tap => {
delete_tap()
}
} }
} }
+1 -1
View File
@@ -1 +1 @@
pub mod wait; pub mod wait;
+2 -2
View File
@@ -1,5 +1,5 @@
use std::sync::Arc;
use std::sync::atomic::{AtomicIsize, Ordering}; use std::sync::atomic::{AtomicIsize, Ordering};
use std::sync::Arc;
use tokio::sync::watch::{channel, Receiver, Sender}; use tokio::sync::watch::{channel, Receiver, Sender};
#[derive(Clone)] #[derive(Clone)]
@@ -41,4 +41,4 @@ impl WaitGroup {
} }
} }
} }
} }
+17 -49
View File
@@ -21,8 +21,8 @@ use winapi::um::winioctl::*;
use winapi::um::winnt::*; use winapi::um::winnt::*;
use winapi::um::winreg::*; use winapi::um::winreg::*;
use std::{io, mem, ptr};
use std::error::Error; use std::error::Error;
use std::{io, mem, ptr};
use winapi::um::minwinbase::OVERLAPPED_u; use winapi::um::minwinbase::OVERLAPPED_u;
#[allow(non_camel_case_types)] #[allow(non_camel_case_types)]
@@ -46,9 +46,7 @@ pub fn string_from_guid(guid: &GUID) -> io::Result<Vec<WCHAR>> {
// GUID_STRING_CHARACTERS + 1 // GUID_STRING_CHARACTERS + 1
let mut string = vec![0; 39]; let mut string = vec![0; 39];
match unsafe { match unsafe { StringFromGUID2(guid, string.as_mut_ptr(), string.len() as _) } {
StringFromGUID2(guid, string.as_mut_ptr(), string.len() as _)
} {
0 => Err(io::Error::new(io::ErrorKind::Other, "Insufficent buffer")), 0 => Err(io::Error::new(io::ErrorKind::Other, "Insufficent buffer")),
_ => Ok(string), _ => Ok(string),
} }
@@ -85,12 +83,8 @@ pub fn luid_to_alias(luid: &NET_LUID) -> io::Result<Vec<WCHAR>> {
// IF_MAX_STRING_SIZE + 1 // IF_MAX_STRING_SIZE + 1
let mut alias = vec![0; 257]; let mut alias = vec![0; 257];
match unsafe { match unsafe { ConvertInterfaceLuidToAlias(luid, alias.as_mut_ptr(), alias.len()) } {
ConvertInterfaceLuidToAlias(luid, alias.as_mut_ptr(), alias.len()) 0 => Ok(alias),
} {
0 => {
Ok(alias)
}
err => Err(io::Error::from_raw_os_error(err as _)), err => Err(io::Error::from_raw_os_error(err as _)),
} }
} }
@@ -140,7 +134,8 @@ pub fn read_file(handle: HANDLE, buffer: &mut [u8]) -> io::Result<DWORD> {
buffer.as_mut_ptr() as _, buffer.as_mut_ptr() as _,
buffer.len() as _, buffer.len() as _,
&mut ret, &mut ret,
&mut ip_overlapped, ) { &mut ip_overlapped,
) {
let e = io::Error::last_os_error(); let e = io::Error::last_os_error();
if e.raw_os_error().unwrap_or(0) == 997 { if e.raw_os_error().unwrap_or(0) == 997 {
if 0 == GetOverlappedResult(handle, &mut ip_overlapped, &mut ret, 1) { if 0 == GetOverlappedResult(handle, &mut ip_overlapped, &mut ret, 1) {
@@ -191,9 +186,7 @@ pub fn create_device_info_list(guid: &GUID) -> io::Result<HDEVINFO> {
} }
pub fn get_class_devs(guid: &GUID, flags: DWORD) -> io::Result<HDEVINFO> { pub fn get_class_devs(guid: &GUID, flags: DWORD) -> io::Result<HDEVINFO> {
match unsafe { match unsafe { SetupDiGetClassDevsW(guid, ptr::null(), ptr::null_mut(), flags) } {
SetupDiGetClassDevsW(guid, ptr::null(), ptr::null_mut(), flags)
} {
INVALID_HANDLE_VALUE => Err(io::Error::last_os_error()), INVALID_HANDLE_VALUE => Err(io::Error::last_os_error()),
devinfo => Ok(devinfo), devinfo => Ok(devinfo),
} }
@@ -248,13 +241,8 @@ pub fn create_device_info(
} }
} }
pub fn set_selected_device( pub fn set_selected_device(devinfo: HDEVINFO, devinfo_data: &SP_DEVINFO_DATA) -> io::Result<()> {
devinfo: HDEVINFO, match unsafe { SetupDiSetSelectedDevice(devinfo, devinfo_data as *const _ as _) } {
devinfo_data: &SP_DEVINFO_DATA,
) -> io::Result<()> {
match unsafe {
SetupDiSetSelectedDevice(devinfo, devinfo_data as *const _ as _)
} {
0 => Err(io::Error::last_os_error()), 0 => Err(io::Error::last_os_error()),
_ => Ok(()), _ => Ok(()),
} }
@@ -308,13 +296,8 @@ pub fn build_driver_info_list(
devinfo_data: &SP_DEVINFO_DATA, devinfo_data: &SP_DEVINFO_DATA,
driver_type: DWORD, driver_type: DWORD,
) -> io::Result<()> { ) -> io::Result<()> {
match unsafe { match unsafe { SetupDiBuildDriverInfoList(devinfo, devinfo_data as *const _ as _, driver_type) }
SetupDiBuildDriverInfoList( {
devinfo,
devinfo_data as *const _ as _,
driver_type,
)
} {
0 => Err(io::Error::last_os_error()), 0 => Err(io::Error::last_os_error()),
_ => Ok(()), _ => Ok(()),
} }
@@ -326,11 +309,7 @@ pub fn destroy_driver_info_list(
driver_type: DWORD, driver_type: DWORD,
) -> io::Result<()> { ) -> io::Result<()> {
match unsafe { match unsafe {
SetupDiDestroyDriverInfoList( SetupDiDestroyDriverInfoList(devinfo, devinfo_data as *const _ as _, driver_type)
devinfo,
devinfo_data as *const _ as _,
driver_type,
)
} { } {
0 => Err(io::Error::last_os_error()), 0 => Err(io::Error::last_os_error()),
_ => Ok(()), _ => Ok(()),
@@ -342,8 +321,7 @@ pub fn get_driver_info_detail(
devinfo_data: &SP_DEVINFO_DATA, devinfo_data: &SP_DEVINFO_DATA,
drvinfo_data: &SP_DRVINFO_DATA_W, drvinfo_data: &SP_DRVINFO_DATA_W,
) -> io::Result<SP_DRVINFO_DETAIL_DATA_W2> { ) -> io::Result<SP_DRVINFO_DETAIL_DATA_W2> {
let mut drvinfo_detail: SP_DRVINFO_DETAIL_DATA_W2 = let mut drvinfo_detail: SP_DRVINFO_DETAIL_DATA_W2 = unsafe { mem::zeroed() };
unsafe { mem::zeroed() };
drvinfo_detail.cbSize = mem::size_of::<SP_DRVINFO_DETAIL_DATA_W>() as _; drvinfo_detail.cbSize = mem::size_of::<SP_DRVINFO_DETAIL_DATA_W>() as _;
match unsafe { match unsafe {
@@ -402,11 +380,7 @@ pub fn call_class_installer(
install_function: DI_FUNCTION, install_function: DI_FUNCTION,
) -> io::Result<()> { ) -> io::Result<()> {
match unsafe { match unsafe {
SetupDiCallClassInstaller( SetupDiCallClassInstaller(install_function, devinfo, devinfo_data as *const _ as _)
install_function,
devinfo,
devinfo_data as *const _ as _,
)
} { } {
0 => Err(io::Error::last_os_error()), 0 => Err(io::Error::last_os_error()),
_ => Ok(()), _ => Ok(()),
@@ -444,16 +418,12 @@ pub fn notify_change_key_value(
notify_filter: DWORD, notify_filter: DWORD,
milliseconds: DWORD, milliseconds: DWORD,
) -> io::Result<()> { ) -> io::Result<()> {
let event = match unsafe { let event = match unsafe { CreateEventW(ptr::null_mut(), FALSE, FALSE, ptr::null()) } {
CreateEventW(ptr::null_mut(), FALSE, FALSE, ptr::null())
} {
INVALID_HANDLE_VALUE => Err(io::Error::last_os_error()), INVALID_HANDLE_VALUE => Err(io::Error::last_os_error()),
event => Ok(event), event => Ok(event),
}?; }?;
match unsafe { match unsafe { RegNotifyChangeKeyValue(key, watch_subtree, notify_filter, event, TRUE) } {
RegNotifyChangeKeyValue(key, watch_subtree, notify_filter, event, TRUE)
} {
0 => Ok(()), 0 => Ok(()),
err => Err(io::Error::from_raw_os_error(err)), err => Err(io::Error::from_raw_os_error(err)),
}?; }?;
@@ -499,9 +469,7 @@ pub fn enum_device_info(
let mut devinfo_data: SP_DEVINFO_DATA = unsafe { mem::zeroed() }; let mut devinfo_data: SP_DEVINFO_DATA = unsafe { mem::zeroed() };
devinfo_data.cbSize = mem::size_of_val(&devinfo_data) as _; devinfo_data.cbSize = mem::size_of_val(&devinfo_data) as _;
match unsafe { match unsafe { SetupDiEnumDeviceInfo(devinfo, member_index, &mut devinfo_data) } {
SetupDiEnumDeviceInfo(devinfo, member_index, &mut devinfo_data)
} {
0 if unsafe { GetLastError() == ERROR_NO_MORE_ITEMS } => None, 0 if unsafe { GetLastError() == ERROR_NO_MORE_ITEMS } => None,
0 => Some(Err(io::Error::last_os_error())), 0 => Some(Err(io::Error::last_os_error())),
_ => Some(Ok(devinfo_data)), _ => Some(Ok(devinfo_data)),
+11 -9
View File
@@ -1,11 +1,11 @@
#![cfg(windows)] #![cfg(windows)]
mod tap;
mod tun;
mod ffi; mod ffi;
mod netsh; mod netsh;
mod route; mod route;
use std::{io, net}; mod tap;
mod tun;
use std::io;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
pub use tap::TapDevice; pub use tap::TapDevice;
pub use tun::*; pub use tun::*;
@@ -33,13 +33,15 @@ pub trait IFace {
/// 设置ip /// 设置ip
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()>; fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()>;
/// 设置路由 /// 设置路由
fn add_route(&self, dest: Ipv4Addr, fn add_route(
netmask: Ipv4Addr, &self,
gateway: Ipv4Addr, metric: u16) -> io::Result<()>; dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()>;
/// 删除路由 /// 删除路由
fn delete_route(&self, dest: Ipv4Addr, fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()>;
netmask: Ipv4Addr,
gateway: Ipv4Addr, ) -> io::Result<()>;
/// 设置最大传输单元 /// 设置最大传输单元
fn set_mtu(&self, mtu: u16) -> io::Result<()>; fn set_mtu(&self, mtu: u16) -> io::Result<()>;
/// 设置跃点 /// 设置跃点
+25 -10
View File
@@ -4,14 +4,17 @@ use std::os::windows::process::CommandExt;
/// 设置网卡名称 /// 设置网卡名称
pub fn set_interface_name(old_name: &str, new_name: &str) -> io::Result<()> { pub fn set_interface_name(old_name: &str, new_name: &str) -> io::Result<()> {
let cmd = format!(" netsh interface set interface name={:?} newname={:?}", old_name, new_name); let cmd = format!(
" netsh interface set interface name={:?} newname={:?}",
old_name, new_name
);
let out = std::process::Command::new("cmd") let out = std::process::Command::new("cmd")
.creation_flags(0x08000000) //winapi-0.3.9/src/um/winbase.rs:283 .creation_flags(0x08000000) //winapi-0.3.9/src/um/winbase.rs:283
.arg("/C") .arg("/C")
.arg(&cmd) .arg(&cmd)
.output()?; .output()?;
if !out.status.success() { if !out.status.success() {
log::warn!("修改网卡名称失败:cmd={:?},out={:?}",cmd,out); log::warn!("修改网卡名称失败:cmd={:?},out={:?}", cmd, out);
return Err(io::Error::new(io::ErrorKind::Other, "修改网卡名称失败")); return Err(io::Error::new(io::ErrorKind::Other, "修改网卡名称失败"));
} }
Ok(()) Ok(())
@@ -28,8 +31,11 @@ pub fn set_interface_ip(index: u32, address: &Ipv4Addr, netmask: &Ipv4Addr) -> i
.arg(&set_address) .arg(&set_address)
.output()?; .output()?;
if !out.status.success() { if !out.status.success() {
log::error!("cmd={:?},out={:?}",set_address,out); log::error!("cmd={:?},out={:?}", set_address, out);
return Err(io::Error::new(io::ErrorKind::Other, format!("设置网络地址失败: {:?}", out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置网络地址失败: {:?}", out),
));
} }
Ok(()) Ok(())
} }
@@ -45,21 +51,30 @@ pub fn set_interface_mtu(index: u32, mtu: u16) -> io::Result<()> {
.arg(&set_mtu) .arg(&set_mtu)
.output()?; .output()?;
if !out.status.success() { if !out.status.success() {
log::error!("cmd={:?},out={:?}",set_mtu,out); log::error!("cmd={:?},out={:?}", set_mtu, out);
return Err(io::Error::new(io::ErrorKind::Other, format!("设置mtu失败: {:?}", out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置mtu失败: {:?}", out),
));
} }
Ok(()) Ok(())
} }
pub fn set_interface_metric(index: u32, metric: u16) -> io::Result<()> { pub fn set_interface_metric(index: u32, metric: u16) -> io::Result<()> {
let set_metric = format!("netsh interface ip set interface {} metric={}", index,metric); let set_metric = format!(
"netsh interface ip set interface {} metric={}",
index, metric
);
let out = std::process::Command::new("cmd") let out = std::process::Command::new("cmd")
.creation_flags(0x08000000) .creation_flags(0x08000000)
.arg("/C") .arg("/C")
.arg(&set_metric) .arg(&set_metric)
.output()?; .output()?;
if !out.status.success() { if !out.status.success() {
log::error!("cmd={:?},out={:?}",set_metric,out); log::error!("cmd={:?},out={:?}", set_metric, out);
return Err(io::Error::new(io::ErrorKind::Other, format!("设置metric失败: {:?}", out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置metric失败: {:?}", out),
));
} }
Ok(()) Ok(())
} }
+27 -9
View File
@@ -3,9 +3,13 @@ use std::net::Ipv4Addr;
use std::os::windows::process::CommandExt; use std::os::windows::process::CommandExt;
/// 添加路由 /// 添加路由
pub fn add_route(index: u32, dest: Ipv4Addr, pub fn add_route(
netmask: Ipv4Addr, index: u32,
gateway: Ipv4Addr, metric: u16) -> io::Result<()> { dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()> {
let set_route = format!( let set_route = format!(
"route add {:?} mask {:?} {:?} metric {} if {}", "route add {:?} mask {:?} {:?} metric {} if {}",
dest, netmask, gateway, metric, index dest, netmask, gateway, metric, index
@@ -18,16 +22,27 @@ pub fn add_route(index: u32, dest: Ipv4Addr,
.output() .output()
.unwrap(); .unwrap();
if !out.status.success() { if !out.status.success() {
log::error!("cmd={:?},out={:?}",set_route,out); log::error!("cmd={:?},out={:?}", set_route, out);
return Err(io::Error::new(io::ErrorKind::Other, format!("添加路由失败: {:?}", out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("添加路由失败: {:?}", out),
));
} }
Ok(()) Ok(())
} }
/// 删除路由 /// 删除路由
pub fn delete_route(index: u32, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> { pub fn delete_route(
index: u32,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
) -> io::Result<()> {
if index == 0 { if index == 0 {
return Err(io::Error::new(io::ErrorKind::Other, format!("网络接口索引错误: {:?}", index))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("网络接口索引错误: {:?}", index),
));
} }
let delete_route = format!( let delete_route = format!(
"route delete {:?} mask {:?} {:?} if {}", "route delete {:?} mask {:?} {:?} if {}",
@@ -41,7 +56,10 @@ pub fn delete_route(index: u32, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4
.output() .output()
.unwrap(); .unwrap();
if !out.status.success() { if !out.status.success() {
return Err(io::Error::new(io::ErrorKind::Other, format!("删除路由失败: {:?}", out))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("删除路由失败: {:?}", out),
));
} }
Ok(()) Ok(())
} }
+35 -66
View File
@@ -51,22 +51,15 @@ pub fn create_interface() -> io::Result<NET_LUID> {
ffi::build_driver_info_list(devinfo, &devinfo_data, SPDIT_COMPATDRIVER)?; ffi::build_driver_info_list(devinfo, &devinfo_data, SPDIT_COMPATDRIVER)?;
let _guard = guard((), |_| { let _guard = guard((), |_| {
let _ = ffi::destroy_driver_info_list( let _ = ffi::destroy_driver_info_list(devinfo, &devinfo_data, SPDIT_COMPATDRIVER);
devinfo,
&devinfo_data,
SPDIT_COMPATDRIVER,
);
}); });
let mut driver_version = 0; let mut driver_version = 0;
let mut member_index = 0; let mut member_index = 0;
while let Some(drvinfo_data) = ffi::enum_driver_info( while let Some(drvinfo_data) =
devinfo, ffi::enum_driver_info(devinfo, &devinfo_data, SPDIT_COMPATDRIVER, member_index)
&devinfo_data, {
SPDIT_COMPATDRIVER,
member_index,
) {
member_index += 1; member_index += 1;
let drvinfo_data = match drvinfo_data { let drvinfo_data = match drvinfo_data {
@@ -78,14 +71,11 @@ pub fn create_interface() -> io::Result<NET_LUID> {
continue; continue;
} }
let drvinfo_detail = match ffi::get_driver_info_detail( let drvinfo_detail =
devinfo, match ffi::get_driver_info_detail(devinfo, &devinfo_data, &drvinfo_data) {
&devinfo_data, Ok(drvinfo_detail) => drvinfo_detail,
&drvinfo_data, _ => continue,
) { };
Ok(drvinfo_detail) => drvinfo_detail,
_ => continue,
};
let is_compatible = drvinfo_detail let is_compatible = drvinfo_detail
.HardwareID .HardwareID
@@ -115,16 +105,8 @@ pub fn create_interface() -> io::Result<NET_LUID> {
ffi::call_class_installer(devinfo, &devinfo_data, DIF_REGISTERDEVICE)?; ffi::call_class_installer(devinfo, &devinfo_data, DIF_REGISTERDEVICE)?;
let _ = ffi::call_class_installer( let _ = ffi::call_class_installer(devinfo, &devinfo_data, DIF_REGISTER_COINSTALLERS);
devinfo, let _ = ffi::call_class_installer(devinfo, &devinfo_data, DIF_INSTALLINTERFACES);
&devinfo_data,
DIF_REGISTER_COINSTALLERS,
);
let _ = ffi::call_class_installer(
devinfo,
&devinfo_data,
DIF_INSTALLINTERFACES,
);
ffi::call_class_installer(devinfo, &devinfo_data, DIF_INSTALLDEVICE)?; ffi::call_class_installer(devinfo, &devinfo_data, DIF_INSTALLDEVICE)?;
@@ -140,21 +122,11 @@ pub fn create_interface() -> io::Result<NET_LUID> {
let key = RegKey::predef(key); let key = RegKey::predef(key);
while let Err(_) = key.get_value::<DWORD, &str>("*IfType") { while let Err(_) = key.get_value::<DWORD, &str>("*IfType") {
ffi::notify_change_key_value( ffi::notify_change_key_value(key.raw_handle(), TRUE, REG_NOTIFY_CHANGE_NAME, 2000)?;
key.raw_handle(),
TRUE,
REG_NOTIFY_CHANGE_NAME,
2000,
)?;
} }
while let Err(_) = key.get_value::<DWORD, &str>("NetLuidIndex") { while let Err(_) = key.get_value::<DWORD, &str>("NetLuidIndex") {
ffi::notify_change_key_value( ffi::notify_change_key_value(key.raw_handle(), TRUE, REG_NOTIFY_CHANGE_NAME, 2000)?;
key.raw_handle(),
TRUE,
REG_NOTIFY_CHANGE_NAME,
2000,
)?;
} }
let if_type: DWORD = key.get_value("*IfType")?; let if_type: DWORD = key.get_value("*IfType")?;
@@ -181,8 +153,7 @@ pub fn check_interface(luid: &NET_LUID) -> io::Result<()> {
let mut member_index = 0; let mut member_index = 0;
while let Some(devinfo_data) = ffi::enum_device_info(devinfo, member_index) while let Some(devinfo_data) = ffi::enum_device_info(devinfo, member_index) {
{
member_index += 1; member_index += 1;
let devinfo_data = match devinfo_data { let devinfo_data = match devinfo_data {
@@ -190,14 +161,11 @@ pub fn check_interface(luid: &NET_LUID) -> io::Result<()> {
Err(_) => continue, Err(_) => continue,
}; };
let hardware_id = match ffi::get_device_registry_property( let hardware_id =
devinfo, match ffi::get_device_registry_property(devinfo, &devinfo_data, SPDRP_HARDWAREID) {
&devinfo_data, Ok(hardware_id) => hardware_id,
SPDRP_HARDWAREID, Err(_) => continue,
) { };
Ok(hardware_id) => hardware_id,
Err(_) => continue,
};
if !decode_utf16(&hardware_id).eq_ignore_ascii_case(HARDWARE_ID) { if !decode_utf16(&hardware_id).eq_ignore_ascii_case(HARDWARE_ID) {
continue; continue;
@@ -238,7 +206,10 @@ pub fn check_interface(luid: &NET_LUID) -> io::Result<()> {
return Ok(()); return Ok(());
} }
Err(io::Error::new(io::ErrorKind::NotFound, "TAP Device not found")) Err(io::Error::new(
io::ErrorKind::NotFound,
"TAP Device not found",
))
} }
/// Deletes an existing interface /// Deletes an existing interface
@@ -251,8 +222,7 @@ pub fn delete_interface(luid: &NET_LUID) -> io::Result<()> {
let mut member_index = 0; let mut member_index = 0;
while let Some(devinfo_data) = ffi::enum_device_info(devinfo, member_index) while let Some(devinfo_data) = ffi::enum_device_info(devinfo, member_index) {
{
member_index += 1; member_index += 1;
let devinfo_data = match devinfo_data { let devinfo_data = match devinfo_data {
@@ -260,14 +230,11 @@ pub fn delete_interface(luid: &NET_LUID) -> io::Result<()> {
Err(_) => continue, Err(_) => continue,
}; };
let hardware_id = match ffi::get_device_registry_property( let hardware_id =
devinfo, match ffi::get_device_registry_property(devinfo, &devinfo_data, SPDRP_HARDWAREID) {
&devinfo_data, Ok(hardware_id) => hardware_id,
SPDRP_HARDWAREID, Err(_) => continue,
) { };
Ok(hardware_id) => hardware_id,
Err(_) => continue,
};
if !decode_utf16(&hardware_id).eq_ignore_ascii_case(HARDWARE_ID) { if !decode_utf16(&hardware_id).eq_ignore_ascii_case(HARDWARE_ID) {
continue; continue;
@@ -308,13 +275,15 @@ pub fn delete_interface(luid: &NET_LUID) -> io::Result<()> {
return ffi::call_class_installer(devinfo, &devinfo_data, DIF_REMOVE); return ffi::call_class_installer(devinfo, &devinfo_data, DIF_REMOVE);
} }
Err(io::Error::new(io::ErrorKind::NotFound, "TAP Device not found")) Err(io::Error::new(
io::ErrorKind::NotFound,
"TAP Device not found",
))
} }
/// Open an handle to an interface /// Open an handle to an interface
pub fn open_interface(luid: &NET_LUID) -> io::Result<HANDLE> { pub fn open_interface(luid: &NET_LUID) -> io::Result<HANDLE> {
let guid = ffi::luid_to_guid(luid) let guid = ffi::luid_to_guid(luid).and_then(|guid| ffi::string_from_guid(&guid))?;
.and_then(|guid| ffi::string_from_guid(&guid))?;
let path = format!(r"\\.\Global\{}.tap", &decode_utf16(&guid)); let path = format!(r"\\.\Global\{}.tap", &decode_utf16(&guid));
@@ -323,6 +292,6 @@ pub fn open_interface(luid: &NET_LUID) -> io::Result<HANDLE> {
GENERIC_READ | GENERIC_WRITE, GENERIC_READ | GENERIC_WRITE,
FILE_SHARE_READ | FILE_SHARE_WRITE, FILE_SHARE_READ | FILE_SHARE_WRITE,
OPEN_EXISTING, OPEN_EXISTING,
FILE_ATTRIBUTE_SYSTEM | FILE_FLAG_OVERLAPPED,//FILE_ATTRIBUTE_SYSTEM, FILE_ATTRIBUTE_SYSTEM | FILE_FLAG_OVERLAPPED, //FILE_ATTRIBUTE_SYSTEM,
) )
} }
+24 -16
View File
@@ -1,11 +1,11 @@
use std::{io, time};
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use std::{io, time};
use winapi::shared::ifdef::NET_LUID; use winapi::shared::ifdef::NET_LUID;
use winapi::um::winioctl::*; use winapi::um::winioctl::*;
use winapi::um::winnt::HANDLE; use winapi::um::winnt::HANDLE;
use crate::{decode_utf16, encode_utf16, ffi, IFace, netsh, route}; use crate::{decode_utf16, encode_utf16, ffi, netsh, route, IFace};
mod iface; mod iface;
@@ -13,7 +13,6 @@ pub struct TapDevice {
index: u32, index: u32,
luid: NET_LUID, luid: NET_LUID,
handle: HANDLE, handle: HANDLE,
} }
unsafe impl Send for TapDevice {} unsafe impl Send for TapDevice {}
@@ -31,7 +30,7 @@ impl TapDevice {
&(), &(),
&mut mac, &mut mac,
) )
.map(|_| mac) .map(|_| mac)
} }
/// Retrieve the version of the driver /// Retrieve the version of the driver
@@ -44,7 +43,7 @@ impl TapDevice {
&(), &(),
&mut version, &mut version,
) )
.map(|_| version) .map(|_| version)
} }
/// Retieve the mtu of the interface /// Retieve the mtu of the interface
@@ -57,10 +56,9 @@ impl TapDevice {
&(), &(),
&mut mtu, &mut mtu,
) )
.map(|_| mtu) .map(|_| mtu)
} }
/// Set the status of the interface, true for connected, /// Set the status of the interface, true for connected,
/// false for disconnected. /// false for disconnected.
pub fn set_status(&self, status: bool) -> io::Result<()> { pub fn set_status(&self, status: bool) -> io::Result<()> {
@@ -98,7 +96,11 @@ impl TapDevice {
}; };
}; };
let index = ffi::luid_to_index(&luid).map(|index| index as u32)?; let index = ffi::luid_to_index(&luid).map(|index| index as u32)?;
Ok(Self { index, luid, handle }) Ok(Self {
index,
luid,
handle,
})
} }
pub fn open(name: &str) -> io::Result<Self> { pub fn open(name: &str) -> io::Result<Self> {
@@ -109,7 +111,11 @@ impl TapDevice {
let handle = iface::open_interface(&luid)?; let handle = iface::open_interface(&luid)?;
let index = ffi::luid_to_index(&luid).map(|index| index as u32)?; let index = ffi::luid_to_index(&luid).map(|index| index as u32)?;
Ok(Self { index, luid, handle }) Ok(Self {
index,
luid,
handle,
})
} }
pub fn delete(self) -> io::Result<()> { pub fn delete(self) -> io::Result<()> {
@@ -140,12 +146,18 @@ impl IFace for TapDevice {
netsh::set_interface_ip(index, &address, &mask) netsh::set_interface_ip(index, &address, &mask)
} }
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr, metric: u16) -> io::Result<()> { fn add_route(
&self,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()> {
let index = self.get_index()?; let index = self.get_index()?;
route::add_route(index, dest, netmask, gateway,metric) route::add_route(index, dest, netmask, gateway, metric)
} }
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> { fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> {
let index = self.get_index()?; let index = self.get_index()?;
route::delete_route(index, dest, netmask, gateway) route::delete_route(index, dest, netmask, gateway)
} }
@@ -161,7 +173,6 @@ impl IFace for TapDevice {
} }
} }
impl TapDevice { impl TapDevice {
pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> { pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
ffi::read_file(self.handle, buf).map(|res| res as _) ffi::read_file(self.handle, buf).map(|res| res as _)
@@ -177,6 +188,3 @@ impl Drop for TapDevice {
let _ = iface::delete_interface(&self.luid); let _ = iface::delete_interface(&self.luid);
} }
} }
+1 -1
View File
@@ -1,8 +1,8 @@
use log::*; use log::*;
use crate::tun::wintun_raw;
use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::atomic::{AtomicBool, Ordering};
use widestring::U16CStr; use widestring::U16CStr;
use crate::tun::wintun_raw;
/// Sets the logger wintun will use when logging. Maps to the WintunSetLogger C function /// Sets the logger wintun will use when logging. Maps to the WintunSetLogger C function
pub fn set_logger(win_tun: &wintun_raw::wintun, f: wintun_raw::WINTUN_LOGGER_CALLBACK) { pub fn set_logger(win_tun: &wintun_raw::wintun, f: wintun_raw::WINTUN_LOGGER_CALLBACK) {
+72 -31
View File
@@ -3,11 +3,11 @@ use std::net::Ipv4Addr;
use winapi::um::{handleapi, synchapi, winbase, winnt}; use winapi::um::{handleapi, synchapi, winbase, winnt};
use crate::{decode_utf16, encode_utf16, ffi, IFace, netsh, route}; use crate::{decode_utf16, encode_utf16, ffi, netsh, route, IFace};
use rand::Rng; use rand::Rng;
mod wintun_raw;
mod log; mod log;
pub mod packet; pub mod packet;
mod wintun_raw;
/// The maximum size of wintun's internal ring buffer (in bytes) /// The maximum size of wintun's internal ring buffer (in bytes)
pub const MAX_RING_CAPACITY: u32 = 0x400_0000; pub const MAX_RING_CAPACITY: u32 = 0x400_0000;
@@ -18,7 +18,6 @@ pub const MIN_RING_CAPACITY: u32 = 0x2_0000;
/// Maximum pool name length including zero terminator /// Maximum pool name length including zero terminator
pub const MAX_POOL: usize = 256; pub const MAX_POOL: usize = 256;
pub struct TunDevice { pub struct TunDevice {
pub(crate) luid: u64, pub(crate) luid: u64,
pub(crate) index: u32, pub(crate) index: u32,
@@ -38,7 +37,6 @@ pub struct TunDevice {
/// The adapter that owns this session /// The adapter that owns this session
pub(crate) adapter: wintun_raw::WINTUN_ADAPTER_HANDLE, pub(crate) adapter: wintun_raw::WINTUN_ADAPTER_HANDLE,
} }
unsafe impl Send for TunDevice {} unsafe impl Send for TunDevice {}
@@ -47,20 +45,31 @@ unsafe impl Sync for TunDevice {}
impl TunDevice { impl TunDevice {
pub unsafe fn create<L>(library: L, pool: &str, name: &str) -> io::Result<Self> pub unsafe fn create<L>(library: L, pool: &str, name: &str) -> io::Result<Self>
where L: Into<libloading::Library>, { where
L: Into<libloading::Library>,
{
let win_tun = match wintun_raw::wintun::from_library(library) { let win_tun = match wintun_raw::wintun::from_library(library) {
Ok(win_tun) => { win_tun } Ok(win_tun) => win_tun,
Err(e) => { Err(e) => {
return Err(io::Error::new(io::ErrorKind::Other, format!("library error {:?} ", e))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("library error {:?} ", e),
));
} }
}; };
let pool_utf16 = encode_utf16(pool); let pool_utf16 = encode_utf16(pool);
if pool_utf16.len() > MAX_POOL { if pool_utf16.len() > MAX_POOL {
return Err(io::Error::new(io::ErrorKind::Other, format!("长度大于{}:{:?}", MAX_POOL, pool))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("长度大于{}:{:?}", MAX_POOL, pool),
));
} }
let name_utf16 = encode_utf16(name); let name_utf16 = encode_utf16(name);
if name_utf16.len() > MAX_POOL { if name_utf16.len() > MAX_POOL {
return Err(io::Error::new(io::ErrorKind::Other, format!("长度大于{}:{:?}", MAX_POOL, pool))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("长度大于{}:{:?}", MAX_POOL, pool),
));
} }
let mut guid_bytes: [u8; 16] = [0u8; 16]; let mut guid_bytes: [u8; 16] = [0u8; 16];
rand::thread_rng().fill(&mut guid_bytes); rand::thread_rng().fill(&mut guid_bytes);
@@ -76,22 +85,32 @@ impl TunDevice {
//SAFETY: the function is loaded from the wintun dll properly, we are providing valid //SAFETY: the function is loaded from the wintun dll properly, we are providing valid
//pointers, and all the strings are correct null terminated UTF-16. This safety rationale //pointers, and all the strings are correct null terminated UTF-16. This safety rationale
//applies for all Wintun* functions below //applies for all Wintun* functions below
let adapter = win_tun.WintunCreateAdapter(pool_utf16.as_ptr(), name_utf16.as_ptr(), guid_ptr); let adapter =
win_tun.WintunCreateAdapter(pool_utf16.as_ptr(), name_utf16.as_ptr(), guid_ptr);
if adapter.is_null() { if adapter.is_null() {
return Err(io::Error::new(io::ErrorKind::Other, "Failed to crate adapter")); return Err(io::Error::new(
io::ErrorKind::Other,
"Failed to crate adapter",
));
} }
Self::init(win_tun, adapter) Self::init(win_tun, adapter)
} }
pub unsafe fn init(win_tun: wintun_raw::wintun, adapter: wintun_raw::WINTUN_ADAPTER_HANDLE) -> io::Result<Self> { pub unsafe fn init(
win_tun: wintun_raw::wintun,
adapter: wintun_raw::WINTUN_ADAPTER_HANDLE,
) -> io::Result<Self> {
// 开启session // 开启session
let session = win_tun.WintunStartSession(adapter, 128 * 1024); let session = win_tun.WintunStartSession(adapter, 128 * 1024);
if session.is_null() { if session.is_null() {
return Err(io::Error::new(io::ErrorKind::Other, "WintunStartSession failed")); return Err(io::Error::new(
io::ErrorKind::Other,
"WintunStartSession failed",
));
} }
//SAFETY: We follow the contract required by CreateEventA. See MSDN //SAFETY: We follow the contract required by CreateEventA. See MSDN
//(the pointers are allowed to be null, and 0 is okay for the others) //(the pointers are allowed to be null, and 0 is okay for the others)
let shutdown_event = synchapi::CreateEventA(std::ptr::null_mut(), let shutdown_event =
0, 0, std::ptr::null_mut()); synchapi::CreateEventA(std::ptr::null_mut(), 0, 0, std::ptr::null_mut());
let read_event = win_tun.WintunGetReadWaitEvent(session) as winnt::HANDLE; let read_event = win_tun.WintunGetReadWaitEvent(session) as winnt::HANDLE;
let mut luid: wintun_raw::NET_LUID = std::mem::zeroed(); let mut luid: wintun_raw::NET_LUID = std::mem::zeroed();
win_tun.WintunGetAdapterLUID(adapter, &mut luid as *mut wintun_raw::NET_LUID); win_tun.WintunGetAdapterLUID(adapter, &mut luid as *mut wintun_raw::NET_LUID);
@@ -107,18 +126,26 @@ impl TunDevice {
}) })
} }
pub unsafe fn delete_for_name<L>(library: L, name: &str) -> io::Result<()> pub unsafe fn delete_for_name<L>(library: L, name: &str) -> io::Result<()>
where L: Into<libloading::Library>, { where
L: Into<libloading::Library>,
{
let win_tun = match wintun_raw::wintun::from_library(library) { let win_tun = match wintun_raw::wintun::from_library(library) {
Ok(win_tun) => win_tun, Ok(win_tun) => win_tun,
Err(e) => { Err(e) => {
return Err(io::Error::new(io::ErrorKind::Other, format!("library error {:?} ", e))); return Err(io::Error::new(
io::ErrorKind::Other,
format!("library error {:?} ", e),
));
} }
}; };
log::set_default_logger_if_unset(&win_tun); log::set_default_logger_if_unset(&win_tun);
let name_utf16 = encode_utf16(name); let name_utf16 = encode_utf16(name);
let adapter = win_tun.WintunOpenAdapter(name_utf16.as_ptr()); let adapter = win_tun.WintunOpenAdapter(name_utf16.as_ptr());
if adapter.is_null() { if adapter.is_null() {
return Err(io::Error::new(io::ErrorKind::Other, "Failed to open adapter")); return Err(io::Error::new(
io::ErrorKind::Other,
"Failed to open adapter",
));
} }
win_tun.WintunCloseAdapter(adapter); win_tun.WintunCloseAdapter(adapter);
win_tun.WintunDeleteDriver(); win_tun.WintunDeleteDriver();
@@ -131,7 +158,10 @@ impl TunDevice {
pub fn version(&self) -> io::Result<Version> { pub fn version(&self) -> io::Result<Version> {
let version = unsafe { self.win_tun.WintunGetRunningDriverVersion() }; let version = unsafe { self.win_tun.WintunGetRunningDriverVersion() };
if version == 0 { if version == 0 {
return Err(io::Error::new(io::ErrorKind::Other, "WintunGetRunningDriverVersion")); return Err(io::Error::new(
io::ErrorKind::Other,
"WintunGetRunningDriverVersion",
));
} else { } else {
Ok(Version { Ok(Version {
major: ((version >> 16) & 0xFF) as u16, major: ((version >> 16) & 0xFF) as u16,
@@ -155,7 +185,6 @@ pub struct Version {
// } // }
// } // }
impl IFace for TunDevice { impl IFace for TunDevice {
fn shutdown(&self) -> io::Result<()> { fn shutdown(&self) -> io::Result<()> {
let _ = unsafe { synchapi::SetEvent(self.shutdown_event) }; let _ = unsafe { synchapi::SetEvent(self.shutdown_event) };
@@ -169,9 +198,7 @@ impl IFace for TunDevice {
fn get_name(&self) -> io::Result<String> { fn get_name(&self) -> io::Result<String> {
let luid = self.luid; let luid = self.luid;
ffi::luid_to_alias(&unsafe { std::mem::transmute(luid) }).map(|name| { ffi::luid_to_alias(&unsafe { std::mem::transmute(luid) }).map(|name| decode_utf16(&name))
decode_utf16(&name)
})
} }
fn set_name(&self, new_name: &str) -> io::Result<()> { fn set_name(&self, new_name: &str) -> io::Result<()> {
@@ -179,15 +206,21 @@ impl IFace for TunDevice {
netsh::set_interface_name(&name, new_name) netsh::set_interface_name(&name, new_name)
} }
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()>{ fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
netsh::set_interface_ip(self.get_index()?, &address, &mask) netsh::set_interface_ip(self.get_index()?, &address, &mask)
} }
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr, metric: u16) -> io::Result<()> { fn add_route(
&self,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()> {
route::add_route(self.get_index()?, dest, netmask, gateway, metric) route::add_route(self.get_index()?, dest, netmask, gateway, metric)
} }
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> { fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> {
route::delete_route(self.get_index()?, dest, netmask, gateway) route::delete_route(self.get_index()?, dest, netmask, gateway)
} }
@@ -257,14 +290,19 @@ impl TunDevice {
) )
}; };
match result { match result {
winbase::WAIT_FAILED => return Err(io::Error::new(io::ErrorKind::Other, "WAIT_FAILED")), winbase::WAIT_FAILED => {
return Err(io::Error::new(io::ErrorKind::Other, "WAIT_FAILED"))
}
_ => { _ => {
if result == winbase::WAIT_OBJECT_0 { if result == winbase::WAIT_OBJECT_0 {
//We have data! //We have data!
continue; continue;
} else if result == winbase::WAIT_OBJECT_0 + 1 { } else if result == winbase::WAIT_OBJECT_0 + 1 {
//Shutdown event triggered //Shutdown event triggered
return Err(io::Error::new(io::ErrorKind::Other, "Shutdown event triggered")); return Err(io::Error::new(
io::ErrorKind::Other,
"Shutdown event triggered",
));
} }
} }
} }
@@ -275,10 +313,14 @@ impl TunDevice {
impl TunDevice { impl TunDevice {
pub fn allocate_send_packet(&self, size: u16) -> io::Result<packet::TunPacket> { pub fn allocate_send_packet(&self, size: u16) -> io::Result<packet::TunPacket> {
let bytes_ptr = unsafe { let bytes_ptr = unsafe {
self.win_tun.WintunAllocateSendPacket(self.session, size as u32) self.win_tun
.WintunAllocateSendPacket(self.session, size as u32)
}; };
if bytes_ptr.is_null() { if bytes_ptr.is_null() {
Err(io::Error::new(io::ErrorKind::Other, "allocate_send_packet failed")) Err(io::Error::new(
io::ErrorKind::Other,
"allocate_send_packet failed",
))
} else { } else {
Ok(packet::TunPacket { Ok(packet::TunPacket {
kind: packet::Kind::SendPacketPending, kind: packet::Kind::SendPacketPending,
@@ -302,7 +344,6 @@ impl TunDevice {
} }
} }
impl Drop for TunDevice { impl Drop for TunDevice {
fn drop(&mut self) { fn drop(&mut self) {
//Close adapter on drop //Close adapter on drop
+6 -6
View File
@@ -1,4 +1,3 @@
use crate::TunDevice; use crate::TunDevice;
pub(crate) enum Kind { pub(crate) enum Kind {
@@ -12,7 +11,7 @@ pub(crate) enum Kind {
/// Represents a wintun packet /// Represents a wintun packet
pub struct TunPacket<'a> { pub struct TunPacket<'a> {
pub(crate) kind: Kind, pub(crate) kind: Kind,
pub(crate) size:usize, pub(crate) size: usize,
pub(crate) bytes_ptr: *const u8, pub(crate) bytes_ptr: *const u8,
//Share ownership of session to prevent the session from being dropped before packets that //Share ownership of session to prevent the session from being dropped before packets that
@@ -20,7 +19,7 @@ pub struct TunPacket<'a> {
pub(crate) tun_device: Option<&'a TunDevice>, pub(crate) tun_device: Option<&'a TunDevice>,
} }
impl <'a>TunPacket<'a> { impl<'a> TunPacket<'a> {
/// Returns the bytes this packet holds as &mut. /// Returns the bytes this packet holds as &mut.
/// The lifetime of the bytes is tied to the lifetime of this packet. /// The lifetime of the bytes is tied to the lifetime of this packet.
pub fn bytes_mut(&mut self) -> &mut [u8] { pub fn bytes_mut(&mut self) -> &mut [u8] {
@@ -30,11 +29,11 @@ impl <'a>TunPacket<'a> {
/// Returns an immutable reference to the bytes this packet holds. /// Returns an immutable reference to the bytes this packet holds.
/// The lifetime of the bytes is tied to the lifetime of this packet. /// The lifetime of the bytes is tied to the lifetime of this packet.
pub fn bytes(&self) -> &[u8] { pub fn bytes(&self) -> &[u8] {
unsafe { std::slice::from_raw_parts(self.bytes_ptr,self.size) } unsafe { std::slice::from_raw_parts(self.bytes_ptr, self.size) }
} }
} }
impl <'a>Drop for TunPacket<'a> { impl<'a> Drop for TunPacket<'a> {
fn drop(&mut self) { fn drop(&mut self) {
match self.kind { match self.kind {
Kind::ReceivePacket => { Kind::ReceivePacket => {
@@ -46,7 +45,8 @@ impl <'a>Drop for TunPacket<'a> {
// ring buffer that the wintun session owns. We return that region of // ring buffer that the wintun session owns. We return that region of
// memory back to wintun here // memory back to wintun here
let tun_device = self.tun_device.unwrap(); let tun_device = self.tun_device.unwrap();
tun_device.win_tun tun_device
.win_tun
.WintunReleaseReceivePacket(tun_device.session, self.bytes_ptr) .WintunReleaseReceivePacket(tun_device.session, self.bytes_ptr)
}; };
} }
+35 -35
View File
@@ -11,8 +11,8 @@ impl<Storage> __BindgenBitfieldUnit<Storage> {
} }
} }
impl<Storage> __BindgenBitfieldUnit<Storage> impl<Storage> __BindgenBitfieldUnit<Storage>
where where
Storage: AsRef<[u8]> + AsMut<[u8]>, Storage: AsRef<[u8]> + AsMut<[u8]>,
{ {
#[inline] #[inline]
pub fn get_bit(&self, index: usize) -> bool { pub fn get_bit(&self, index: usize) -> bool {
@@ -112,40 +112,40 @@ fn bindgen_test_layout__GUID() {
unsafe { &(*(::std::ptr::null::<_GUID>())).Data1 as *const _ as usize }, unsafe { &(*(::std::ptr::null::<_GUID>())).Data1 as *const _ as usize },
0usize, 0usize,
concat!( concat!(
"Offset of field: ", "Offset of field: ",
stringify!(_GUID), stringify!(_GUID),
"::", "::",
stringify!(Data1) stringify!(Data1)
) )
); );
assert_eq!( assert_eq!(
unsafe { &(*(::std::ptr::null::<_GUID>())).Data2 as *const _ as usize }, unsafe { &(*(::std::ptr::null::<_GUID>())).Data2 as *const _ as usize },
4usize, 4usize,
concat!( concat!(
"Offset of field: ", "Offset of field: ",
stringify!(_GUID), stringify!(_GUID),
"::", "::",
stringify!(Data2) stringify!(Data2)
) )
); );
assert_eq!( assert_eq!(
unsafe { &(*(::std::ptr::null::<_GUID>())).Data3 as *const _ as usize }, unsafe { &(*(::std::ptr::null::<_GUID>())).Data3 as *const _ as usize },
6usize, 6usize,
concat!( concat!(
"Offset of field: ", "Offset of field: ",
stringify!(_GUID), stringify!(_GUID),
"::", "::",
stringify!(Data3) stringify!(Data3)
) )
); );
assert_eq!( assert_eq!(
unsafe { &(*(::std::ptr::null::<_GUID>())).Data4 as *const _ as usize }, unsafe { &(*(::std::ptr::null::<_GUID>())).Data4 as *const _ as usize },
8usize, 8usize,
concat!( concat!(
"Offset of field: ", "Offset of field: ",
stringify!(_GUID), stringify!(_GUID),
"::", "::",
stringify!(Data4) stringify!(Data4)
) )
); );
} }
@@ -248,20 +248,20 @@ fn bindgen_test_layout__NET_LUID_LH() {
unsafe { &(*(::std::ptr::null::<_NET_LUID_LH>())).Value as *const _ as usize }, unsafe { &(*(::std::ptr::null::<_NET_LUID_LH>())).Value as *const _ as usize },
0usize, 0usize,
concat!( concat!(
"Offset of field: ", "Offset of field: ",
stringify!(_NET_LUID_LH), stringify!(_NET_LUID_LH),
"::", "::",
stringify!(Value) stringify!(Value)
) )
); );
assert_eq!( assert_eq!(
unsafe { &(*(::std::ptr::null::<_NET_LUID_LH>())).Info as *const _ as usize }, unsafe { &(*(::std::ptr::null::<_NET_LUID_LH>())).Info as *const _ as usize },
0usize, 0usize,
concat!( concat!(
"Offset of field: ", "Offset of field: ",
stringify!(_NET_LUID_LH), stringify!(_NET_LUID_LH),
"::", "::",
stringify!(Info) stringify!(Info)
) )
); );
} }
@@ -310,33 +310,33 @@ pub struct wintun {
pub WintunCloseAdapter: unsafe extern "C" fn(arg1: WINTUN_ADAPTER_HANDLE), pub WintunCloseAdapter: unsafe extern "C" fn(arg1: WINTUN_ADAPTER_HANDLE),
pub WintunOpenAdapter: unsafe extern "C" fn(arg1: LPCWSTR) -> WINTUN_ADAPTER_HANDLE, pub WintunOpenAdapter: unsafe extern "C" fn(arg1: LPCWSTR) -> WINTUN_ADAPTER_HANDLE,
pub WintunGetAdapterLUID: pub WintunGetAdapterLUID:
unsafe extern "C" fn(arg1: WINTUN_ADAPTER_HANDLE, arg2: *mut NET_LUID), unsafe extern "C" fn(arg1: WINTUN_ADAPTER_HANDLE, arg2: *mut NET_LUID),
pub WintunGetRunningDriverVersion: unsafe extern "C" fn() -> DWORD, pub WintunGetRunningDriverVersion: unsafe extern "C" fn() -> DWORD,
pub WintunDeleteDriver: unsafe extern "C" fn() -> BOOL, pub WintunDeleteDriver: unsafe extern "C" fn() -> BOOL,
pub WintunSetLogger: unsafe extern "C" fn(arg1: WINTUN_LOGGER_CALLBACK), pub WintunSetLogger: unsafe extern "C" fn(arg1: WINTUN_LOGGER_CALLBACK),
pub WintunStartSession: pub WintunStartSession:
unsafe extern "C" fn(arg1: WINTUN_ADAPTER_HANDLE, arg2: DWORD) -> WINTUN_SESSION_HANDLE, unsafe extern "C" fn(arg1: WINTUN_ADAPTER_HANDLE, arg2: DWORD) -> WINTUN_SESSION_HANDLE,
pub WintunEndSession: unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE), pub WintunEndSession: unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE),
pub WintunGetReadWaitEvent: unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE) -> HANDLE, pub WintunGetReadWaitEvent: unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE) -> HANDLE,
pub WintunReceivePacket: pub WintunReceivePacket:
unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: *mut DWORD) -> *mut BYTE, unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: *mut DWORD) -> *mut BYTE,
pub WintunReleaseReceivePacket: pub WintunReleaseReceivePacket:
unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: *const BYTE), unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: *const BYTE),
pub WintunAllocateSendPacket: pub WintunAllocateSendPacket:
unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: DWORD) -> *mut BYTE, unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: DWORD) -> *mut BYTE,
pub WintunSendPacket: unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: *const BYTE), pub WintunSendPacket: unsafe extern "C" fn(arg1: WINTUN_SESSION_HANDLE, arg2: *const BYTE),
} }
impl wintun { impl wintun {
pub unsafe fn new<P>(path: P) -> Result<Self, ::libloading::Error> pub unsafe fn new<P>(path: P) -> Result<Self, ::libloading::Error>
where where
P: AsRef<::std::ffi::OsStr>, P: AsRef<::std::ffi::OsStr>,
{ {
let library = ::libloading::Library::new(path)?; let library = ::libloading::Library::new(path)?;
Self::from_library(library) Self::from_library(library)
} }
pub unsafe fn from_library<L>(library: L) -> Result<Self, ::libloading::Error> pub unsafe fn from_library<L>(library: L) -> Result<Self, ::libloading::Error>
where where
L: Into<::libloading::Library>, L: Into<::libloading::Library>,
{ {
let __library = library.into(); let __library = library.into();
let WintunCreateAdapter = __library.get(b"WintunCreateAdapter\0").map(|sym| *sym)?; let WintunCreateAdapter = __library.get(b"WintunCreateAdapter\0").map(|sym| *sym)?;