diff --git a/Cargo.lock b/Cargo.lock index cd67c2e..c8a796b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -195,6 +195,10 @@ name = "cc" version = "1.0.94" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "17f6e324229dc011159fcc089755d1e2e216a90d43a7dea6853ca740b84f35e7" +dependencies = [ + "jobserver", + "libc", +] [[package]] name = "cesu8" @@ -208,6 +212,36 @@ version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "baf1de4339761588bc0619e3cbc0120ee582ebb74b53b4efbf79117bd2da40fd" +[[package]] +name = "cfg_aliases" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" + +[[package]] +name = "chacha20" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures", +] + +[[package]] +name = "chacha20poly1305" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10cd79432192d1c0f4e1a0fef9527696cc039165d729fb41b3f4f4f354c2dc35" +dependencies = [ + "aead", + "chacha20", + "cipher", + "poly1305", + "zeroize", +] + [[package]] name = "chrono" version = "0.4.38" @@ -230,6 +264,7 @@ checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ "crypto-common", "inout", + "zeroize", ] [[package]] @@ -244,7 +279,7 @@ dependencies = [ [[package]] name = "common" -version = "1.2.9" +version = "1.2.10" [[package]] name = "console" @@ -606,6 +641,15 @@ version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8eaf4bc02d17cbdd7ff4c7438cafcdf7fb9a4613313ad11b4f8fefe7d3fa0130" +[[package]] +name = "jobserver" +version = "0.1.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2b099aaa34a9751c5bf0878add70444e1ed2dd73f347be99003d4577277de6e" +dependencies = [ + "libc", +] + [[package]] name = "js-sys" version = "0.3.69" @@ -720,6 +764,12 @@ dependencies = [ "winapi", ] +[[package]] +name = "lz4_flex" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75761162ae2b0e580d7e7c390558127e5f01b4194debd6221fd8c207fc80e3f5" + [[package]] name = "memchr" version = "2.7.2" @@ -948,6 +998,17 @@ version = "0.3.30" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d231b230927b5e4ad203db57bbcbee2802f6bce620b1e4a9024a07d94e2907ec" +[[package]] +name = "poly1305" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf" +dependencies = [ + "cpufeatures", + "opaque-debug", + "universal-hash", +] + [[package]] name = "polyval" version = "0.6.2" @@ -1305,6 +1366,16 @@ dependencies = [ "digest", ] +[[package]] +name = "signal-hook" +version = "0.3.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8621587d4798caf8eb44879d42e56b9a93ea5dcd315a6487c357130095b62801" +dependencies = [ + "libc", + "signal-hook-registry", +] + [[package]] name = "signal-hook-registry" version = "1.4.2" @@ -1586,13 +1657,16 @@ checksum = "49874b5167b65d7193b8aba1567f5c7d93d001cafc34600cee003eda787e483f" [[package]] name = "vnt" -version = "1.2.9" +version = "1.2.10" dependencies = [ "aes", "aes-gcm", "anyhow", "bytes", "cbc", + "cfg_aliases", + "chacha20", + "chacha20poly1305", "crossbeam-epoch", "crossbeam-queue", "crossbeam-utils", @@ -1602,6 +1676,7 @@ dependencies = [ "libloading", "libsm", "log", + "lz4_flex", "mio", "openssl-sys", "packet", @@ -1619,11 +1694,12 @@ dependencies = [ "thiserror", "tokio", "tun", + "zstd", ] [[package]] name = "vnt-cli" -version = "1.2.9" +version = "1.2.10" dependencies = [ "anyhow", "chrono", @@ -1637,6 +1713,7 @@ dependencies = [ "rand", "serde", "serde_yaml", + "signal-hook", "sudo", "uuid", "vnt", @@ -1645,7 +1722,7 @@ dependencies = [ [[package]] name = "vnt-jni" -version = "1.2.9" +version = "1.2.10" dependencies = [ "android_logger", "common", @@ -2003,3 +2080,31 @@ name = "zeroize" version = "1.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "525b4ec142c6b68a2d10f01f7bbf6755599ca3f81ea53b8431b7dd348f5fdb2d" + +[[package]] +name = "zstd" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d789b1514203a1120ad2429eae43a7bd32b90976a7bb8a05f7ec02fa88cc23a" +dependencies = [ + "zstd-safe", +] + +[[package]] +name = "zstd-safe" +version = "7.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1cd99b45c6bc03a018c8b8a86025678c87e55526064e38f9df301989dce7ec0a" +dependencies = [ + "zstd-sys", +] + +[[package]] +name = "zstd-sys" +version = "2.0.10+zstd.1.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c253a4914af5bafc8fa8c86ee400827e83cf6ec01195ec1f1ed8441bf00d65aa" +dependencies = [ + "cc", + "pkg-config", +] diff --git a/README.md b/README.md index 7017090..26a92ba 100644 --- a/README.md +++ b/README.md @@ -73,21 +73,24 @@ cargo build -p vnt-cli --no-default-features features说明 -| feature | 说明 | 是否默认 | -|------------------|----------------------|------| -| openssl | 使用openssl中的aes_ecb算法 | 否 | -| openssl-vendored | 从源码编译openssl | 否 | -| ring-cipher | 使用ring中的aes_gcm算法 | 否 | -| aes_cbc | 支持aes_cbc加密 | 是 | -| aes_ecb | 支持aes_ecb加密 | 是 | -| aes_gcm | 支持aes_gcm加密 | 是 | -| sm4_cbc | 支持sm4_cbc加密 | 是 | -| server_encrypt | 支持服务端加密 | 是 | -| ip_proxy | 内置ip代理 | 是 | -| port_mapping | 端口映射 | 是 | -| log | 日志 | 是 | -| command | list、route等命令 | 是 | -| file_config | yaml配置文件 | 是 | +| feature | 说明 | 是否默认 | +|-------------------|--------------------------------|------| +| openssl | 使用openssl中的加密算法 | 否 | +| openssl-vendored | 从源码编译openssl | 否 | +| ring-cipher | 使用ring中的加密算法 | 否 | +| aes_cbc | 支持aes_cbc加密 | 是 | +| aes_ecb | 支持aes_ecb加密 | 是 | +| aes_gcm | 支持aes_gcm加密 | 是 | +| sm4_cbc | 支持sm4_cbc加密 | 是 | +| chacha20_poly1305 | 支持chacha20和chacha20_poly1305加密 | 是 | +| server_encrypt | 支持服务端加密 | 是 | +| ip_proxy | 内置ip代理 | 是 | +| port_mapping | 端口映射 | 是 | +| log | 日志 | 是 | +| command | list、route等命令 | 是 | +| file_config | yaml配置文件 | 是 | +| lz4 | lz4压缩 | 是 | +| zstd | zstd压缩 | 否 | ### ip转发/代理 @@ -171,12 +174,14 @@ sudo pfctl -f /etc/pf.conf -e - Mac - Linux - - Arch Linux `yay -Syu vnt` - Windows - 默认使用tun网卡 依赖wintun.dll([win-tun](https://www.wintun.net/))(将dll放到同目录下,建议使用版本0.14.1) - 使用tap网卡 依赖tap-windows([win-tap](https://build.openvpn.net/downloads/releases/))(建议使用版本9.24.7) - Android - - [VntApp](https://github.com/lbl8603/VntApp) + +### GUI + +支持安卓和Windows [下载](https://github.com/lbl8603/VntApp/releases/) ### 特性 @@ -280,13 +285,20 @@ vnt默认使用10.26.0.0/24网段,和本地网络适配器的ip冲突 ### 交流群 +对VNT有任何问题均可以加群联系作者 + QQ: 1034868233 +### 赞助 +如果VNT对你有帮助,欢迎打赏作者 + + ### 其他 可使用社区小伙伴搭建的中继服务器 1. -s vnt.8443.eu.org:29871 +2. -s vnt.wherewego.top:29872 ### 参与贡献 diff --git a/common/Cargo.toml b/common/Cargo.toml index 72c4ae8..16e39e1 100644 --- a/common/Cargo.toml +++ b/common/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "common" -version = "1.2.9" +version = "1.2.10" edition = "2021" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html diff --git a/vnt-cli/Cargo.toml b/vnt-cli/Cargo.toml index e591254..dfe9120 100644 --- a/vnt-cli/Cargo.toml +++ b/vnt-cli/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "vnt-cli" -version = "1.2.9" +version = "1.2.10" edition = "2021" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html @@ -24,12 +24,13 @@ features = [ [target.'cfg(any(target_os = "linux",target_os = "macos"))'.dependencies] sudo = "0.6.0" +signal-hook = "0.3.17" [target.'cfg(target_os = "windows")'.dependencies] winapi = { version = "0.3.9", features = ["handleapi", "processthreadsapi", "winnt", "securitybaseapi", "impl-default"] } [features] -default = ["server_encrypt", "aes_gcm", "aes_cbc", "aes_ecb", "sm4_cbc", "ip_proxy", "port_mapping", "log", "command", "file_config"] +default = ["server_encrypt", "aes_gcm", "aes_cbc", "aes_ecb", "sm4_cbc", "chacha20_poly1305", "ip_proxy", "port_mapping", "log", "command", "file_config", "lz4"] openssl = ["vnt/openssl"] openssl-vendored = ["vnt/openssl-vendored"] ring-cipher = ["vnt/ring-cipher"] @@ -37,9 +38,12 @@ aes_cbc = ["vnt/aes_cbc"] aes_ecb = ["vnt/aes_ecb"] sm4_cbc = ["vnt/sm4_cbc"] aes_gcm = ["vnt/aes_gcm"] +chacha20_poly1305 = ["vnt/chacha20_poly1305"] server_encrypt = ["vnt/server_encrypt"] ip_proxy = ["vnt/ip_proxy"] port_mapping = ["vnt/port_mapping"] +lz4 = ["vnt/lz4_compress"] +zstd = ["vnt/zstd_compress"] log = ["log4rs"] command = [] file_config = [] diff --git a/vnt-cli/README.md b/vnt-cli/README.md index 97c5ea2..34f9d55 100644 --- a/vnt-cli/README.md +++ b/vnt-cli/README.md @@ -1,33 +1,54 @@ ## 模块介绍 + 体积小,可以在服务器、路由器等环境使用 + ## 详细参数说明 + ### -k `` + 一个虚拟局域网的标识,在同一服务器下,相同token的设备会组建一个局域网 + ### -n `` + 设备名称,方便区分不同设备 + ### -d `` + 设备id,每台设备的唯一标识,注意不要重复 + ### -c + 关闭控制台交互式命令,后台运行时可以加此参数 + ### -s `` + 注册和中继服务器地址,注册和转发数据,以'TXT:'开头表示解析TXT记录,TXT记录内容必须是'host:port'形式的服务器地址 + ### -e `` + 使用stun服务探测客户端NAT类型,不同类型有不同的打洞策略 + ### -a + 加了此参数表示使用tap网卡,默认使用tun网卡,tun网卡效率更高 注意:仅在windows上支持使用tap,用于兼容低版本windows系统(低版本windows不支持wintun) + ### --nic `` + 指定虚拟网卡名称,默认tun模式使用vnt-tun,tap模式使用vnt-tap + ### -i ``、-o `` -配置点对网(IP代理)时使用,例如A(虚拟ip:10.26.0.2)通过B(虚拟ip:10.26.0.3,本地出口ip:192.168.0.10)访问C(目标网段192.168.0.0/24), +配置点对网(IP代理)时使用,例如A(虚拟ip:10.26.0.2)通过B(虚拟ip:10.26.0.3,本地出口ip:192.168.0.10)访问C( +目标网段192.168.0.0/24), 则在A配置 **'-i 192.168.0.0/24,10.26.0.3'** ,表示将192.168.0.0/24网段的数据都转发到10.26.0.3节点 在B配置 **'-o 192.168.0.0/24'** ,表示允许将数据转发到 192.168.0.0/24 ,允许转发所有网段可以使用 **'-o 0.0.0.0/0'** --i和-o参数均可使用多次,来指定不同网段,例如 **'-o 192.168.1.0/24 -o 192.168.2.0/24'** 表示允许转发目标为192.168.1.0/24或192.168.2.0/24这两个网段的数据 +-i和-o参数均可使用多次,来指定不同网段,例如 **'-o 192.168.1.0/24 -o 192.168.2.0/24'** +表示允许转发目标为192.168.1.0/24或192.168.2.0/24这两个网段的数据 ### -w `` @@ -39,9 +60,11 @@ | 大于等于8 | AES256-GCM | ### -W + 开启和服务端通信的数据加密,采用rsa+aes256gcm加密客户端和服务端之间通信的数据,可以避免token泄漏、中间人攻击 注意: + 1. -w ``是用于客户端-客户端之间的加密,password不会传递到服务端,只添加这个参数不会加密客户端-服务端通信的数据 2. -W 用于开启客户端-服务端之间的加密 @@ -49,55 +72,93 @@ 设置虚拟网卡的mtu值,大多数情况下使用默认值效率会更高,也可根据实际情况微调这个值,不加密默认为1450,加密默认为1410 -### --tcp +### --tcp + 和服务端使用tcp通信。有些网络提供商对UDP限制比较大,这个时候可以选择使用TCP模式,提高稳定性。一般来说udp延迟和消耗更低 + ### --ip `` + 指定虚拟ip,指定的ip不能和其他设备重复,必须有效并且在服务端所属网段下,默认情况由服务端分配 + ### --par `` + 任务并行度(必须为正整数),默认值为1,该值表示处理网卡读写的任务数,组网设备数较多、处理延迟较大时可适当调大此值 + ### --model `` -加密模式,可选值 aes_gcm/aes_cbc/aes_ecb/sm4_cbc,默认使用aes_gcm,通常情况aes_gcm安全性高、aes_ecb性能更好,但是在低性能设备上sm4_cbc也许速度会更快; +加密模式,可选值 +aes_gcm/aes_cbc/aes_ecb/sm4_cbc/chacha20_poly1305/chacha20/xor,默认使用aes_gcm,通常情况aes_gcm和chacha20_poly1305安全性高。 +各种加密模式的安全性和速度都不相同,请按需选取 -| 密码位数 | model | 加密算法 | -|-------|---------|------------| -| 1~8位 | aes_gcm | AES128-GCM | -| `>=`8 | aes_gcm | AES256-GCM | -| 1~8位 | aes_cbc | AES128-CBC | -| `>=`8 | aes_cbc | AES256-CBC | -| 1~8位 | aes_ecb | AES128-ECB | -| `>=`8 | aes_ecb | AES256-ECB | -| `>0` | sm4_cbc | SM4-CBC | -### --finger +特别说明:xor只是对数据进行简单异或,仅仅避免了明文传输,安全性很差,同时对性能影响也极小; + +| 密码位数 | model | 加密算法 | +|-------|-------------------|-------------------| +| 1~8位 | aes_gcm | AES128-GCM | +| `>=`8 | aes_gcm | AES256-GCM | +| 1~8位 | aes_cbc | AES128-CBC | +| `>=`8 | aes_cbc | AES256-CBC | +| 1~8位 | aes_ecb | AES128-ECB | +| `>=`8 | aes_ecb | AES256-ECB | +| `>0` | sm4_cbc | SM4-CBC | +| `>0` | chacha20_poly1305 | ChaCha20-Poly1305 | +| `>0` | chacha20 | ChaCha20 | +| `>0` | xor | 简单异或混淆 | + +### --finger 开启数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启,开启会损耗一部分性能 注意:默认情况下服务端不会对中转的数据做校验,如果要对中转的数据做校验,则需要客户端、服务端都开启此参数 + ### --punch `` + 取值ipv4/ipv6,选择只使用ipv4打洞或者只使用ipv6打洞,默认两者都会使用 + ### --ports `` + 指定本地监听的端口组,多个端口使用逗号分隔,多个端口可以分摊流量,增加并发、减缓流量限制,tcp会监听端口组的第一个端口,用于tcp直连 - 例1:‘--ports 12345,12346,12347’ 表示udp监听12345、12346、12347这三个端口,tcp监听12345端口 - 例2:‘--ports 0,0’ 表示udp监听两个未使用的端口,tcp监听一个未使用的端口 + ### --cmd + 开启交互式命令,开启后可以直接在窗口下输入命令,如需后台运行请勿开启 + ### --first_latency + 优先使用低延迟通道,默认情况下优先使用p2p通道,某些情况下可能p2p比客户端中继延迟更高,可使用此参数进行优化传输 + ### --no-proxy + 关闭内置的ip代理,内置的代理较为简单,而且一般来说直接使用网卡NAT转发性能会更高, 有需要可以自行配置NAT转发,[可参考‘编译’小节中的NAT配置](https://github.com/lbl8603/vnt#%E7%BC%96%E8%AF%91) + ### --dns `<223.5.5.5>` + 设置域名解析服务器地址,可以设置多个。如果使用TXT记录的域名,则dns默认使用223.5.5.5和114.114.114.114,端口省略值为53 当地址解析失败时,会依次尝试后面的dns,直到有A记录、AAAA记录(或TXT记录)的解析结果 ### --mapping `10.26.0.10:80>` + 端口映射,可以设置多个映射地址,例如 '--mapping udp:0.0.0.0:80->10.26.0.10:80 --mapping tcp:0.0.0.0:80->10.26.0.11:81' 表示将本地udp 80端口的数据转发到10.26.0.10:80,将本地tcp 80端口的数据转发到10.26.0.11:81,转发的目的地址可以使用域名+端口 + +### --compressor `` + +启用压缩,默认仅支持lz4压缩,开启压缩后,如果数据包长度大于等于128,则会使用压缩,否则还是会按原数据发送 + +也支持开启zstd压缩,但是需要自行编译,编译时加入参数--features zstd + +如果宽度速度比较慢,可以考虑使用高级别的压缩 + ### -f `` + 指定配置文件 配置文件采用yaml格式,可参考: + ```yaml # 全部参数 tap: false #是否使用tap 仅在windows上支持使用tap @@ -105,7 +166,7 @@ token: xxx #组网token device_id: xxx #当前设备id name: windows 11 #当前设备名称 server_address: ip:port #注册和中继服务器 -stun_server: #stun服务器 +stun_server: #stun服务器 - stun1.l.google.com:19302 - stun2.l.google.com:19302 in_ips: #代理ip入站 @@ -122,7 +183,7 @@ parallel: 1 #任务并行度 cipher_model: aes_gcm #客户端加密算法 finger: false #关闭数据指纹 punch_model: ipv4 #打洞模式,表示只使用ipv4地址打洞,默认会同时使用v6和v4 -ports: +ports: - 0 #使用随机端口,tcp监听此端口 - 0 cmd: false #关闭控制台输入 @@ -141,25 +202,41 @@ mapping: ``` 或者需要哪个配置就加哪个,当然token是必须的 + ```yaml # 部分参数 token: xxx #组网token ``` + ### --use-channel `` + - relay:仅中继模式,会禁止打洞/p2p直连,只使用服务器转发 - p2p:仅直连模式,会禁止网络数据从服务器/客户端转发,只会使用服务器转发控制包 + ### --packet-loss `<0>` + 模拟丢包,取值0~1之间的小数,程序会按设定的概率主动丢包。在模拟弱网环境时会有帮助。 + ### --packet-delay `<0>` + 模拟延迟,整数,单位毫秒(ms),程序会按设定的值延迟发包,可用于模拟弱网 ### --list + 在后台运行时,查看其他设备列表 + ### --all + 在后台运行时,查看其他设备完整信息 + ### --info + 在后台运行时,查看当前设备信息 -### --route + +### --route + 在后台运行时,查看数据转发路径 + ### --stop + 停止后台运行 diff --git a/vnt-cli/src/config/file_config.rs b/vnt-cli/src/config/file_config.rs new file mode 100644 index 0000000..0fbe436 --- /dev/null +++ b/vnt-cli/src/config/file_config.rs @@ -0,0 +1,171 @@ +use anyhow::anyhow; +use std::net::Ipv4Addr; +use std::str::FromStr; + +use serde::{Deserialize, Serialize}; + +use crate::config::get_device_id; +use vnt::channel::punch::PunchModel; +use vnt::channel::UseChannelType; +use vnt::cipher::CipherModel; +use vnt::compression::Compressor; +use vnt::core::Config; + +#[derive(Serialize, Deserialize, Debug)] +#[serde(default)] +pub struct FileConfig { + #[cfg(target_os = "windows")] + pub tap: bool, + pub token: String, + pub device_id: String, + pub name: String, + pub server_address: String, + pub stun_server: Vec, + pub dns: Vec, + pub in_ips: Vec, + pub out_ips: Vec, + pub password: Option, + pub mtu: Option, + pub tcp: bool, + pub ip: Option, + pub use_channel: String, + #[cfg(feature = "ip_proxy")] + pub no_proxy: bool, + pub server_encrypt: bool, + pub parallel: usize, + pub cipher_model: Option, + pub finger: bool, + pub punch_model: String, + pub ports: Option>, + pub cmd: bool, + pub first_latency: bool, + pub device_name: Option, + pub packet_loss: Option, + pub packet_delay: u32, + #[cfg(feature = "port_mapping")] + pub mapping: Vec, + pub compressor: Option, +} + +impl Default for FileConfig { + fn default() -> Self { + Self { + #[cfg(target_os = "windows")] + tap: false, + token: "".to_string(), + device_id: get_device_id(), + name: os_info::get().to_string(), + server_address: "nat1.wherewego.top:29872".to_string(), + stun_server: vec![ + "stun1.l.google.com:19302".to_string(), + "stun2.l.google.com:19302".to_string(), + "stun.miwifi.com:3478".to_string(), + ], + dns: vec![], + in_ips: vec![], + out_ips: vec![], + password: None, + mtu: None, + tcp: false, + ip: None, + use_channel: "all".to_string(), + #[cfg(feature = "ip_proxy")] + no_proxy: false, + server_encrypt: false, + parallel: 1, + cipher_model: None, + finger: false, + punch_model: "all".to_string(), + ports: None, + cmd: false, + first_latency: false, + device_name: None, + packet_loss: None, + packet_delay: 0, + #[cfg(feature = "port_mapping")] + mapping: vec![], + compressor: None, + } + } +} + +pub fn read_config(file_path: &str) -> anyhow::Result<(Config, bool)> { + let conf = std::fs::read_to_string(file_path)?; + let file_conf = match serde_yaml::from_str::(&conf) { + Ok(val) => val, + Err(e) => { + log::error!("{:?}", e); + return Err(anyhow!("{}", e)); + } + }; + if file_conf.token.is_empty() { + return Err(anyhow!("token is_empty")); + } + + let in_ips = match common::args_parse::ips_parse(&file_conf.in_ips) { + Ok(in_ips) => in_ips, + Err(e) => { + return Err(anyhow!("in_ips {:?} error:{}", &file_conf.in_ips, e)); + } + }; + let out_ips = match common::args_parse::out_ips_parse(&file_conf.out_ips) { + Ok(out_ips) => out_ips, + Err(e) => { + return Err(anyhow!("out_ips {:?} error:{}", &file_conf.out_ips, e)); + } + }; + let virtual_ip = match file_conf.ip.clone().map(|v| Ipv4Addr::from_str(&v)) { + None => None, + Some(r) => Some(r.map_err(|e| anyhow!("ip {:?} error:{}", &file_conf.ip, e))?), + }; + let cipher_model = { + #[cfg(not(any(feature = "aes_gcm", feature = "server_encrypt")))] + if file_conf.password.is_some() && file_conf.cipher_model.is_none() { + Err(anyhow!("cipher_model undefined"))? + } + #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] + CipherModel::AesGcm + }; + + let punch_model = PunchModel::from_str(&file_conf.punch_model).map_err(|e| anyhow!("{}", e))?; + let use_channel_type = + UseChannelType::from_str(&file_conf.use_channel).map_err(|e| anyhow!("{}", e))?; + let compressor = if let Some(compressor) = file_conf.compressor.as_ref() { + Compressor::from_str(compressor).map_err(|e| anyhow!("{}", e))? + } else { + Compressor::None + }; + let config = Config::new( + #[cfg(target_os = "windows")] + file_conf.tap, + file_conf.token, + file_conf.device_id, + file_conf.name, + file_conf.server_address, + file_conf.dns, + file_conf.stun_server, + in_ips, + out_ips, + file_conf.password, + file_conf.mtu, + file_conf.tcp, + virtual_ip, + #[cfg(feature = "ip_proxy")] + file_conf.no_proxy, + file_conf.server_encrypt, + file_conf.parallel, + cipher_model, + file_conf.finger, + punch_model, + file_conf.ports, + file_conf.first_latency, + file_conf.device_name, + use_channel_type, + file_conf.packet_loss, + file_conf.packet_delay, + #[cfg(feature = "port_mapping")] + file_conf.mapping, + compressor, + )?; + Ok((config, file_conf.cmd)) +} diff --git a/vnt-cli/src/main.rs b/vnt-cli/src/main.rs index 9ad1521..fbdd09c 100644 --- a/vnt-cli/src/main.rs +++ b/vnt-cli/src/main.rs @@ -1,3 +1,4 @@ +use anyhow::anyhow; use std::io; use std::net::Ipv4Addr; use std::path::PathBuf; @@ -10,6 +11,7 @@ use common::args_parse::{ips_parse, out_ips_parse}; use vnt::channel::punch::PunchModel; use vnt::channel::UseChannelType; use vnt::cipher::CipherModel; +use vnt::compression::Compressor; use vnt::core::{Config, Vnt}; #[cfg(feature = "command")] @@ -78,6 +80,7 @@ fn main() { opts.optmulti("", "dns", "dns", ""); opts.optmulti("", "mapping", "mapping", ""); opts.optopt("f", "", "配置文件", ""); + opts.optopt("", "compressor", "压缩算法", ""); //"后台运行时,查看其他设备列表" opts.optflag("", "list", "后台运行时,查看其他设备列表"); opts.optflag("", "all", "后台运行时,查看其他设备完整信息"); @@ -230,19 +233,6 @@ fn main() { let cipher_model = match matches.opt_get::("model") { Ok(model) => { - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] - { - if password.is_some() && model.is_none() { - println!("Encryption not supported"); - return; - } - } #[cfg(not(any(feature = "aes_gcm", feature = "server_encrypt")))] { if password.is_some() && model.is_none() { @@ -294,6 +284,13 @@ fn main() { .unwrap_or(0); #[cfg(feature = "port_mapping")] let port_mapping_list = matches.opt_strs("mapping"); + let compressor = if let Some(compressor) = matches.opt_str("compressor").as_ref() { + Compressor::from_str(compressor) + .map_err(|e| anyhow!("{}", e)) + .unwrap() + } else { + Compressor::None + }; let config = match Config::new( #[cfg(target_os = "windows")] tap, @@ -324,10 +321,11 @@ fn main() { packet_delay, #[cfg(feature = "port_mapping")] port_mapping_list, + compressor, ) { Ok(config) => config, Err(e) => { - println!("config error: {}", e); + println!("config.toml error: {}", e); return; } }; @@ -356,6 +354,29 @@ fn main0(config: Config, _show_cmd: bool) { } } let vnt_util = Vnt::new(config, callback::VntHandler {}).unwrap(); + #[cfg(any(target_os = "linux", target_os = "macos"))] + { + let vnt_c = vnt_util.clone(); + let mut signals = signal_hook::iterator::Signals::new(&[ + signal_hook::consts::SIGINT, + signal_hook::consts::SIGTERM, + ]) + .unwrap(); + let handle = signals.handle(); + std::thread::spawn(move || { + for sig in signals.forever() { + match sig { + signal_hook::consts::SIGINT | signal_hook::consts::SIGTERM => { + println!("Received SIGINT, {}", sig); + vnt_c.stop(); + handle.close(); + break; + } + _ => {} + } + } + }); + } #[cfg(feature = "command")] { let vnt_c = vnt_util.clone(); @@ -389,6 +410,7 @@ fn main0(config: Config, _show_cmd: bool) { vnt_util.wait() } + #[cfg(feature = "command")] fn command(cmd: &str, vnt: &Vnt) -> bool { if cmd.is_empty() { @@ -441,33 +463,8 @@ fn print_usage(program: &str, _opts: Options) { println!(" -i 配置点对网(IP代理)时使用,-i 192.168.0.0/24,10.26.0.3表示允许接收网段192.168.0.0/24的数据"); println!(" 并转发到10.26.0.3,可指定多个网段"); println!(" -o 配置点对网时使用,-o 192.168.0.0/24表示允许将数据转发到192.168.0.0/24,可指定多个网段"); - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] - let enums = String::new(); - #[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - ))] - let mut enums = String::new(); - #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] - enums.push_str("/aes_gcm"); - #[cfg(feature = "aes_cbc")] - enums.push_str("/aes_cbc"); - #[cfg(feature = "aes_ecb")] - enums.push_str("/aes_ecb"); - #[cfg(feature = "sm4_cbc")] - enums.push_str("/sm4_cbc"); - if !enums.is_empty() { - println!(" -w 使用该密码生成的密钥对客户端数据进行加密,并且服务端无法解密,使用相同密码的客户端才能通信"); - } + + println!(" -w 使用该密码生成的密钥对客户端数据进行加密,并且服务端无法解密,使用相同密码的客户端才能通信"); #[cfg(feature = "server_encrypt")] println!(" -W 加密当前客户端和服务端通信的数据,请留意服务端指纹是否正确"); println!(" -u 自定义mtu(不加密默认为1450,加密默认为1410)"); @@ -477,15 +474,31 @@ fn print_usage(program: &str, _opts: Options) { println!(" --tcp 和服务端使用tcp通信,默认使用udp,遇到udp qos时可指定使用tcp"); println!(" --ip 指定虚拟ip,指定的ip不能和其他设备重复,必须有效并且在服务端所属网段下,默认情况由服务端分配"); println!(" --par 任务并行度(必须为正整数),默认值为1"); - if !enums.is_empty() { - println!( - " --model 加密模式(默认aes_gcm),可选值{}", - &enums[1..] - ); - } - if !enums.is_empty() { - println!(" --finger 增加数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启"); - } + let mut enums = String::new(); + #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] + enums.push_str("/aes_gcm"); + #[cfg(feature = "chacha20_poly1305")] + enums.push_str("/chacha20_poly1305/chacha20"); + #[cfg(feature = "aes_cbc")] + enums.push_str("/aes_cbc"); + #[cfg(feature = "aes_ecb")] + enums.push_str("/aes_ecb"); + #[cfg(feature = "sm4_cbc")] + enums.push_str("/sm4_cbc"); + enums.push_str("/xor"); + println!( + " --model 加密模式(默认aes_gcm),可选值{}", + &enums[1..] + ); + #[cfg(any( + feature = "aes_gcm", + feature = "chacha20_poly1305", + feature = "server_encrypt", + feature = "aes_cbc", + feature = "aes_ecb", + feature = "sm4_cbc" + ))] + println!(" --finger 增加数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启"); println!(" --punch 取值ipv4/ipv6/all,ipv4表示仅使用ipv4打洞"); println!(" --ports 取值0~65535,指定本地监听的一组端口,默认监听两个随机端口,使用过多端口会增加网络负担"); #[cfg(feature = "command")] @@ -502,7 +515,14 @@ fn print_usage(program: &str, _opts: Options) { println!(" --dns DNS服务器地址,可使用多个dns,不指定时使用系统解析"); #[cfg(feature = "port_mapping")] println!(" --mapping 端口映射,例如 --mapping udp:0.0.0.0:80->10.26.0.10:80 --mapping tcp:0.0.0.0:80->10.26.0.10:80"); - + #[cfg(all(feature = "lz4", feature = "zstd"))] + println!(" --compressor 启用压缩,可选值lz4/zstd<,level>,level为压缩级别,例如 --compressor lz4 或--compressor zstd,10"); + #[cfg(feature = "lz4")] + #[cfg(not(feature = "zstd"))] + println!(" --compressor 启用压缩,可选值lz4,例如 --compressor lz4"); + #[cfg(feature = "zstd")] + #[cfg(not(feature = "lz4"))] + println!(" --compressor 启用压缩,可选值zstd<,level>,level为压缩级别,例如 --compressor zstd,10"); println!(); #[cfg(feature = "command")] { diff --git a/vnt-jni/Cargo.toml b/vnt-jni/Cargo.toml index 339c3f5..7ac63b5 100644 --- a/vnt-jni/Cargo.toml +++ b/vnt-jni/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "vnt-jni" -version = "1.2.9" +version = "1.2.10" edition = "2021" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html diff --git a/vnt-jni/src/config.rs b/vnt-jni/src/config.rs index 2a94cab..20ccf4b 100644 --- a/vnt-jni/src/config.rs +++ b/vnt-jni/src/config.rs @@ -7,6 +7,7 @@ use jni::JNIEnv; use vnt::channel::punch::PunchModel; use vnt::channel::UseChannelType; use vnt::cipher::CipherModel; +use vnt::compression::Compressor; use vnt::core::Config; use crate::utils::*; @@ -118,6 +119,7 @@ pub fn new_config(env: &mut JNIEnv, config: JObject) -> Result { packet_loss_rate, packet_delay, port_mapping, + Compressor::None, ) { Ok(config) => config, Err(e) => { diff --git a/vnt/Cargo.toml b/vnt/Cargo.toml index a58f0a4..3a1ee64 100644 --- a/vnt/Cargo.toml +++ b/vnt/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "vnt" -version = "1.2.9" +version = "1.2.10" edition = "2021" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html @@ -23,6 +23,8 @@ aes-gcm = { version = "0.10.2", optional = true } ring = { version = "0.17.0", optional = true } cbc = { version = "0.1.2", optional = true } ecb = { version = "0.1.2", optional = true } +chacha20poly1305 = { version = "0.10.1", optional = true } +chacha20 = { version = "0.9.1", optional = true } aes = "0.8.3" stun-format = { version = "1.0.1", features = ["fmt", "rfc3489"] } rsa = { version = "0.9.2", features = [], optional = true } @@ -37,6 +39,8 @@ dns-parser = "0.8.0" tokio = { version = "1.37.0", features = ["full"], optional = true } +lz4_flex = { version = "0.11", default-features = false, optional = true } +zstd = { version = "0.13.1", optional = true } [target.'cfg(target_os = "windows")'.dependencies] libloading = "0.8.0" @@ -45,9 +49,10 @@ libloading = "0.8.0" [build-dependencies] protobuf-codegen = "3.2.0" protoc-bin-vendored = "3.0.0" +cfg_aliases = "0.2.1" [features] -default = ["server_encrypt", "aes_gcm", "aes_cbc", "aes_ecb", "sm4_cbc", "ip_proxy", "port_mapping"] +default = ["server_encrypt", "aes_gcm", "aes_cbc", "aes_ecb", "sm4_cbc", "chacha20_poly1305", "ip_proxy", "port_mapping", "lz4_compress", "zstd_compress"] openssl = ["openssl-sys"] # 从源码编译 openssl-vendored = ["openssl-sys/vendored"] @@ -56,6 +61,9 @@ aes_cbc = ["cbc"] aes_ecb = ["ecb"] sm4_cbc = ["libsm"] aes_gcm = ["aes-gcm"] +chacha20_poly1305 = ["chacha20poly1305", "chacha20"] server_encrypt = ["aes-gcm", "rsa", "spki"] ip_proxy = ["tokio"] port_mapping = ["tokio"] +lz4_compress = ["lz4_flex"] +zstd_compress = ["zstd"] diff --git a/vnt/build.rs b/vnt/build.rs index 8e9f1ab..5dd2e36 100644 --- a/vnt/build.rs +++ b/vnt/build.rs @@ -1,4 +1,17 @@ +use cfg_aliases::cfg_aliases; + fn main() { + cfg_aliases! { + cipher: { + any(feature = "aes_gcm", + feature = "chacha20_poly1305", + feature = "server_encrypt", + feature = "aes_cbc", + feature = "aes_ecb", + feature = "sm4_cbc" + )}, + } + std::fs::create_dir_all("src/proto").unwrap(); protobuf_codegen::Codegen::new() .pure() diff --git a/vnt/src/channel/context.rs b/vnt/src/channel/context.rs index 200747a..47de7da 100644 --- a/vnt/src/channel/context.rs +++ b/vnt/src/channel/context.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use std::net::{Ipv4Addr, SocketAddr, SocketAddrV6, UdpSocket}; use std::ops::Deref; -use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; use std::time::{Duration, Instant}; use std::{io, thread}; @@ -48,7 +48,6 @@ impl ChannelContext { tcp_map: RwLock::new(HashMap::with_capacity(64)), route_table: RouteTable::new(use_channel_type, first_latency, channel_num), is_tcp, - state: AtomicBool::new(true), packet_loss_rate, packet_delay, main_index: AtomicUsize::new(0), @@ -86,8 +85,6 @@ pub struct ContextInner { pub route_table: RouteTable, // 是否使用tcp连接服务器 is_tcp: bool, - //状态 - state: AtomicBool, //控制丢包率,取值v=[0,100_0000] 丢包率r=v/100_0000 packet_loss_rate: u32, //控制延迟 @@ -100,12 +97,6 @@ impl ContextInner { pub fn use_channel_type(&self) -> UseChannelType { self.route_table.use_channel_type } - pub fn is_stop(&self) -> bool { - !self.state.load(Ordering::Acquire) - } - pub fn stop(&self) { - self.state.store(false, Ordering::Release); - } /// 通过sub_udp_socket是否为空来判断是否为锥形网络 pub fn is_cone(&self) -> bool { self.sub_udp_socket.read().is_empty() diff --git a/vnt/src/channel/handler.rs b/vnt/src/channel/handler.rs index d9e8c0e..4715fb9 100644 --- a/vnt/src/channel/handler.rs +++ b/vnt/src/channel/handler.rs @@ -2,5 +2,11 @@ use crate::channel::context::ChannelContext; use crate::channel::RouteKey; pub trait RecvChannelHandler: Clone + Send + 'static { - fn handle(&mut self, buf: &mut [u8], route_key: RouteKey, context: &ChannelContext); + fn handle( + &mut self, + buf: &mut [u8], + extend: &mut [u8], + route_key: RouteKey, + context: &ChannelContext, + ); } diff --git a/vnt/src/channel/mod.rs b/vnt/src/channel/mod.rs index 5373f48..7536edb 100644 --- a/vnt/src/channel/mod.rs +++ b/vnt/src/channel/mod.rs @@ -18,7 +18,7 @@ pub mod sender; pub mod tcp_channel; pub mod udp_channel; -const BUFFER_SIZE: usize = 1024 * 16; +pub const BUFFER_SIZE: usize = 1024 * 16; #[derive(Debug, Copy, Clone, Eq, PartialEq)] pub enum UseChannelType { Relay, diff --git a/vnt/src/channel/tcp_channel.rs b/vnt/src/channel/tcp_channel.rs index 5628937..c5a0219 100644 --- a/vnt/src/channel/tcp_channel.rs +++ b/vnt/src/channel/tcp_channel.rs @@ -31,7 +31,7 @@ pub fn tcp_listen( stop_manager: StopManager, recv_handler: H, context: ChannelContext, -) -> io::Result>)>> +) -> anyhow::Result>)>> where H: RecvChannelHandler, { @@ -75,7 +75,7 @@ fn tcp_listen0( accept_tcp_receiver: Receiver<(TcpStream, SocketAddr, Option>)>, mut recv_handler: H, context: ChannelContext, -) -> io::Result<()> +) -> anyhow::Result<()> where H: RecvChannelHandler, { @@ -87,6 +87,7 @@ where let mut read_map: HashMap, usize)> = HashMap::with_capacity(32); + let mut extend = [0; BUFFER_SIZE]; loop { poll.poll(&mut events, None)?; for event in events.iter() { @@ -108,7 +109,7 @@ where if e.kind() == io::ErrorKind::WouldBlock { break; } - return Err(e); + return Err(e)?; } } }, @@ -132,9 +133,13 @@ where } token => { if event.is_readable() { - if let Err(e) = - readable_handle(&token, &mut read_map, &mut recv_handler, &context) - { + if let Err(e) = readable_handle( + &token, + &mut read_map, + &mut recv_handler, + &context, + &mut extend, + ) { closed_handle_r(&token, &mut read_map); log::warn!("{:?}", e); if let Err(e) = write_waker.notify(token, false) { @@ -159,7 +164,7 @@ fn init_writable_handler( receiver: Receiver<(TcpStream, Token, SocketAddr, Option>)>, stop_manager: StopManager, context: ChannelContext, -) -> io::Result { +) -> anyhow::Result { let poll = Poll::new()?; let writable_notify = WritableNotify::new(Waker::new(poll.registry(), NOTIFY)?); let worker = { @@ -339,6 +344,7 @@ fn readable_handle( map: &mut HashMap, usize)>, recv_handler: &mut H, context: &ChannelContext, + extend: &mut [u8], ) -> io::Result<()> where H: RecvChannelHandler, @@ -360,7 +366,7 @@ where } *begin += len; if end > 4 && *begin == end { - recv_handler.handle(&mut buf[4..end], *route_key, context); + recv_handler.handle(&mut buf[4..end], extend, *route_key, context); *begin = 0; } } diff --git a/vnt/src/channel/udp_channel.rs b/vnt/src/channel/udp_channel.rs index 3f3b1ec..4a0a26a 100644 --- a/vnt/src/channel/udp_channel.rs +++ b/vnt/src/channel/udp_channel.rs @@ -17,7 +17,7 @@ pub fn udp_listen( stop_manager: StopManager, recv_handler: H, context: ChannelContext, -) -> io::Result>>> +) -> anyhow::Result>>> where H: RecvChannelHandler, { @@ -31,7 +31,7 @@ fn sub_udp_listen( stop_manager: StopManager, recv_handler: H, context: ChannelContext, -) -> io::Result>>> +) -> anyhow::Result>>> where H: RecvChannelHandler, { @@ -70,6 +70,7 @@ where { let mut events = Events::with_capacity(1024); let mut buf = [0; BUFFER_SIZE]; + let mut extend = [0; BUFFER_SIZE]; let mut read_map: HashMap = HashMap::with_capacity(32); loop { poll.poll(&mut events, None)?; @@ -115,6 +116,7 @@ where Ok((len, addr)) => { recv_handler.handle( &mut buf[..len], + &mut extend, RouteKey::new(false, token.0, addr), &context, ); @@ -206,7 +208,7 @@ fn main_udp_listen( stop_manager: StopManager, recv_handler: H, context: ChannelContext, -) -> io::Result<()> +) -> anyhow::Result<()> where H: RecvChannelHandler, { @@ -252,6 +254,7 @@ where } let mut events = Events::with_capacity(udps.len()); + let mut extend = [0; BUFFER_SIZE]; loop { poll.poll(&mut events, None)?; for x in events.iter() { @@ -270,6 +273,7 @@ where Ok((len, addr)) => { recv_handler.handle( &mut buf[..len], + &mut extend, RouteKey::new(false, index, addr), &context, ); diff --git a/vnt/src/cipher/aes_cbc/mod.rs b/vnt/src/cipher/aes_cbc/mod.rs new file mode 100644 index 0000000..6286a84 --- /dev/null +++ b/vnt/src/cipher/aes_cbc/mod.rs @@ -0,0 +1,2 @@ +mod rs_aes_cbc; +pub use rs_aes_cbc::*; diff --git a/vnt/src/cipher/aes_cbc.rs b/vnt/src/cipher/aes_cbc/rs_aes_cbc.rs similarity index 89% rename from vnt/src/cipher/aes_cbc.rs rename to vnt/src/cipher/aes_cbc/rs_aes_cbc.rs index 708fd55..b16d33a 100644 --- a/vnt/src/cipher/aes_cbc.rs +++ b/vnt/src/cipher/aes_cbc/rs_aes_cbc.rs @@ -1,6 +1,5 @@ -use std::io; - use aes::cipher::{block_padding::Pkcs7, BlockDecryptMut, BlockEncryptMut, KeyIvInit}; +use anyhow::anyhow; use rand::RngCore; use crate::cipher::Finger; @@ -50,14 +49,14 @@ impl AesCbcCipher { pub fn decrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if !net_packet.is_encrypt() { //未加密的数据直接丢弃 - return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); + return Err(anyhow!("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")); + return Err(anyhow!("aes_cbc data err")); } let mut iv = [0; 16]; iv[0..4].copy_from_slice(&net_packet.source().octets()); @@ -75,7 +74,7 @@ impl AesCbcCipher { 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")); + return Err(anyhow!("aes_cbc finger err")); } } let rs = match &self.cipher { @@ -92,10 +91,7 @@ impl AesCbcCipher { net_packet.set_data_len(HEAD_LEN + len - 4)?; Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("解密失败:{}", e), - )), + Err(e) => Err(anyhow!("aes_cbc 解密失败:{}", e)), } } /// net_packet 必须预留足够长度 @@ -103,7 +99,7 @@ impl AesCbcCipher { pub fn encrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let data_len = net_packet.data_len(); let mut iv = [0; 16]; iv[0..4].copy_from_slice(&net_packet.source().octets()); @@ -146,10 +142,7 @@ impl AesCbcCipher { net_packet.set_encrypt_flag(true); Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("加密失败:{}", e), - )), + Err(e) => Err(anyhow!("aes_cbc 加密失败:{}", e)), }; } } diff --git a/vnt/src/cipher/aes_ecb/mod.rs b/vnt/src/cipher/aes_ecb/mod.rs new file mode 100644 index 0000000..eab811f --- /dev/null +++ b/vnt/src/cipher/aes_ecb/mod.rs @@ -0,0 +1,9 @@ +#[cfg(not(any(feature = "openssl-vendored", feature = "openssl")))] +mod rs_aes_ecb; +#[cfg(not(any(feature = "openssl-vendored", feature = "openssl")))] +pub use rs_aes_ecb::*; + +#[cfg(any(feature = "openssl-vendored", feature = "openssl"))] +mod openssl_aes_ecb; +#[cfg(any(feature = "openssl-vendored", feature = "openssl"))] +pub use openssl_aes_ecb::*; diff --git a/vnt/src/cipher/openssl_aes_ecb.rs b/vnt/src/cipher/aes_ecb/openssl_aes_ecb.rs similarity index 88% rename from vnt/src/cipher/openssl_aes_ecb.rs rename to vnt/src/cipher/aes_ecb/openssl_aes_ecb.rs index e0e5b23..649fa04 100644 --- a/vnt/src/cipher/openssl_aes_ecb.rs +++ b/vnt/src/cipher/aes_ecb/openssl_aes_ecb.rs @@ -1,8 +1,11 @@ -use crate::cipher::Finger; -use crate::protocol::{NetPacket, HEAD_LEN}; +use std::ptr; + +use anyhow::anyhow; use libc::c_int; use openssl_sys::EVP_CIPHER_CTX; -use std::{io, ptr}; + +use crate::cipher::Finger; +use crate::protocol::{NetPacket, HEAD_LEN}; pub struct AesEcbCipher { key: Vec, @@ -100,10 +103,10 @@ impl AesEcbCipher { pub fn decrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if !net_packet.is_encrypt() { //未加密的数据直接丢弃 - return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); + return Err(anyhow!("not encrypt")); } if let Some(finger) = &self.finger { @@ -116,18 +119,18 @@ impl AesEcbCipher { nonce_raw[11] = net_packet.source_ttl(); let len = net_packet.payload().len(); if len < 12 { - return Err(io::Error::new(io::ErrorKind::Other, "data len err")); + return Err(anyhow!("data len err")); } let secret_body = &net_packet.payload()[..len - 12]; let finger = finger.calculate_finger(&nonce_raw, secret_body); if &finger != &net_packet.payload()[len - 12..] { - return Err(io::Error::new(io::ErrorKind::Other, "finger err")); + return Err(anyhow!("finger err")); } net_packet.set_data_len(net_packet.data_len() - finger.len())?; } if net_packet.payload().len() < 16 { log::error!("数据异常,长度{}小于{}", net_packet.payload().len(), 16); - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } let input = net_packet.payload(); let mut out = [0u8; 1024 * 5]; @@ -147,22 +150,22 @@ impl AesEcbCipher { //校验头部 let src_net_packet = NetPacket::new(text)?; if src_net_packet.source() != net_packet.source() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.destination() != net_packet.destination() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.protocol() != net_packet.protocol() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.transport_protocol() != net_packet.transport_protocol() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.is_gateway() != net_packet.is_gateway() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.source_ttl() != net_packet.source_ttl() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } } net_packet.set_encrypt_flag(false); @@ -175,7 +178,7 @@ impl AesEcbCipher { pub fn encrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let input = net_packet.buffer(); let mut out = [0u8; 1024 * 5]; let mut out_len = 0; @@ -191,7 +194,7 @@ impl AesEcbCipher { } let out_len = out_len as usize; if out_len == 0 { - return Err(io::Error::new(io::ErrorKind::Other, "ciphertext len err")); + return Err(anyhow!("ciphertext len err")); } //密文 let ciphertext = &out[..out_len]; diff --git a/vnt/src/cipher/aes_ecb.rs b/vnt/src/cipher/aes_ecb/rs_aes_ecb.rs similarity index 82% rename from vnt/src/cipher/aes_ecb.rs rename to vnt/src/cipher/aes_ecb/rs_aes_ecb.rs index 364c605..5fd9485 100644 --- a/vnt/src/cipher/aes_ecb.rs +++ b/vnt/src/cipher/aes_ecb/rs_aes_ecb.rs @@ -1,7 +1,8 @@ +use aes::cipher::{block_padding::Pkcs7, BlockDecryptMut, BlockEncryptMut, KeyInit}; +use anyhow::anyhow; + use crate::cipher::Finger; use crate::protocol::{NetPacket, HEAD_LEN}; -use aes::cipher::{block_padding::Pkcs7, BlockDecryptMut, BlockEncryptMut, KeyInit}; -use std::io; type Aes128EcbEnc = ecb::Encryptor; type Aes128EcbDec = ecb::Decryptor; @@ -46,10 +47,10 @@ impl AesEcbCipher { pub fn decrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if !net_packet.is_encrypt() { //未加密的数据直接丢弃 - return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); + return Err(anyhow!("not encrypt")); } if let Some(finger) = &self.finger { @@ -62,18 +63,18 @@ impl AesEcbCipher { nonce_raw[11] = net_packet.source_ttl(); let len = net_packet.payload().len(); if len < 12 { - return Err(io::Error::new(io::ErrorKind::Other, "payload len <12")); + return Err(anyhow!("payload len <12")); } let secret_body = &net_packet.payload()[..len - 12]; let finger = finger.calculate_finger(&nonce_raw, secret_body); if &finger != &net_packet.payload()[len - 12..] { - return Err(io::Error::new(io::ErrorKind::Other, "finger err")); + return Err(anyhow!("finger err")); } net_packet.set_data_len(net_packet.data_len() - finger.len())?; } if net_packet.payload().len() < 16 { log::error!("数据异常,长度{}小于{}", net_packet.payload().len(), 16); - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } let mut out = [0u8; 1024 * 5]; let rs = match self.key { @@ -87,32 +88,29 @@ impl AesEcbCipher { //校验头部 let src_net_packet = NetPacket::new(buf)?; if src_net_packet.source() != net_packet.source() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.destination() != net_packet.destination() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.protocol() != net_packet.protocol() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.transport_protocol() != net_packet.transport_protocol() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.is_gateway() != net_packet.is_gateway() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.source_ttl() != net_packet.source_ttl() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } net_packet.set_data_len(buf.len())?; net_packet.set_payload(src_net_packet.payload())?; net_packet.set_encrypt_flag(false); Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("aes_ecb解密失败:{}", e), - )), + Err(e) => Err(anyhow!("aes_ecb解密失败:{}", e)), } } /// net_packet 必须预留足够长度 @@ -120,7 +118,7 @@ impl AesEcbCipher { pub fn encrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let mut out = [0u8; 1024 * 5]; let rs = match self.key { AesEcbEnum::AES128ECB(key) => Aes128EcbEnc::new(&key.into()) @@ -152,10 +150,7 @@ impl AesEcbCipher { } Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("aes_ecb加密失败:{}", e), - )), + Err(e) => Err(anyhow!("aes_ecb加密失败:{}", e)), }; } } diff --git a/vnt/src/cipher/aes_gcm_cipher.rs b/vnt/src/cipher/aes_gcm/aes_gcm_cipher.rs similarity index 87% rename from vnt/src/cipher/aes_gcm_cipher.rs rename to vnt/src/cipher/aes_gcm/aes_gcm_cipher.rs index 9920c10..e261a49 100644 --- a/vnt/src/cipher/aes_gcm_cipher.rs +++ b/vnt/src/cipher/aes_gcm/aes_gcm_cipher.rs @@ -1,8 +1,7 @@ -use std::io; - use aes_gcm::aead::consts::{U12, U16}; use aes_gcm::aead::generic_array::GenericArray; use aes_gcm::{AeadInPlace, Aes128Gcm, Aes256Gcm, Key, KeyInit, Nonce, Tag}; +use anyhow::anyhow; use rand::RngCore; use crate::cipher::finger::Finger; @@ -39,14 +38,14 @@ impl AesGcmCipher { pub fn decrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if !net_packet.is_encrypt() { //未加密的数据直接丢弃 - return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); + return Err(anyhow!("not encrypt")); } if net_packet.payload().len() < AES_GCM_ENCRYPTION_RESERVED { log::error!("数据异常,长度小于{}", AES_GCM_ENCRYPTION_RESERVED); - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } let mut nonce_raw = [0; 12]; nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); @@ -62,7 +61,7 @@ impl AesGcmCipher { if let Some(finger) = &self.finger { let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body()); if &finger != secret_body.finger() { - return Err(io::Error::new(io::ErrorKind::Other, "finger err")); + return Err(anyhow!("finger err")); } } let tag: GenericArray = Tag::clone_from_slice(tag); @@ -75,10 +74,7 @@ impl AesGcmCipher { } }; if let Err(e) = rs { - return Err(io::Error::new( - io::ErrorKind::Other, - format!("解密失败:{}", e), - )); + return Err(anyhow!("解密失败:{}", e)); } net_packet.set_encrypt_flag(false); net_packet.set_data_len(net_packet.data_len() - AES_GCM_ENCRYPTION_RESERVED)?; @@ -89,9 +85,9 @@ impl AesGcmCipher { pub fn encrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if net_packet.reserve() < AES_GCM_ENCRYPTION_RESERVED { - return Err(io::Error::new(io::ErrorKind::Other, "too short")); + return Err(anyhow!("too short")); } let mut nonce_raw = [0; 12]; nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); @@ -123,10 +119,7 @@ impl AesGcmCipher { net_packet.set_encrypt_flag(true); Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("加密失败:{}", e), - )), + Err(e) => Err(anyhow!("加密失败:{}", e)), }; } } diff --git a/vnt/src/cipher/aes_gcm/mod.rs b/vnt/src/cipher/aes_gcm/mod.rs new file mode 100644 index 0000000..ade96e4 --- /dev/null +++ b/vnt/src/cipher/aes_gcm/mod.rs @@ -0,0 +1,9 @@ +#[cfg(feature = "ring-cipher")] +mod ring_aes_gcm_cipher; +#[cfg(feature = "ring-cipher")] +pub use ring_aes_gcm_cipher::*; + +#[cfg(not(feature = "ring-cipher"))] +mod aes_gcm_cipher; +#[cfg(not(feature = "ring-cipher"))] +pub use aes_gcm_cipher::*; diff --git a/vnt/src/cipher/ring_aes_gcm_cipher.rs b/vnt/src/cipher/aes_gcm/ring_aes_gcm_cipher.rs similarity index 86% rename from vnt/src/cipher/ring_aes_gcm_cipher.rs rename to vnt/src/cipher/aes_gcm/ring_aes_gcm_cipher.rs index bc65a18..4123d83 100644 --- a/vnt/src/cipher/ring_aes_gcm_cipher.rs +++ b/vnt/src/cipher/aes_gcm/ring_aes_gcm_cipher.rs @@ -1,9 +1,9 @@ -use crate::cipher::Finger; +use anyhow::anyhow; use rand::RngCore; use ring::aead; use ring::aead::{LessSafeKey, UnboundKey}; -use std::io; +use crate::cipher::Finger; use crate::protocol::body::{SecretBody, AES_GCM_ENCRYPTION_RESERVED}; use crate::protocol::NetPacket; @@ -53,14 +53,14 @@ impl AesGcmCipher { pub fn decrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if !net_packet.is_encrypt() { //未加密的数据直接丢弃 - return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); + return Err(anyhow!("not encrypt")); } if net_packet.payload().len() < AES_GCM_ENCRYPTION_RESERVED { log::error!("数据异常,长度小于{}", AES_GCM_ENCRYPTION_RESERVED); - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } let mut nonce_raw = [0; 12]; nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); @@ -74,7 +74,7 @@ impl AesGcmCipher { if let Some(finger) = &self.finger { let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body()); if &finger != secret_body.finger() { - return Err(io::Error::new(io::ErrorKind::Other, "ring aes finger err")); + return Err(anyhow!("ring aes finger err")); } } @@ -87,10 +87,7 @@ impl AesGcmCipher { } }; if let Err(e) = rs { - return Err(io::Error::new( - io::ErrorKind::Other, - format!("解密失败:{}", e), - )); + return Err(anyhow!("解密失败:{}", e)); } net_packet.set_encrypt_flag(false); net_packet.set_data_len(net_packet.data_len() - AES_GCM_ENCRYPTION_RESERVED)?; @@ -102,7 +99,7 @@ impl AesGcmCipher { pub fn encrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let mut nonce_raw = [0; 12]; nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets()); @@ -128,10 +125,7 @@ impl AesGcmCipher { Ok(tag) => { let tag = tag.as_ref(); if tag.len() != 16 { - return Err(io::Error::new( - io::ErrorKind::Other, - format!("加密tag长度错误:{}", tag.len()), - )); + return Err(anyhow!("加密tag长度错误:{}", tag.len())); } secret_body.set_tag(tag)?; if let Some(finger) = &self.finger { @@ -141,10 +135,7 @@ impl AesGcmCipher { net_packet.set_encrypt_flag(true); Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("加密失败:{}", e), - )), + Err(e) => Err(anyhow!("加密失败:{}", e)), }; } } diff --git a/vnt/src/cipher/chacha20/mod.rs b/vnt/src/cipher/chacha20/mod.rs new file mode 100644 index 0000000..b249c3a --- /dev/null +++ b/vnt/src/cipher/chacha20/mod.rs @@ -0,0 +1,2 @@ +mod rs_chacha20; +pub use rs_chacha20::*; diff --git a/vnt/src/cipher/chacha20/rs_chacha20.rs b/vnt/src/cipher/chacha20/rs_chacha20.rs new file mode 100644 index 0000000..5a78d9a --- /dev/null +++ b/vnt/src/cipher/chacha20/rs_chacha20.rs @@ -0,0 +1,114 @@ +use aes::cipher::Iv; +use anyhow::anyhow; +use chacha20::cipher::{Key, KeyIvInit, StreamCipher}; +use chacha20::ChaCha20; + +use crate::cipher::Finger; +use crate::protocol::body::ChaCah20SecretBody; +use crate::protocol::NetPacket; + +#[derive(Clone)] +pub struct ChaCha20Cipher { + key: [u8; 32], + pub(crate) finger: Option, +} + +impl ChaCha20Cipher { + pub fn new_256(key: [u8; 32], finger: Option) -> Self { + Self { key, finger } + } +} + +impl ChaCha20Cipher { + pub fn key(&self) -> &[u8] { + &self.key + } +} + +impl ChaCha20Cipher { + pub fn decrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + if !net_packet.is_encrypt() { + //未加密的数据直接丢弃 + return Err(anyhow!("not encrypt")); + } + let mut iv = [0; 12]; + 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(); + + let mut secret_body = + ChaCah20SecretBody::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(anyhow!("ChaCha20 finger err")); + } + } + + ChaCha20::new( + Key::::from_slice(&self.key), + Iv::::from_slice(&iv), + ) + .apply_keystream(secret_body.en_body_mut()); + let len = secret_body.en_body().len(); + net_packet.set_encrypt_flag(false); + net_packet.set_payload_len(len)?; + Ok(()) + } + pub fn encrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + let data_len = net_packet.data_len(); + let mut iv = [0; 12]; + 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(_) = &self.finger { + net_packet.set_data_len(data_len + 12)?; + } + let mut secret_body = + ChaCah20SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?; + ChaCha20::new( + Key::::from_slice(&self.key), + Iv::::from_slice(&iv), + ) + .apply_keystream(secret_body.en_body_mut()); + if let Some(finger) = &self.finger { + let finger = finger.calculate_finger(&iv[..12], secret_body.en_body_mut()); + let mut secret_body = ChaCah20SecretBody::new(net_packet.payload_mut(), true)?; + secret_body.set_finger(&finger)?; + } + + net_packet.set_encrypt_flag(true); + Ok(()) + } +} + +#[test] +fn test_chacha20() { + let d = ChaCha20Cipher::new_256([0; 32], Some(Finger::new("123"))); + let mut p = + NetPacket::new_encrypt([0; 13 + crate::protocol::body::ENCRYPTION_RESERVED]).unwrap(); + let src = p.buffer().to_vec(); + d.encrypt_ipv4(&mut p).unwrap(); + d.decrypt_ipv4(&mut p).unwrap(); + assert_eq!(p.buffer(), &src); + + let d = ChaCha20Cipher::new_256([0; 32], None); + let mut p = + NetPacket::new_encrypt([0; 13 + crate::protocol::body::ENCRYPTION_RESERVED]).unwrap(); + let src = p.buffer().to_vec(); + d.encrypt_ipv4(&mut p).unwrap(); + d.decrypt_ipv4(&mut p).unwrap(); + assert_eq!(p.buffer(), &src); +} diff --git a/vnt/src/cipher/chacha20_poly1305/mod.rs b/vnt/src/cipher/chacha20_poly1305/mod.rs new file mode 100644 index 0000000..1b4260f --- /dev/null +++ b/vnt/src/cipher/chacha20_poly1305/mod.rs @@ -0,0 +1,9 @@ +#[cfg(feature = "ring-cipher")] +mod ring_chacha20_poly1305; +#[cfg(feature = "ring-cipher")] +pub use ring_chacha20_poly1305::*; + +#[cfg(not(feature = "ring-cipher"))] +mod rs_chacha20_poly1305; +#[cfg(not(feature = "ring-cipher"))] +pub use rs_chacha20_poly1305::*; diff --git a/vnt/src/cipher/chacha20_poly1305/ring_chacha20_poly1305.rs b/vnt/src/cipher/chacha20_poly1305/ring_chacha20_poly1305.rs new file mode 100644 index 0000000..f27e11a --- /dev/null +++ b/vnt/src/cipher/chacha20_poly1305/ring_chacha20_poly1305.rs @@ -0,0 +1,129 @@ +use anyhow::anyhow; + +use ring::aead; +use ring::aead::{LessSafeKey, UnboundKey}; + +use crate::cipher::Finger; +use crate::protocol::body::{SecretBody, AES_GCM_ENCRYPTION_RESERVED}; +use crate::protocol::NetPacket; + +#[derive(Clone)] +pub struct ChaCha20Poly1305Cipher { + key: Vec, + pub(crate) cipher: LessSafeKey, + pub(crate) finger: Option, +} + +impl ChaCha20Poly1305Cipher { + pub fn new_256(key: [u8; 32], finger: Option) -> Self { + let cipher = LessSafeKey::new(UnboundKey::new(&aead::CHACHA20_POLY1305, &key).unwrap()); + Self { + key: key.to_vec(), + cipher, + finger, + } + } +} + +impl ChaCha20Poly1305Cipher { + pub fn key(&self) -> &[u8] { + &self.key + } +} + +impl ChaCha20Poly1305Cipher { + pub fn decrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + if !net_packet.is_encrypt() { + //未加密的数据直接丢弃 + return Err(anyhow!("not encrypt")); + } + if net_packet.payload().len() < AES_GCM_ENCRYPTION_RESERVED { + log::error!("数据异常,长度小于{}", AES_GCM_ENCRYPTION_RESERVED); + return Err(anyhow!("data err")); + } + let mut nonce_raw = [0; 12]; + nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); + nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets()); + nonce_raw[8] = net_packet.protocol().into(); + nonce_raw[9] = net_packet.transport_protocol(); + nonce_raw[10] = net_packet.is_gateway() as u8; + nonce_raw[11] = net_packet.source_ttl(); + let nonce = aead::Nonce::assume_unique_for_key(nonce_raw); + let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?; + if let Some(finger) = &self.finger { + let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body()); + if &finger != secret_body.finger() { + return Err(anyhow!("ring CHACHA20_POLY1305 finger err")); + } + } + + let rs = self + .cipher + .open_in_place(nonce, aead::Aad::empty(), secret_body.en_body_mut()); + if let Err(e) = rs { + return Err(anyhow!("ring CHACHA20_POLY1305 解密失败:{}", e)); + } + net_packet.set_encrypt_flag(false); + net_packet.set_data_len(net_packet.data_len() - AES_GCM_ENCRYPTION_RESERVED)?; + return Ok(()); + } + /// net_packet 必须预留足够长度 + /// data_len是有效载荷的长度 + /// 返回加密后载荷的长度 + pub fn encrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + let mut nonce_raw = [0; 12]; + nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); + nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets()); + nonce_raw[8] = net_packet.protocol().into(); + nonce_raw[9] = net_packet.transport_protocol(); + nonce_raw[10] = net_packet.is_gateway() as u8; + nonce_raw[11] = net_packet.source_ttl(); + let nonce = aead::Nonce::assume_unique_for_key(nonce_raw); + let data_len = net_packet.data_len() + AES_GCM_ENCRYPTION_RESERVED; + net_packet.set_data_len(data_len)?; + let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?; + let rs = self.cipher.seal_in_place_separate_tag( + nonce, + aead::Aad::empty(), + secret_body.body_mut(), + ); + return match rs { + Ok(tag) => { + let tag = tag.as_ref(); + if tag.len() != 16 { + return Err(anyhow!("加密tag长度错误:{}", tag.len())); + } + secret_body.set_tag(tag)?; + if let Some(finger) = &self.finger { + let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body()); + secret_body.set_finger(&finger)?; + } + net_packet.set_encrypt_flag(true); + Ok(()) + } + Err(e) => Err(anyhow!("ring CHACHA20_POLY1305 加密失败:{}", e)), + }; + } +} + +#[test] +fn test_ring_chacha20_poly1305() { + let d = ChaCha20Poly1305Cipher::new_256([0; 32], Some(Finger::new("123"))); + let mut p = NetPacket::new_encrypt([0; 73]).unwrap(); + let src = p.buffer().to_vec(); + d.encrypt_ipv4(&mut p).unwrap(); + d.decrypt_ipv4(&mut p).unwrap(); + assert_eq!(p.buffer(), &src); + let d = ChaCha20Poly1305Cipher::new_256([0; 32], None); + let mut p = NetPacket::new_encrypt([0; 73]).unwrap(); + let src = p.buffer().to_vec(); + d.encrypt_ipv4(&mut p).unwrap(); + d.decrypt_ipv4(&mut p).unwrap(); + assert_eq!(p.buffer(), &src); +} diff --git a/vnt/src/cipher/chacha20_poly1305/rs_chacha20_poly1305.rs b/vnt/src/cipher/chacha20_poly1305/rs_chacha20_poly1305.rs new file mode 100644 index 0000000..9764f70 --- /dev/null +++ b/vnt/src/cipher/chacha20_poly1305/rs_chacha20_poly1305.rs @@ -0,0 +1,121 @@ +use anyhow::anyhow; +use chacha20poly1305::aead::{Nonce, Tag}; +use chacha20poly1305::{AeadInPlace, ChaCha20Poly1305, Key, KeyInit}; + +use crate::cipher::Finger; +use crate::protocol::body::{SecretBody, AES_GCM_ENCRYPTION_RESERVED}; +use crate::protocol::NetPacket; + +#[derive(Clone)] +pub struct ChaCha20Poly1305Cipher { + key: Vec, + pub(crate) cipher: ChaCha20Poly1305, + pub(crate) finger: Option, +} + +impl ChaCha20Poly1305Cipher { + pub fn new_256(key: [u8; 32], finger: Option) -> Self { + let key: &Key = &key.into(); + let cipher = ChaCha20Poly1305::new(key); + Self { + key: key.to_vec(), + cipher, + finger, + } + } +} +impl ChaCha20Poly1305Cipher { + pub fn key(&self) -> &[u8] { + &self.key + } +} + +impl ChaCha20Poly1305Cipher { + pub fn decrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + if !net_packet.is_encrypt() { + //未加密的数据直接丢弃 + return Err(anyhow!("not encrypt")); + } + if net_packet.payload().len() < AES_GCM_ENCRYPTION_RESERVED { + log::error!("数据异常,长度小于{}", AES_GCM_ENCRYPTION_RESERVED); + return Err(anyhow!("data err")); + } + let mut nonce_raw = [0; 12]; + nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); + nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets()); + nonce_raw[8] = net_packet.protocol().into(); + nonce_raw[9] = net_packet.transport_protocol(); + nonce_raw[10] = net_packet.is_gateway() as u8; + nonce_raw[11] = net_packet.source_ttl(); + let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?; + if let Some(finger) = &self.finger { + let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body()); + if &finger != secret_body.finger() { + return Err(anyhow!("rs CHACHA20_POLY1305 finger err")); + } + } + let nonce: Nonce = nonce_raw.into(); + let tag: Tag = + Tag::::from_slice(secret_body.tag()).clone(); + if let Err(e) = + self.cipher + .decrypt_in_place_detached(&nonce, &[], secret_body.body_mut(), &tag) + { + return Err(anyhow!("rs CHACHA20_POLY1305 decrypt_ipv4 {:?}", e)); + } + net_packet.set_encrypt_flag(false); + net_packet.set_data_len(net_packet.data_len() - AES_GCM_ENCRYPTION_RESERVED)?; + Ok(()) + } + /// net_packet 必须预留足够长度 + /// data_len是有效载荷的长度 + /// 返回加密后载荷的长度 + pub fn encrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + let mut nonce_raw = [0; 12]; + nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); + nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets()); + nonce_raw[8] = net_packet.protocol().into(); + nonce_raw[9] = net_packet.transport_protocol(); + nonce_raw[10] = net_packet.is_gateway() as u8; + nonce_raw[11] = net_packet.source_ttl(); + let nonce = nonce_raw.into(); + let data_len = net_packet.data_len() + AES_GCM_ENCRYPTION_RESERVED; + net_packet.set_data_len(data_len)?; + let mut secret_body = SecretBody::new(net_packet.payload_mut(), self.finger.is_some())?; + let rs = self + .cipher + .encrypt_in_place_detached(&nonce, &[], secret_body.body_mut()); + return match rs { + Ok(tag) => { + let tag: &[u8] = tag.as_ref(); + if tag.len() != 16 { + return Err(anyhow!("加密tag长度错误:{}", tag.len(),)); + } + secret_body.set_tag(tag)?; + if let Some(finger) = &self.finger { + let finger = finger.calculate_finger(&nonce_raw, secret_body.en_body()); + secret_body.set_finger(&finger)?; + } + net_packet.set_encrypt_flag(true); + Ok(()) + } + Err(e) => Err(anyhow!("rs CHACHA20_POLY1305 加密失败:{}", e)), + }; + } +} + +#[test] +fn test_rs_chacha20_poly1305() { + let d = ChaCha20Poly1305Cipher::new_256([0; 32], Some(Finger::new("123"))); + let mut p = NetPacket::new_encrypt([0; 73]).unwrap(); + let src = p.buffer().to_vec(); + d.encrypt_ipv4(&mut p).unwrap(); + d.decrypt_ipv4(&mut p).unwrap(); + assert_eq!(p.buffer(), &src); +} diff --git a/vnt/src/cipher/cipher.rs b/vnt/src/cipher/cipher.rs index 3ab71ff..3aba856 100644 --- a/vnt/src/cipher/cipher.rs +++ b/vnt/src/cipher/cipher.rs @@ -1,51 +1,42 @@ -#[cfg(feature = "aes_ecb")] -#[cfg(not(any(feature = "openssl-vendored", feature = "openssl")))] -use crate::cipher::aes_ecb::AesEcbCipher; use std::fmt::Display; +use std::str::FromStr; + +use anyhow::anyhow; +#[cfg(cipher)] +use sha2::Digest; #[cfg(feature = "aes_cbc")] use crate::cipher::aes_cbc::AesCbcCipher; -#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] -#[cfg(not(feature = "ring-cipher"))] -use crate::cipher::aes_gcm_cipher::AesGcmCipher; #[cfg(feature = "aes_ecb")] -#[cfg(any(feature = "openssl-vendored", feature = "openssl"))] -use crate::cipher::openssl_aes_ecb::AesEcbCipher; +use crate::cipher::aes_ecb::AesEcbCipher; #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] -#[cfg(feature = "ring-cipher")] -use crate::cipher::ring_aes_gcm_cipher::AesGcmCipher; +use crate::cipher::aes_gcm::AesGcmCipher; +#[cfg(feature = "chacha20_poly1305")] +use crate::cipher::chacha20::ChaCha20Cipher; +#[cfg(feature = "chacha20_poly1305")] +use crate::cipher::chacha20_poly1305::ChaCha20Poly1305Cipher; #[cfg(feature = "sm4_cbc")] use crate::cipher::sm4_cbc::Sm4CbcCipher; -#[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" -))] +use crate::cipher::xor::XORCipher; +#[cfg(cipher)] use crate::cipher::Finger; use crate::protocol::NetPacket; -#[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" -))] -use sha2::Digest; -use std::io; -use std::str::FromStr; #[derive(Copy, Clone, Eq, PartialEq, Debug)] pub enum CipherModel { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] AesGcm, + #[cfg(feature = "chacha20_poly1305")] + Chacha20Poly1305, + #[cfg(feature = "chacha20_poly1305")] + Chacha20, #[cfg(feature = "aes_cbc")] AesCbc, #[cfg(feature = "aes_ecb")] AesEcb, #[cfg(feature = "sm4_cbc")] Sm4Cbc, + Xor, None, } @@ -54,61 +45,55 @@ impl Display for CipherModel { let str = match self { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] CipherModel::AesGcm => "aes_gcm".to_string(), + #[cfg(feature = "chacha20_poly1305")] + CipherModel::Chacha20Poly1305 => "chacha20_poly1305".to_string(), + #[cfg(feature = "chacha20_poly1305")] + CipherModel::Chacha20 => "chacha20".to_string(), #[cfg(feature = "aes_cbc")] CipherModel::AesCbc => "aes_cbc".to_string(), #[cfg(feature = "aes_ecb")] CipherModel::AesEcb => "aes_ecb".to_string(), #[cfg(feature = "sm4_cbc")] CipherModel::Sm4Cbc => "sm4_cbc".to_string(), + CipherModel::Xor => "xor".to_string(), CipherModel::None => "none".to_string(), }; write!(f, "{}", str) } } + impl FromStr for CipherModel { type Err = String; fn from_str(s: &str) -> Result { - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] - return Err(format!("not match '{}', no encrypt", s)); - #[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - ))] match s.to_lowercase().trim() { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] "aes_gcm" => Ok(CipherModel::AesGcm), + #[cfg(feature = "chacha20_poly1305")] + "chacha20_poly1305" => Ok(CipherModel::Chacha20Poly1305), + #[cfg(feature = "chacha20_poly1305")] + "chacha20" => Ok(CipherModel::Chacha20), #[cfg(feature = "aes_cbc")] "aes_cbc" => Ok(CipherModel::AesCbc), #[cfg(feature = "aes_ecb")] "aes_ecb" => Ok(CipherModel::AesEcb), #[cfg(feature = "sm4_cbc")] "sm4_cbc" => Ok(CipherModel::Sm4Cbc), + "xor" => Ok(CipherModel::Xor), _ => { let mut enums = String::new(); #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] enums.push_str("/aes_gcm"); + #[cfg(feature = "chacha20_poly1305")] + enums.push_str("/chacha20_poly1305/chacha20"); #[cfg(feature = "aes_cbc")] enums.push_str("/aes_cbc"); #[cfg(feature = "aes_ecb")] enums.push_str("/aes_ecb"); #[cfg(feature = "sm4_cbc")] enums.push_str("/sm4_cbc"); - let str = if enums.is_empty() { - "no encrypt" - } else { - &enums[1..] - }; - Err(format!("not match '{}', enum:{}", s, str)) + enums.push_str("/xor"); + Err(format!("not match '{}', enum:{}", s, &enums[1..])) } } } @@ -118,49 +103,37 @@ impl FromStr for CipherModel { pub enum Cipher { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] AesGcm((AesGcmCipher, Vec)), + #[cfg(feature = "chacha20_poly1305")] + Chacha20Poly1305(ChaCha20Poly1305Cipher), + #[cfg(feature = "chacha20_poly1305")] + Chacha20(ChaCha20Cipher), #[cfg(feature = "aes_cbc")] AesCbc(AesCbcCipher), #[cfg(feature = "aes_ecb")] AesEcb(AesEcbCipher), #[cfg(feature = "sm4_cbc")] Sm4Cbc(Sm4CbcCipher), + Xor(XORCipher), None, } + impl Cipher { - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] - pub fn new_password( - _model: CipherModel, - _password: Option, - _token: Option, - ) -> Self { - Cipher::None - } - #[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - ))] pub fn new_password( model: CipherModel, password: Option, token: Option, ) -> Self { - let finger = token.map(|token| Finger::new(&token)); if let Some(password) = password { - let mut hasher = sha2::Sha256::new(); - hasher.update(password.as_bytes()); - let key: [u8; 32] = hasher.finalize().into(); + #[cfg(cipher)] + let key: [u8; 32] = { + let mut hasher = sha2::Sha256::new(); + hasher.update(password.as_bytes()); + hasher.finalize().into() + }; match model { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] CipherModel::AesGcm => { + let finger = token.map(|token| Finger::new(&token)); if password.len() < 8 { let aes = AesGcmCipher::new_128(key[..16].try_into().unwrap(), finger); Cipher::AesGcm((aes, key[..16].to_vec())) @@ -169,8 +142,21 @@ impl Cipher { Cipher::AesGcm((aes, key.to_vec())) } } + #[cfg(feature = "chacha20_poly1305")] + CipherModel::Chacha20Poly1305 => { + let finger = token.map(|token| Finger::new(&token)); + let chacha = ChaCha20Poly1305Cipher::new_256(key, finger); + Cipher::Chacha20Poly1305(chacha) + } + #[cfg(feature = "chacha20_poly1305")] + CipherModel::Chacha20 => { + let finger = token.map(|token| Finger::new(&token)); + let chacha = ChaCha20Cipher::new_256(key, finger); + Cipher::Chacha20(chacha) + } #[cfg(feature = "aes_cbc")] CipherModel::AesCbc => { + let finger = token.map(|token| Finger::new(&token)); if password.len() < 8 { let aes = AesCbcCipher::new_128(key[..16].try_into().unwrap(), finger); Cipher::AesCbc(aes) @@ -181,6 +167,7 @@ impl Cipher { } #[cfg(feature = "aes_ecb")] CipherModel::AesEcb => { + let finger = token.map(|token| Finger::new(&token)); if password.len() < 8 { let aes = AesEcbCipher::new_128(key[..16].try_into().unwrap(), finger); Cipher::AesEcb(aes) @@ -191,126 +178,97 @@ impl Cipher { } #[cfg(feature = "sm4_cbc")] CipherModel::Sm4Cbc => { + let finger = token.map(|token| Finger::new(&token)); let aes = Sm4CbcCipher::new_128(key[..16].try_into().unwrap(), finger); Cipher::Sm4Cbc(aes) } + CipherModel::Xor => { + let _token = token; + Cipher::Xor(XORCipher::new_256(crate::cipher::xor::simple_hash( + &password, + ))) + } CipherModel::None => Cipher::None, } } else { Cipher::None } } - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] - pub fn new_key(_key: [u8; 32], _token: String) -> io::Result { - Err(io::Error::new(io::ErrorKind::Other, "key error")) + #[cfg(not(any(feature = "aes_gcm", feature = "server_encrypt")))] + pub fn new_key(_key: [u8; 32], _token: String) -> anyhow::Result { + Err(anyhow!("key error")) } - #[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - ))] - pub fn new_key(key: [u8; 32], token: String) -> io::Result { + #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] + pub fn new_key(key: [u8; 32], token: String) -> anyhow::Result { let finger = Some(Finger::new(&token)); match key.len() { - #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] 16 => { let aes = AesGcmCipher::new_128(key[..16].try_into().unwrap(), finger); Ok(Cipher::AesGcm((aes, key[..16].to_vec()))) } - #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] 32 => { let aes = AesGcmCipher::new_256(key, finger); Ok(Cipher::AesGcm((aes, key.to_vec()))) } - _ => Err(io::Error::new(io::ErrorKind::Other, "key error")), + _ => Err(anyhow!("key error")), } } pub fn decrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { match self { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] Cipher::AesGcm((aes_gcm, _)) => aes_gcm.decrypt_ipv4(net_packet), #[cfg(feature = "aes_cbc")] Cipher::AesCbc(aes_cbc) => aes_cbc.decrypt_ipv4(net_packet), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20Poly1305(chacha20poly1305) => chacha20poly1305.decrypt_ipv4(net_packet), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20(chacha20) => chacha20.decrypt_ipv4(net_packet), #[cfg(feature = "aes_ecb")] Cipher::AesEcb(aes_ecb) => aes_ecb.decrypt_ipv4(net_packet), #[cfg(feature = "sm4_cbc")] Cipher::Sm4Cbc(sm4_cbc) => sm4_cbc.decrypt_ipv4(net_packet), + Cipher::Xor(xor) => xor.decrypt_ipv4(net_packet), Cipher::None => { if net_packet.is_encrypt() { - return Err(io::Error::new(io::ErrorKind::Other, "not key")); + return Err(anyhow!("not key")); } Ok(()) } } } - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] - pub fn encrypt_ipv4 + AsMut<[u8]>>( - &self, - _net_packet: &mut NetPacket, - ) -> io::Result<()> { - Ok(()) - } - #[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - ))] pub fn encrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { match self { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] Cipher::AesGcm((aes_gcm, _)) => aes_gcm.encrypt_ipv4(net_packet), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20Poly1305(chacha20poly1305) => chacha20poly1305.encrypt_ipv4(net_packet), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20(chacha20) => chacha20.encrypt_ipv4(net_packet), #[cfg(feature = "aes_cbc")] Cipher::AesCbc(aes_cbc) => aes_cbc.encrypt_ipv4(net_packet), #[cfg(feature = "aes_ecb")] Cipher::AesEcb(aes_ecb) => aes_ecb.encrypt_ipv4(net_packet), #[cfg(feature = "sm4_cbc")] Cipher::Sm4Cbc(sm4_cbc) => sm4_cbc.encrypt_ipv4(net_packet), + Cipher::Xor(xor) => xor.encrypt_ipv4(net_packet), Cipher::None => Ok(()), } } - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] + #[cfg(not(cipher))] pub fn check_finger + AsMut<[u8]>>( &self, _net_packet: &NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { Ok(()) } - #[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - ))] - pub fn check_finger>(&self, net_packet: &NetPacket) -> io::Result<()> { + #[cfg(cipher)] + pub fn check_finger>(&self, net_packet: &NetPacket) -> anyhow::Result<()> { match self { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] Cipher::AesGcm((aes_gcm, _)) => aes_gcm @@ -318,6 +276,18 @@ impl Cipher { .as_ref() .map(|f| f.check_finger(net_packet)) .unwrap_or(Ok(())), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20Poly1305(chacha20poly1305) => chacha20poly1305 + .finger + .as_ref() + .map(|f| f.check_finger(net_packet)) + .unwrap_or(Ok(())), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20(chacha20) => chacha20 + .finger + .as_ref() + .map(|f| f.check_finger(net_packet)) + .unwrap_or(Ok(())), #[cfg(feature = "aes_cbc")] Cipher::AesCbc(aes_cbc) => aes_cbc .finger @@ -336,6 +306,7 @@ impl Cipher { .as_ref() .map(|f| f.check_finger(net_packet)) .unwrap_or(Ok(())), + Cipher::Xor(_) => Ok(()), Cipher::None => Ok(()), } } @@ -343,12 +314,17 @@ impl Cipher { match self { #[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] Cipher::AesGcm((_, key)) => Some(key), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20Poly1305(chacha20poly1305) => Some(chacha20poly1305.key()), + #[cfg(feature = "chacha20_poly1305")] + Cipher::Chacha20(chacha20) => Some(chacha20.key()), #[cfg(feature = "aes_cbc")] Cipher::AesCbc(aes_cbc) => Some(aes_cbc.key()), #[cfg(feature = "aes_ecb")] Cipher::AesEcb(aes_ecb) => Some(aes_ecb.key()), #[cfg(feature = "sm4_cbc")] Cipher::Sm4Cbc(sm4_cbc) => Some(sm4_cbc.key()), + Cipher::Xor(xor) => Some(xor.key()), Cipher::None => None, } } diff --git a/vnt/src/cipher/finger.rs b/vnt/src/cipher/finger.rs index acad057..73d6dc3 100644 --- a/vnt/src/cipher/finger.rs +++ b/vnt/src/cipher/finger.rs @@ -1,4 +1,4 @@ -use std::io; +use anyhow::anyhow; use sha2::Digest; @@ -16,15 +16,15 @@ impl Finger { let hash: [u8; 32] = hasher.finalize().into(); Finger { hash } } - pub fn check_finger>(&self, net_packet: &NetPacket) -> io::Result<()> { + pub fn check_finger>(&self, net_packet: &NetPacket) -> anyhow::Result<()> { if !net_packet.is_encrypt() { //未加密的数据直接丢弃 - return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); + return Err(anyhow!("not encrypt")); } let payload_len = net_packet.payload().len(); if payload_len < 12 { log::error!("数据异常,长度小于{}", 12); - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } let mut nonce_raw = [0; 12]; nonce_raw[0..4].copy_from_slice(&net_packet.source().octets()); @@ -36,7 +36,7 @@ impl Finger { let payload = net_packet.payload(); let finger = self.calculate_finger(&nonce_raw, &payload[..payload_len - 12]); if &finger[..] != &payload[payload_len - 12..] { - return Err(io::Error::new(io::ErrorKind::Other, "finger err")); + return Err(anyhow!("finger err")); } Ok(()) } diff --git a/vnt/src/cipher/mod.rs b/vnt/src/cipher/mod.rs index cc31860..1cc12d7 100644 --- a/vnt/src/cipher/mod.rs +++ b/vnt/src/cipher/mod.rs @@ -1,40 +1,32 @@ -#[cfg(feature = "aes_cbc")] -mod aes_cbc; -#[cfg(feature = "aes_ecb")] -#[cfg(not(any(feature = "openssl-vendored", feature = "openssl")))] -mod aes_ecb; -#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] -#[cfg(not(feature = "ring-cipher"))] -mod aes_gcm_cipher; mod cipher; -#[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" -))] +#[cfg(cipher)] mod finger; -#[cfg(feature = "aes_ecb")] -#[cfg(any(feature = "openssl-vendored", feature = "openssl"))] -mod openssl_aes_ecb; -#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] -#[cfg(feature = "ring-cipher")] -mod ring_aes_gcm_cipher; -#[cfg(feature = "sm4_cbc")] -mod sm4_cbc; pub use cipher::Cipher; pub use cipher::CipherModel; -#[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" -))] +#[cfg(cipher)] pub use finger::Finger; #[cfg(feature = "server_encrypt")] mod rsa_cipher; #[cfg(feature = "server_encrypt")] pub use rsa_cipher::RsaCipher; + +#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))] +mod aes_gcm; + +#[cfg(feature = "chacha20_poly1305")] +mod chacha20; +#[cfg(feature = "chacha20_poly1305")] +mod chacha20_poly1305; + +#[cfg(feature = "aes_ecb")] +mod aes_ecb; + +#[cfg(feature = "aes_cbc")] +mod aes_cbc; + +#[cfg(feature = "sm4_cbc")] +mod sm4_cbc; + +mod xor; +pub use xor::simple_hash; diff --git a/vnt/src/cipher/sm4_cbc/mod.rs b/vnt/src/cipher/sm4_cbc/mod.rs new file mode 100644 index 0000000..77a92eb --- /dev/null +++ b/vnt/src/cipher/sm4_cbc/mod.rs @@ -0,0 +1,2 @@ +mod rs_sm4_cbc; +pub use rs_sm4_cbc::*; diff --git a/vnt/src/cipher/sm4_cbc.rs b/vnt/src/cipher/sm4_cbc/rs_sm4_cbc.rs similarity index 81% rename from vnt/src/cipher/sm4_cbc.rs rename to vnt/src/cipher/sm4_cbc/rs_sm4_cbc.rs index 0c302a5..68b1615 100644 --- a/vnt/src/cipher/sm4_cbc.rs +++ b/vnt/src/cipher/sm4_cbc/rs_sm4_cbc.rs @@ -1,9 +1,9 @@ use crate::cipher::Finger; use crate::protocol::{NetPacket, HEAD_LEN}; +use anyhow::anyhow; use libsm::sm4::cipher_mode::CipherMode; use libsm::sm4::Sm4CipherMode; use rand::RngCore; -use std::io; pub struct Sm4CbcCipher { key: [u8; 16], @@ -41,10 +41,10 @@ impl Sm4CbcCipher { pub fn decrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if !net_packet.is_encrypt() { //未加密的数据直接丢弃 - return Err(io::Error::new(io::ErrorKind::Other, "not encrypt")); + return Err(anyhow!("not encrypt")); } if let Some(finger) = &self.finger { @@ -57,12 +57,12 @@ impl Sm4CbcCipher { nonce_raw[11] = net_packet.source_ttl(); let len = net_packet.payload().len(); if len < 12 { - return Err(io::Error::new(io::ErrorKind::Other, "payload len <12")); + return Err(anyhow!("payload len <12")); } let secret_body = &net_packet.payload()[..len - 12]; let finger = finger.calculate_finger(&nonce_raw, secret_body); if &finger != &net_packet.payload()[len - 12..] { - return Err(io::Error::new(io::ErrorKind::Other, "finger err")); + return Err(anyhow!("finger err")); } net_packet.set_data_len(net_packet.data_len() - finger.len())?; } @@ -70,7 +70,7 @@ impl Sm4CbcCipher { let len = payload.len(); if len < 16 || len > 1024 * 4 { log::error!("数据异常,长度{}小于16或大于4096", len); - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } let mut out = [0u8; 1024 * 4]; let data = &payload[..len - 16]; @@ -79,32 +79,29 @@ impl Sm4CbcCipher { Ok(len) => { let src_net_packet = NetPacket::new(&out[..len])?; if src_net_packet.source() != net_packet.source() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.destination() != net_packet.destination() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.protocol() != net_packet.protocol() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.transport_protocol() != net_packet.transport_protocol() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.is_gateway() != net_packet.is_gateway() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } if src_net_packet.source_ttl() != net_packet.source_ttl() { - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } net_packet.set_data_len(len)?; net_packet.set_payload(src_net_packet.payload())?; net_packet.set_encrypt_flag(false); Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("sm4_cbc解密失败:{}", e), - )), + Err(e) => Err(anyhow!("sm4_cbc解密失败:{}", e)), } } /// net_packet 必须预留足够长度 @@ -112,7 +109,7 @@ impl Sm4CbcCipher { pub fn encrypt_ipv4 + AsMut<[u8]>>( &self, net_packet: &mut NetPacket, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let mut out = [0u8; 1024 * 4]; let mut iv = [0u8; 16]; rand::thread_rng().fill_bytes(&mut iv); @@ -121,7 +118,7 @@ impl Sm4CbcCipher { "数据异常,长度{}大于1024 * 4 - 32", net_packet.buffer().len() ); - return Err(io::Error::new(io::ErrorKind::Other, "data err")); + return Err(anyhow!("data err")); } match self.cipher.encrypt(net_packet.buffer(), &iv, &mut out) { Ok(len) => { @@ -146,10 +143,7 @@ impl Sm4CbcCipher { net_packet.set_encrypt_flag(true); Ok(()) } - Err(e) => Err(io::Error::new( - io::ErrorKind::Other, - format!("sm4_cbc加密失败:{}", e), - )), + Err(e) => Err(anyhow!("sm4_cbc加密失败:{}", e)), } } } diff --git a/vnt/src/cipher/xor/mod.rs b/vnt/src/cipher/xor/mod.rs new file mode 100644 index 0000000..fe38376 --- /dev/null +++ b/vnt/src/cipher/xor/mod.rs @@ -0,0 +1,2 @@ +mod xor; +pub use xor::*; diff --git a/vnt/src/cipher/xor/xor.rs b/vnt/src/cipher/xor/xor.rs new file mode 100644 index 0000000..7a76241 --- /dev/null +++ b/vnt/src/cipher/xor/xor.rs @@ -0,0 +1,84 @@ +use anyhow::anyhow; + +use crate::protocol::NetPacket; + +pub fn simple_hash(input: &str) -> [u8; 32] { + let mut result = [0u8; 32]; + let bytes = input.as_bytes(); + for (index, v) in result.iter_mut().enumerate() { + *v = bytes[index % bytes.len()]; + } + + let mut state = 0u8; + + for (i, &byte) in bytes.iter().enumerate() { + let combined = byte.wrapping_add(state).rotate_left((i % 8) as u32); + result[i % 32] ^= combined; + state = state.wrapping_add(byte).rotate_left(3); + } + + for i in 0..32 { + result[i] = result[i] + .rotate_left((result[(i + 1) % 32] % 8) as u32) + .wrapping_add(state); + state = state.wrapping_add(result[i]).rotate_left(3); + } + + result +} + +#[derive(Clone)] +pub struct XORCipher { + key: [u8; 32], +} + +impl XORCipher { + pub fn new_256(key: [u8; 32]) -> Self { + Self { key } + } +} + +impl XORCipher { + pub fn key(&self) -> &[u8] { + &self.key + } +} + +impl XORCipher { + pub fn decrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + if !net_packet.is_encrypt() { + //未加密的数据直接丢弃 + return Err(anyhow!("not encrypt")); + } + let key = &self.key; + for (i, byte) in net_packet.payload_mut().iter_mut().enumerate() { + *byte ^= key[i & 31]; + } + net_packet.set_encrypt_flag(false); + Ok(()) + } + pub fn encrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> anyhow::Result<()> { + net_packet.set_encrypt_flag(true); + let key = &self.key; + for (i, byte) in net_packet.payload_mut().iter_mut().enumerate() { + *byte ^= key[i & 31]; + } + Ok(()) + } +} + +#[test] +fn test_xor() { + let d = XORCipher::new_256(simple_hash("password")); + let mut p = NetPacket::new_encrypt([0; 1000]).unwrap(); + let src = p.buffer().to_vec(); + d.encrypt_ipv4(&mut p).unwrap(); + d.decrypt_ipv4(&mut p).unwrap(); + assert_eq!(p.buffer(), &src) +} diff --git a/vnt/src/compression/lz4_compress.rs b/vnt/src/compression/lz4_compress.rs new file mode 100644 index 0000000..6c28abc --- /dev/null +++ b/vnt/src/compression/lz4_compress.rs @@ -0,0 +1,33 @@ +use anyhow::anyhow; + +use crate::protocol::NetPacket; + +#[derive(Clone)] +pub struct Lz4Compressor; + +impl Lz4Compressor { + pub fn compress, O: AsRef<[u8]> + AsMut<[u8]>>( + in_net_packet: &NetPacket, + out: &mut NetPacket, + ) -> anyhow::Result<()> { + out.set_data_len_max(); + let len = match lz4_flex::compress_into(in_net_packet.payload(), out.payload_mut()) { + Ok(len) => len, + Err(e) => Err(anyhow!("Lz4 compress {}", e))?, + }; + out.set_payload_len(len)?; + Ok(()) + } + pub fn decompress, O: AsRef<[u8]> + AsMut<[u8]>>( + in_net_packet: &NetPacket, + out: &mut NetPacket, + ) -> anyhow::Result<()> { + out.set_data_len_max(); + let len = match lz4_flex::decompress_into(in_net_packet.payload(), out.payload_mut()) { + Ok(len) => len, + Err(e) => Err(anyhow!("Lz4 decompress {}", e))?, + }; + out.set_payload_len(len)?; + Ok(()) + } +} diff --git a/vnt/src/compression/mod.rs b/vnt/src/compression/mod.rs new file mode 100644 index 0000000..797ed42 --- /dev/null +++ b/vnt/src/compression/mod.rs @@ -0,0 +1,218 @@ +use std::str::FromStr; + +use anyhow::anyhow; + +#[cfg(feature = "lz4_compress")] +use crate::compression::lz4_compress::Lz4Compressor; +#[cfg(feature = "zstd_compress")] +use crate::compression::zstd_compress::ZstdCompressor; +use crate::protocol::extension::CompressionAlgorithm; +#[cfg(feature = "zstd_compress")] +use zstd::zstd_safe::CompressionLevel; + +use crate::protocol::NetPacket; + +#[cfg(feature = "lz4_compress")] +mod lz4_compress; +#[cfg(feature = "zstd_compress")] +mod zstd_compress; + +#[derive(Clone, Copy, Debug)] +pub enum Compressor { + #[cfg(feature = "lz4_compress")] + Lz4, + #[cfg(feature = "zstd_compress")] + Zstd(CompressionLevel), + None, +} + +impl FromStr for Compressor { + type Err = String; + #[cfg(not(any(feature = "lz4_compress", feature = "zstd_compress")))] + fn from_str(s: &str) -> Result { + Err(format!("not match '{}', Compression not supported", s)) + } + #[cfg(any(feature = "lz4_compress", feature = "zstd_compress"))] + fn from_str(s: &str) -> Result { + let str = s.trim().to_lowercase(); + match str.as_str() { + #[cfg(feature = "lz4_compress")] + "lz4" => Ok(Compressor::Lz4), + #[cfg(feature = "zstd_compress")] + "zstd" => Ok(Compressor::Zstd(9)), + "none" => Ok(Compressor::None), + _ => { + #[cfg(feature = "zstd_compress")] + { + let string_array: Vec = str.split(',').map(|s| s.to_string()).collect(); + if string_array.len() != 2 || string_array[0] != "zstd" { + return Err(format!("not match '{}', exp: zstd,10", s)); + } + return match CompressionLevel::from_str(&string_array[1]) { + Ok(level) => Ok(Compressor::Zstd(level)), + Err(_) => Err(format!("not match '{}', exp: zstd,10", s)), + }; + } + #[cfg(not(feature = "zstd_compress"))] + #[cfg(feature = "lz4_compress")] + return Err(format!("not match '{}', exp: lz4", s)); + } + } + } +} + +#[cfg(not(any(feature = "lz4_compress", feature = "zstd_compress")))] +impl Compressor { + pub fn compress, O: AsRef<[u8]> + AsMut<[u8]>>( + &self, + _in_net_packet: &NetPacket, + _out: &mut NetPacket, + ) -> anyhow::Result { + Ok(false) + } + pub fn decompress, O: AsRef<[u8]> + AsMut<[u8]>>( + _algorithm: CompressionAlgorithm, + _in_net_packet: &NetPacket, + _out: &mut NetPacket, + ) -> anyhow::Result<()> { + Err(anyhow!("Unsupported decompress")) + } +} + +#[cfg(any(feature = "lz4_compress", feature = "zstd_compress"))] +impl Compressor { + pub fn compress, O: AsRef<[u8]> + AsMut<[u8]>>( + &self, + in_net_packet: &NetPacket, + out: &mut NetPacket, + ) -> anyhow::Result { + match self { + #[cfg(feature = "lz4_compress")] + Compressor::Lz4 => { + if in_net_packet.data_len() < 128 { + return Ok(false); + } + Lz4Compressor::compress(in_net_packet, out)?; + let mut compression_extension_tail = out.append_compression_extension_tail()?; + compression_extension_tail.set_algorithm(CompressionAlgorithm::Lz4); + //压缩没效果,则放弃压缩 + if out.data_len() >= in_net_packet.data_len() - 16 { + return Ok(false); + } + return Ok(true); + } + #[cfg(feature = "zstd_compress")] + Compressor::Zstd(level) => { + if in_net_packet.data_len() < 128 { + return Ok(false); + } + ZstdCompressor::compress(*level, in_net_packet, out)?; + let mut compression_extension_tail = out.append_compression_extension_tail()?; + compression_extension_tail.set_algorithm(CompressionAlgorithm::Zstd); + //压缩没效果,则放弃压缩 + if out.data_len() >= in_net_packet.data_len() - 16 { + return Ok(false); + } + return Ok(true); + } + Compressor::None => {} + } + Ok(false) + } + pub fn decompress, O: AsRef<[u8]> + AsMut<[u8]>>( + algorithm: CompressionAlgorithm, + in_net_packet: &NetPacket, + out: &mut NetPacket, + ) -> anyhow::Result<()> { + match algorithm { + #[cfg(feature = "lz4_compress")] + CompressionAlgorithm::Lz4 => Lz4Compressor::decompress(in_net_packet, out), + #[cfg(feature = "zstd_compress")] + CompressionAlgorithm::Zstd => ZstdCompressor::decompress(in_net_packet, out), + _ => Err(anyhow!("Unknown decompress {:?}", algorithm)), + } + } +} + +#[test] +fn test_lz4() { + use crate::protocol::extension::{CompressionAlgorithm, ExtensionTailPacket}; + let lz4 = Compressor::Lz4; + let in_packet = NetPacket::new([ + 65, 108, 105, 99, 101, 32, 119, 97, 116, 32, 98, 101, 103, 105, 110, 110, 105, 110, 103, + 32, 116, 111, 32, 103, 101, 116, 32, 118, 101, 114, 121, 32, 116, 105, 114, 101, 100, 32, + 111, 102, 32, 115, 105, 116, 116, 105, 110, 103, 32, 98, 121, 32, 104, 101, 114, 32, 115, + 105, 115, 116, 101, 114, 32, 111, 110, 32, 116, 104, 101, 32, 98, 97, 110, 107, 44, 32, 97, + 110, 100, 32, 111, 102, 32, 104, 97, 118, 105, 110, 103, 32, 110, 111, 116, 104, 105, 110, + 103, 32, 116, 111, 32, 100, 111, 58, 32, 111, 110, 99, 101, 32, 111, 114, 32, 116, 119, + 105, 99, 101, 32, 115, 104, 101, 32, 104, 97, 100, 32, 112, 101, 101, 112, 101, 100, 32, + 105, 110, 116, 111, 32, 116, 104, 101, 32, 98, 111, 111, 107, 32, 104, 101, 114, 32, 115, + 105, 115, 116, 101, 114, 32, 119, 97, 115, 32, 114, 101, 97, 100, 105, 110, 103, 44, 32, + 98, 117, 116, 32, 105, 116, 32, 104, 97, 100, 32, 110, 111, 32, 112, 105, 99, 116, 117, + 114, 101, 115, 32, 111, 114, 32, 99, 111, 110, 118, 101, 114, 115, 97, 116, 105, + ]) + .unwrap(); + let mut out_packet = NetPacket::new([0; 1000]).unwrap(); + let mut src_out_packet = NetPacket::new([0; 1000]).unwrap(); + lz4.compress(&in_packet, &mut out_packet).unwrap(); + let tail = out_packet.split_tail_packet().unwrap(); + match tail { + ExtensionTailPacket::Compression(c) => match c.algorithm() { + CompressionAlgorithm::Lz4 => { + Compressor::decompress(CompressionAlgorithm::Lz4, &out_packet, &mut src_out_packet) + .unwrap(); + } + _ => { + unimplemented!() + } + }, + ExtensionTailPacket::Unknown => { + unimplemented!() + } + } + assert!(!out_packet.is_extension()); + assert_eq!(in_packet.payload(), src_out_packet.payload()) +} +#[test] +fn test_zstd() { + use crate::protocol::extension::{CompressionAlgorithm, ExtensionTailPacket}; + let zstd = Compressor::Zstd(22); + let in_packet = NetPacket::new([ + 65, 108, 105, 99, 101, 32, 119, 97, 115, 32, 98, 101, 103, 105, 110, 110, 105, 110, 103, + 32, 116, 111, 32, 103, 101, 116, 32, 118, 101, 114, 121, 32, 116, 105, 114, 101, 100, 32, + 111, 102, 32, 115, 105, 116, 116, 105, 110, 103, 32, 98, 121, 32, 104, 101, 114, 32, 115, + 105, 115, 116, 101, 114, 32, 111, 110, 32, 116, 104, 101, 32, 98, 97, 110, 107, 44, 32, 97, + 110, 100, 32, 111, 102, 32, 104, 97, 118, 105, 110, 103, 32, 110, 111, 116, 104, 105, 110, + 103, 32, 116, 111, 32, 100, 111, 58, 32, 111, 110, 99, 101, 32, 111, 114, 32, 116, 119, + 105, 99, 101, 32, 115, 104, 101, 32, 104, 97, 100, 32, 112, 101, 101, 112, 101, 100, 32, + 105, 110, 116, 111, 32, 116, 104, 101, 32, 98, 111, 111, 107, 32, 104, 101, 114, 32, 115, + 105, 115, 116, 101, 114, 32, 119, 97, 115, 32, 114, 101, 97, 100, 105, 110, 103, 44, 32, + 98, 117, 116, 32, 105, 116, 32, 104, 97, 100, 32, 110, 111, 32, 112, 105, 99, 116, 117, + 114, 101, 115, 32, 111, 114, 32, 99, 111, 110, 118, 101, 114, 115, 97, 116, 105, + ]) + .unwrap(); + let mut out_packet = NetPacket::new([0; 1000]).unwrap(); + let mut src_out_packet = NetPacket::new([0; 1000]).unwrap(); + zstd.compress(&in_packet, &mut out_packet).unwrap(); + let tail = out_packet.split_tail_packet().unwrap(); + match tail { + ExtensionTailPacket::Compression(c) => match c.algorithm() { + CompressionAlgorithm::Zstd => { + Compressor::decompress( + CompressionAlgorithm::Zstd, + &out_packet, + &mut src_out_packet, + ) + .unwrap(); + } + _ => { + unimplemented!() + } + }, + ExtensionTailPacket::Unknown => { + unimplemented!() + } + } + assert!(!out_packet.is_extension()); + assert_eq!(in_packet.payload(), src_out_packet.payload()) +} diff --git a/vnt/src/compression/zstd_compress.rs b/vnt/src/compression/zstd_compress.rs new file mode 100644 index 0000000..d647ecb --- /dev/null +++ b/vnt/src/compression/zstd_compress.rs @@ -0,0 +1,38 @@ +use crate::protocol::NetPacket; +use anyhow::anyhow; +use zstd::zstd_safe::CompressionLevel; + +#[derive(Clone)] +pub struct ZstdCompressor; + +impl ZstdCompressor { + pub fn compress, O: AsRef<[u8]> + AsMut<[u8]>>( + compression_level: CompressionLevel, + in_net_packet: &NetPacket, + out: &mut NetPacket, + ) -> anyhow::Result<()> { + out.set_data_len_max(); + let len = match zstd::zstd_safe::compress( + out.payload_mut(), + in_net_packet.payload(), + compression_level, + ) { + Ok(len) => len, + Err(e) => Err(anyhow!("zstd compress {}", e))?, + }; + out.set_payload_len(len)?; + Ok(()) + } + pub fn decompress, O: AsRef<[u8]> + AsMut<[u8]>>( + in_net_packet: &NetPacket, + out: &mut NetPacket, + ) -> anyhow::Result<()> { + out.set_data_len_max(); + let len = match zstd::zstd_safe::decompress(out.payload_mut(), in_net_packet.payload()) { + Ok(len) => len, + Err(e) => Err(anyhow!("zstd decompress {}", e))?, + }; + out.set_payload_len(len)?; + Ok(()) + } +} diff --git a/vnt/src/core/conn.rs b/vnt/src/core/conn.rs index 411e9f6..bd80078 100644 --- a/vnt/src/core/conn.rs +++ b/vnt/src/core/conn.rs @@ -38,7 +38,7 @@ pub struct Vnt { current_device: Arc>, nat_test: NatTest, device_list: Arc)>>, - context: ChannelContext, + context: Arc>>, peer_nat_info_map: Arc>>, down_count_watcher: WatchU64Adder, up_count_watcher: WatchSingleU64Adder, @@ -47,7 +47,7 @@ pub struct Vnt { impl Vnt { pub fn new(config: Config, callback: Call) -> anyhow::Result { - log::info!("config:{:?}", config); + log::info!("config.toml:{:?}", config); //服务端非对称加密 #[cfg(feature = "server_encrypt")] let rsa_cipher: Arc>> = Arc::new(Mutex::new(None)); @@ -132,8 +132,11 @@ impl Vnt { // pc上先创建虚拟网卡 #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] let device = { + log::info!("开始创建tun"); let device = tun_tap_device::create_device(&config)?; + log::info!("创建tun成功"); let tun_info = DeviceInfo::new(device.name()?, device.version()?); + log::info!("tun信息{:?}", tun_info); callback.create_tun(tun_info); device }; @@ -177,6 +180,7 @@ impl Vnt { config.parallel, up_counter, device_list.clone(), + config.compressor, ); #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] let device_adapter = DeviceAdapter::new(device.clone()); @@ -273,7 +277,7 @@ impl Vnt { current_device, nat_test, device_list, - context, + context: Arc::new(Mutex::new(Some(context))), peer_nat_info_map, down_count_watcher, up_count_watcher, @@ -390,16 +394,24 @@ impl Vnt { device_list } pub fn route(&self, ip: &Ipv4Addr) -> Option { - self.context.route_table.route_one(ip) + self.context.lock().as_ref()?.route_table.route_one(ip) } pub fn is_gateway(&self, ip: &Ipv4Addr) -> bool { self.current_device.load().is_gateway(ip) } pub fn route_key(&self, route_key: &RouteKey) -> Option { - self.context.route_table.route_to_id(route_key) + self.context + .lock() + .as_ref()? + .route_table + .route_to_id(route_key) } pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec)> { - self.context.route_table.route_table() + if let Some(context) = self.context.lock().as_ref() { + context.route_table.route_table() + } else { + vec![] + } } pub fn up_stream(&self) -> u64 { self.up_count_watcher.get() @@ -408,6 +420,8 @@ impl Vnt { self.down_count_watcher.get() } pub fn stop(&self) { + //退出协助回收资源 + let _ = self.context.lock().take(); self.stop_manager.stop() } pub fn wait(&self) { diff --git a/vnt/src/core/mod.rs b/vnt/src/core/mod.rs index 9ade95b..70fd235 100644 --- a/vnt/src/core/mod.rs +++ b/vnt/src/core/mod.rs @@ -7,6 +7,7 @@ pub use conn::Vnt; use crate::channel::punch::PunchModel; use crate::channel::UseChannelType; use crate::cipher::CipherModel; +use crate::compression::Compressor; use crate::util::{address_choose, dns_query_all}; mod conn; @@ -46,6 +47,7 @@ pub struct Config { // 端口映射 #[cfg(feature = "port_mapping")] pub port_mapping_list: Vec<(bool, SocketAddr, String)>, + pub compressor: Compressor, } impl Config { @@ -77,6 +79,7 @@ impl Config { packet_delay: u32, // 例如 [udp:127.0.0.1:80->10.26.0.10:8080,tcp:127.0.0.1:80->10.26.0.10:8080] #[cfg(feature = "port_mapping")] port_mapping_list: Vec, + compressor: Compressor, ) -> anyhow::Result { for x in stun_server.iter_mut() { if !x.contains(":") { @@ -140,36 +143,33 @@ impl Config { packet_delay, #[cfg(feature = "port_mapping")] port_mapping_list, + compressor, }) } } + impl Config { - #[cfg(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - ))] pub fn password_hash(&self) -> Option<[u8; 16]> { - self.password.as_ref().map(|v| { - use sha2::Digest; - let mut hasher = sha2::Sha256::new(); - hasher.update(self.cipher_model.to_string().as_bytes()); - hasher.update(v.as_bytes()); - hasher.update(self.token.as_bytes()); - let key: [u8; 32] = hasher.finalize().into(); - key[16..].try_into().unwrap() - }) - } - #[cfg(not(any( - feature = "aes_gcm", - feature = "server_encrypt", - feature = "aes_cbc", - feature = "aes_ecb", - feature = "sm4_cbc" - )))] - pub fn password_hash(&self) -> Option<[u8; 16]> { - None + if let Some(p) = self.password.as_ref() { + match self.cipher_model { + CipherModel::Xor => { + let key = crate::cipher::simple_hash(&format!("Xor{}{}", p, self.token)); + Some(key[16..].try_into().unwrap()) + } + CipherModel::None => None, + #[cfg(cipher)] + _ => { + use sha2::Digest; + let mut hasher = sha2::Sha256::new(); + hasher.update(self.cipher_model.to_string().as_bytes()); + hasher.update(p.as_bytes()); + hasher.update(self.token.as_bytes()); + let key: [u8; 32] = hasher.finalize().into(); + Some(key[16..].try_into().unwrap()) + } + } + } else { + None + } } } diff --git a/vnt/src/handle/extension/mod.rs b/vnt/src/handle/extension/mod.rs new file mode 100644 index 0000000..45e13cb --- /dev/null +++ b/vnt/src/handle/extension/mod.rs @@ -0,0 +1,24 @@ +use crate::compression::Compressor; +use crate::protocol::extension::ExtensionTailPacket; +use crate::protocol::NetPacket; +use anyhow::anyhow; + +pub fn handle_extension_tail + AsMut<[u8]>, O: AsRef<[u8]> + AsMut<[u8]>>( + in_net_packet: &mut NetPacket, + out: &mut NetPacket, +) -> anyhow::Result { + if in_net_packet.is_extension() { + let tail_packet = in_net_packet.split_tail_packet()?; + match tail_packet { + ExtensionTailPacket::Compression(extension) => { + let compression_algorithm = extension.algorithm(); + Compressor::decompress(compression_algorithm, &in_net_packet, out)?; + out.head_mut().copy_from_slice(in_net_packet.head()); + Ok(true) + } + ExtensionTailPacket::Unknown => Err(anyhow!("Unknown decompress")), + } + } else { + Ok(false) + } +} diff --git a/vnt/src/handle/maintain/heartbeat.rs b/vnt/src/handle/maintain/heartbeat.rs index d9ce1dc..e5b4a94 100644 --- a/vnt/src/handle/maintain/heartbeat.rs +++ b/vnt/src/handle/maintain/heartbeat.rs @@ -1,4 +1,3 @@ -use std::io; use std::net::Ipv4Addr; use std::sync::Arc; use std::time::Duration; @@ -167,7 +166,7 @@ fn client_relay0( current_device: &CurrentDeviceInfo, device_list: &Mutex<(u16, Vec)>, client_cipher: &Cipher, -) -> io::Result<()> { +) -> anyhow::Result<()> { // 离线了不再探测 if current_device.status.offline() { return Ok(()); @@ -211,7 +210,7 @@ fn client_relay0( fn heartbeat_packet( src: Ipv4Addr, dest: Ipv4Addr, -) -> io::Result> { +) -> anyhow::Result> { let mut net_packet = NetPacket::new_encrypt([0u8; 12 + 4 + ENCRYPTION_RESERVED])?; net_packet.set_default_version(); net_packet.set_protocol(Protocol::Control); @@ -228,7 +227,7 @@ fn heartbeat_packet_client( client_cipher: &Cipher, src: Ipv4Addr, dest: Ipv4Addr, -) -> io::Result> { +) -> anyhow::Result> { let mut net_packet = heartbeat_packet(src, dest)?; client_cipher.encrypt_ipv4(&mut net_packet)?; Ok(net_packet) @@ -239,7 +238,7 @@ fn heartbeat_packet_server( server_cipher: &Cipher, src: Ipv4Addr, dest: Ipv4Addr, -) -> io::Result> { +) -> anyhow::Result> { let mut net_packet = heartbeat_packet(src, dest)?; let mut ping = PingPacket::new(net_packet.payload_mut())?; ping.set_epoch(device_list.lock().0); diff --git a/vnt/src/handle/maintain/punch.rs b/vnt/src/handle/maintain/punch.rs index 8c7c226..8c00f45 100644 --- a/vnt/src/handle/maintain/punch.rs +++ b/vnt/src/handle/maintain/punch.rs @@ -2,9 +2,10 @@ use std::collections::HashMap; use std::net::Ipv4Addr; use std::sync::mpsc::{sync_channel, Receiver, SyncSender}; use std::sync::Arc; +use std::thread; use std::time::Duration; -use std::{io, thread}; +use anyhow::anyhow; use crossbeam_utils::atomic::AtomicCell; use parking_lot::Mutex; use protobuf::Message; @@ -222,7 +223,7 @@ fn punch0( punch_record: &Mutex>, last_punch_record: &mut HashMap, total_count: usize, -) -> io::Result<()> { +) -> anyhow::Result<()> { let nat_info = nat_test.nat_info(); if total_count < 10 && (nat_info.public_ips.is_empty() @@ -297,7 +298,7 @@ fn punch_packet( virtual_ip: Ipv4Addr, nat_info: &NatInfo, dest: Ipv4Addr, -) -> io::Result>> { +) -> anyhow::Result>> { let mut punch_reply = PunchInfo::new(); punch_reply.reply = false; punch_reply.public_ip_list = nat_info @@ -320,7 +321,7 @@ fn punch_packet( log::info!("请求打洞={:?}", punch_reply); let bytes = punch_reply .write_to_bytes() - .map_err(|e| io::Error::new(io::ErrorKind::Other, format!("punch_packet {:?}", e)))?; + .map_err(|e| anyhow!("punch_packet {:?}", e))?; let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?; net_packet.set_default_version(); net_packet.set_protocol(Protocol::OtherTurn); diff --git a/vnt/src/handle/mod.rs b/vnt/src/handle/mod.rs index fc0c10d..e59631d 100644 --- a/vnt/src/handle/mod.rs +++ b/vnt/src/handle/mod.rs @@ -2,6 +2,7 @@ use crossbeam_utils::atomic::AtomicCell; use std::net::{Ipv4Addr, SocketAddr}; pub mod callback; +mod extension; pub mod handshaker; pub mod maintain; pub mod recv_data; diff --git a/vnt/src/handle/recv_data/client.rs b/vnt/src/handle/recv_data/client.rs index ee748c6..5d619d6 100644 --- a/vnt/src/handle/recv_data/client.rs +++ b/vnt/src/handle/recv_data/client.rs @@ -1,19 +1,23 @@ -use parking_lot::RwLock; -use protobuf::Message; +use anyhow::anyhow; use std::collections::HashMap; -use std::io; use std::net::{Ipv4Addr, Ipv6Addr}; use std::sync::Arc; +use parking_lot::RwLock; +use protobuf::Message; + use packet::icmp::{icmp, Kind}; use packet::ip::ipv4; use packet::ip::ipv4::packet::IpV4Packet; +#[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] +use tun::device::IFace; use crate::channel::context::ChannelContext; use crate::channel::punch::NatInfo; use crate::channel::{Route, RouteKey}; use crate::cipher::Cipher; use crate::external_route::AllowExternalRoute; +use crate::handle::extension::handle_extension_tail; use crate::handle::maintain::PunchSender; use crate::handle::recv_data::PacketHandler; use crate::handle::CurrentDeviceInfo; @@ -27,8 +31,6 @@ use crate::protocol::{ control_packet, ip_turn_packet, other_turn_packet, NetPacket, Protocol, MAX_TTL, }; use crate::tun_tap_device::tun_create_helper::DeviceAdapter; -#[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] -use tun::device::IFace; /// 处理来源于客户端的包 #[derive(Clone)] @@ -70,14 +72,26 @@ impl PacketHandler for ClientPacketHandler { fn handle( &self, mut net_packet: NetPacket<&mut [u8]>, + mut extend: NetPacket<&mut [u8]>, route_key: RouteKey, context: &ChannelContext, current_device: &CurrentDeviceInfo, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { self.client_cipher.decrypt_ipv4(&mut net_packet)?; context .route_table .update_read_time(&net_packet.source(), &route_key); + //处理扩展 + let net_packet = if net_packet.is_extension() { + //这样重用数组,减少一次数据拷贝 + if handle_extension_tail(&mut net_packet, &mut extend)? { + extend + } else { + net_packet + } + } else { + net_packet + }; match net_packet.protocol() { Protocol::Service => {} Protocol::Error => {} @@ -103,7 +117,7 @@ impl ClientPacketHandler { context: &ChannelContext, current_device: &CurrentDeviceInfo, route_key: RouteKey, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let destination = net_packet.destination(); let source = net_packet.source(); match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) { @@ -190,7 +204,7 @@ impl ClientPacketHandler { current_device: &CurrentDeviceInfo, mut net_packet: NetPacket<&mut [u8]>, route_key: RouteKey, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let metric = net_packet.source_ttl() - net_packet.ttl() + 1; let source = net_packet.source(); match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? { @@ -278,17 +292,15 @@ impl ClientPacketHandler { current_device: &CurrentDeviceInfo, net_packet: NetPacket<&mut [u8]>, route_key: RouteKey, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if context.use_channel_type().is_only_relay() { return Ok(()); } let source = net_packet.source(); match other_turn_packet::Protocol::from(net_packet.transport_protocol()) { other_turn_packet::Protocol::Punch => { - let mut punch_info = - PunchInfo::parse_from_bytes(net_packet.payload()).map_err(|e| { - io::Error::new(io::ErrorKind::Other, format!("PunchInfo {:?}", e)) - })?; + let mut punch_info = PunchInfo::parse_from_bytes(net_packet.payload()) + .map_err(|e| anyhow!("PunchInfo {:?}", e))?; let public_ips = punch_info .public_ip_list .iter() @@ -348,9 +360,9 @@ impl ClientPacketHandler { punch_reply.ipv6 = ipv6.octets().to_vec(); punch_reply.ipv6_port = nat_info.udp_ports[0] as u32; } - let bytes = punch_reply.write_to_bytes().map_err(|e| { - io::Error::new(io::ErrorKind::Other, format!("punch_reply {:?}", e)) - })?; + let bytes = punch_reply + .write_to_bytes() + .map_err(|e| anyhow!("punch_reply {:?}", e))?; let mut punch_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?; punch_packet.set_default_version(); diff --git a/vnt/src/handle/recv_data/mod.rs b/vnt/src/handle/recv_data/mod.rs index b644eda..bb14c20 100644 --- a/vnt/src/handle/recv_data/mod.rs +++ b/vnt/src/handle/recv_data/mod.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use std::net::Ipv4Addr; use std::sync::Arc; -use std::{io, thread}; +use std::thread; use crossbeam_utils::atomic::AtomicCell; use parking_lot::{Mutex, RwLock}; @@ -43,7 +43,13 @@ pub struct RecvDataHandler { } impl RecvChannelHandler for RecvDataHandler { - fn handle(&mut self, buf: &mut [u8], route_key: RouteKey, context: &ChannelContext) { + fn handle( + &mut self, + buf: &mut [u8], + extend: &mut [u8], + route_key: RouteKey, + context: &ChannelContext, + ) { //判断stun响应包 if !route_key.is_tcp() { if let Ok(rs) = self @@ -55,8 +61,13 @@ impl RecvChannelHandler for RecvDataHandler { } } } - if let Err(e) = self.handle0(buf, route_key, context) { - log::error!("[{}]-{:?}", thread::current().name().unwrap_or(""), e); + if let Err(e) = self.handle0(buf, extend, route_key, context) { + log::error!( + "[{}]-{:?}-{:?}", + thread::current().name().unwrap_or(""), + route_key.addr, + e + ); } } } @@ -116,12 +127,14 @@ impl RecvDataHandler { fn handle0( &mut self, buf: &mut [u8], + extend: &mut [u8], route_key: RouteKey, context: &ChannelContext, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { // 统计流量 self.counter.add(buf.len() as _); let net_packet = NetPacket::new(buf)?; + let extend = NetPacket::unchecked(extend); if net_packet.ttl() == 0 || net_packet.source_ttl() < net_packet.ttl() { log::warn!("丢弃过时包:{:?}", net_packet.head()); return Ok(()); @@ -139,16 +152,16 @@ impl RecvDataHandler { if net_packet.is_gateway() { //服务端-客户端包 self.server - .handle(net_packet, route_key, context, ¤t_device) + .handle(net_packet, extend, route_key, context, ¤t_device) } else { //客户端-客户端包 self.client - .handle(net_packet, route_key, context, ¤t_device) + .handle(net_packet, extend, route_key, context, ¤t_device) } } else { //转发包 self.turn - .handle(net_packet, route_key, context, ¤t_device) + .handle(net_packet, extend, route_key, context, ¤t_device) } } } @@ -157,8 +170,9 @@ pub trait PacketHandler { fn handle( &self, net_packet: NetPacket<&mut [u8]>, + extend: NetPacket<&mut [u8]>, route_key: RouteKey, context: &ChannelContext, current_device: &CurrentDeviceInfo, - ) -> io::Result<()>; + ) -> anyhow::Result<()>; } diff --git a/vnt/src/handle/recv_data/server.rs b/vnt/src/handle/recv_data/server.rs index 877ecd9..747e3be 100644 --- a/vnt/src/handle/recv_data/server.rs +++ b/vnt/src/handle/recv_data/server.rs @@ -1,3 +1,4 @@ +use anyhow::anyhow; use std::io; use std::net::Ipv4Addr; use std::sync::Arc; @@ -11,6 +12,8 @@ use protobuf::Message; use packet::icmp::{icmp, Kind}; use packet::ip::ipv4; use packet::ip::ipv4::packet::IpV4Packet; +#[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] +use tun::device::IFace; use crate::channel::context::ChannelContext; use crate::channel::{Route, RouteKey}; @@ -34,8 +37,6 @@ use crate::protocol::error_packet::InErrorPacket; use crate::protocol::{ip_turn_packet, service_packet, NetPacket, Protocol, MAX_TTL}; use crate::tun_tap_device::tun_create_helper::DeviceAdapter; use crate::{proto, PeerClientInfo}; -#[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] -use tun::device::IFace; /// 处理来源于服务端的包 #[derive(Clone)] @@ -94,10 +95,11 @@ impl PacketHandler for ServerPacketHandler { fn handle( &self, mut net_packet: NetPacket<&mut [u8]>, + _extend: NetPacket<&mut [u8]>, route_key: RouteKey, context: &ChannelContext, current_device: &CurrentDeviceInfo, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { context .route_table .update_read_time(&net_packet.source(), &route_key); @@ -135,10 +137,8 @@ impl PacketHandler for ServerPacketHandler { } else if net_packet.protocol() == Protocol::Service && net_packet.transport_protocol() == service_packet::Protocol::HandshakeResponse.into() { - let response = - HandshakeResponse::parse_from_bytes(net_packet.payload()).map_err(|e| { - io::Error::new(io::ErrorKind::Other, format!("HandshakeResponse {:?}", e)) - })?; + let response = HandshakeResponse::parse_from_bytes(net_packet.payload()) + .map_err(|e| anyhow!("HandshakeResponse {:?}", e))?; log::info!("握手响应:{:?},{}", route_key, response); //如果开启了加密,则发送加密握手请求 #[cfg(feature = "server_encrypt")] @@ -253,7 +253,7 @@ impl ServerPacketHandler { current_device: &CurrentDeviceInfo, net_packet: NetPacket<&mut [u8]>, route_key: RouteKey, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { match service_packet::Protocol::from(net_packet.transport_protocol()) { service_packet::Protocol::RegistrationResponse => { let response = RegistrationResponse::parse_from_bytes(net_packet.payload()) @@ -439,7 +439,7 @@ impl ServerPacketHandler { &self, current_device: &CurrentDeviceInfo, context: &ChannelContext, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { if current_device.status.online() { log::info!("已连接的不需要注册,{:?}", self.config_info); return Ok(()); @@ -468,7 +468,8 @@ impl ServerPacketHandler { )?; log::info!("发送注册请求,{:?}", self.config_info); //注册请求只发送到默认通道 - context.send_default(response.buffer(), current_device.connect_server) + context.send_default(response.buffer(), current_device.connect_server)?; + Ok(()) } fn error( &self, @@ -526,7 +527,7 @@ impl ServerPacketHandler { current_device: &CurrentDeviceInfo, net_packet: NetPacket<&mut [u8]>, route_key: RouteKey, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? { ControlPacket::PongPacket(pong_packet) => { let current_time = crate::handle::now_time() as u16; diff --git a/vnt/src/handle/recv_data/turn.rs b/vnt/src/handle/recv_data/turn.rs index fa7e814..f4d70ff 100644 --- a/vnt/src/handle/recv_data/turn.rs +++ b/vnt/src/handle/recv_data/turn.rs @@ -3,6 +3,7 @@ use crate::channel::RouteKey; use crate::handle::recv_data::PacketHandler; use crate::handle::CurrentDeviceInfo; use crate::protocol::NetPacket; +use anyhow::Context; /// 处理客户端中转包 #[derive(Clone)] @@ -18,10 +19,11 @@ impl PacketHandler for TurnPacketHandler { fn handle( &self, mut net_packet: NetPacket<&mut [u8]>, + _extend: NetPacket<&mut [u8]>, route_key: RouteKey, context: &ChannelContext, _current_device: &CurrentDeviceInfo, - ) -> std::io::Result<()> { + ) -> anyhow::Result<()> { // ttl减一 let ttl = net_packet.incr_ttl(); if ttl > 0 { @@ -33,7 +35,9 @@ impl PacketHandler for TurnPacketHandler { return Ok(()); } if route.metric <= ttl { - return context.send_by_key(net_packet.buffer(), route.route_key()); + return context + .send_by_key(net_packet.buffer(), route.route_key()) + .context("转发失败"); } } //其他没有路由的不转发 diff --git a/vnt/src/handle/registrar.rs b/vnt/src/handle/registrar.rs index 160850d..417fea5 100644 --- a/vnt/src/handle/registrar.rs +++ b/vnt/src/handle/registrar.rs @@ -1,4 +1,4 @@ -use std::io; +use anyhow::anyhow; use std::net::Ipv4Addr; use protobuf::Message; @@ -19,7 +19,7 @@ pub fn registration_request_packet( is_fast: bool, allow_ip_change: bool, client_secret_hash: Option<&[u8]>, -) -> io::Result>> { +) -> anyhow::Result>> { let mut request = RegistrationRequest::new(); request.token = token; request.device_id = device_id; @@ -36,9 +36,9 @@ pub fn registration_request_packet( .client_secret_hash .extend_from_slice(client_secret_hash); } - let bytes = request.write_to_bytes().map_err(|e| { - io::Error::new(io::ErrorKind::Other, format!("RegistrationRequest {:?}", e)) - })?; + let bytes = request + .write_to_bytes() + .map_err(|e| anyhow!("RegistrationRequest {:?}", e))?; let buf = vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED]; let mut net_packet = NetPacket::new_encrypt(buf)?; net_packet.set_destination(GATEWAY_IP); diff --git a/vnt/src/handle/tun_tap/tun_handler.rs b/vnt/src/handle/tun_tap/tun_handler.rs index 94e1bd7..9729f67 100644 --- a/vnt/src/handle/tun_tap/tun_handler.rs +++ b/vnt/src/handle/tun_tap/tun_handler.rs @@ -13,7 +13,9 @@ use tun::device::IFace; use tun::Device; use crate::channel::context::ChannelContext; +use crate::channel::BUFFER_SIZE; use crate::cipher::Cipher; +use crate::compression::Compressor; use crate::external_route::ExternalRoute; use crate::handle::tun_tap::channel_group::channel_group; use crate::handle::{check_dest, CurrentDeviceInfo, PeerDeviceInfo}; @@ -27,7 +29,7 @@ use crate::protocol::ip_turn_packet::BroadcastPacket; use crate::protocol::{ip_turn_packet, NetPacket, MAX_TTL}; use crate::util::{SingleU64Adder, StopManager}; -fn icmp(device_writer: &Device, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> io::Result<()> { +fn icmp(device_writer: &Device, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> anyhow::Result<()> { if ipv4_packet.protocol() == Protocol::Icmp { let mut icmp = IcmpPacket::new(ipv4_packet.payload_mut())?; if icmp.kind() == Kind::EchoRequest { @@ -43,43 +45,6 @@ fn icmp(device_writer: &Device, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> io::R Ok(()) } -/// 接收tun数据,并且转发到udp上 -pub(crate) fn handle( - context: &ChannelContext, - data: &mut [u8], - len: usize, - device_writer: &Device, - current_device: CurrentDeviceInfo, - ip_route: &ExternalRoute, - #[cfg(feature = "ip_proxy")] proxy_map: &Option, - client_cipher: &Cipher, - server_cipher: &Cipher, - device_list: &Mutex<(u16, Vec)>, -) -> io::Result<()> { - //忽略掉结构不对的情况(ipv6数据、win tap会读到空数据),不然日志打印太多了 - let ipv4_packet = match IpV4Packet::new(&mut data[12..len]) { - Ok(packet) => packet, - Err(_) => return Ok(()), - }; - let src_ip = ipv4_packet.source_ip(); - let dest_ip = ipv4_packet.destination_ip(); - if src_ip == dest_ip { - return icmp(&device_writer, ipv4_packet); - } - return base_handle( - context, - data, - len, - current_device, - ip_route, - #[cfg(feature = "ip_proxy")] - proxy_map, - client_cipher, - server_cipher, - device_list, - ); -} - pub fn start( stop_manager: StopManager, context: ChannelContext, @@ -92,6 +57,7 @@ pub fn start( parallel: usize, mut up_counter: SingleU64Adder, device_list: Arc)>>, + compressor: Compressor, ) -> io::Result<()> { if parallel > 1 { let (sender, receivers) = channel_group::<(Vec, usize)>(parallel, 16); @@ -108,6 +74,7 @@ pub fn start( thread::Builder::new() .name(format!("tunHandler-{}", index)) .spawn(move || { + let mut extend = [0; BUFFER_SIZE]; while let Ok((mut buf, len)) = receiver.recv() { #[cfg(not(target_os = "macos"))] let start = 0; @@ -117,6 +84,7 @@ pub fn start( &context, &mut buf[start..], len, + &mut extend, &device, current_device.load(), &ip_route, @@ -125,6 +93,7 @@ pub fn start( &client_cipher, &server_cipher, &device_list, + &compressor, ) { Ok(_) => {} Err(e) => { @@ -162,6 +131,7 @@ pub fn start( server_cipher, &mut up_counter, device_list, + compressor, ) { log::warn!("stop:{}", e); } @@ -176,7 +146,7 @@ fn broadcast( net_packet: &mut NetPacket<&mut [u8]>, current_device: &CurrentDeviceInfo, device_list: &Mutex<(u16, Vec)>, -) -> io::Result<()> { +) -> anyhow::Result<()> { let list: Vec = device_list .lock() .1 @@ -250,29 +220,43 @@ fn broadcast( broadcast.set_address(&p2p_ips)?; broadcast.set_data(net_packet.buffer())?; server_cipher.encrypt_ipv4(&mut server_packet)?; - sender.send_default(server_packet.buffer(), current_device.connect_server) + sender.send_default(server_packet.buffer(), current_device.connect_server)?; + Ok(()) } +/// 接收tun数据,并且转发到udp上 /// 实现一个原地发送,必须保证是如下结构 /// |12字节开头|ip报文|至少1024字节结尾| /// -#[inline] -fn base_handle( +pub(crate) fn handle( context: &ChannelContext, buf: &mut [u8], data_len: usize, //数据总长度=12+ip包长度 + extend: &mut [u8], + device_writer: &Device, current_device: CurrentDeviceInfo, ip_route: &ExternalRoute, #[cfg(feature = "ip_proxy")] proxy_map: &Option, client_cipher: &Cipher, server_cipher: &Cipher, device_list: &Mutex<(u16, Vec)>, -) -> io::Result<()> { - let ipv4_packet = IpV4Packet::new(&buf[12..data_len])?; + compressor: &Compressor, +) -> anyhow::Result<()> { + //忽略掉结构不对的情况(ipv6数据、win tap会读到空数据),不然日志打印太多了 + let ipv4_packet = match IpV4Packet::new(&mut buf[12..data_len]) { + Ok(packet) => packet, + Err(_) => return Ok(()), + }; + let src_ip = ipv4_packet.source_ip(); + let dest_ip = ipv4_packet.destination_ip(); + if src_ip == dest_ip { + return icmp(&device_writer, ipv4_packet); + } let protocol = ipv4_packet.protocol(); let src_ip = ipv4_packet.source_ip(); let mut dest_ip = ipv4_packet.destination_ip(); let mut net_packet = NetPacket::new0(data_len, buf)?; + let mut out = NetPacket::unchecked(extend); net_packet.set_default_version(); net_packet.set_protocol(protocol::Protocol::IpTurn); net_packet.set_transport_protocol(ip_turn_packet::Protocol::Ipv4.into()); @@ -288,11 +272,49 @@ fn base_handle( } return Ok(()); } + if !dest_ip.is_multicast() && !dest_ip.is_broadcast() && current_device.broadcast_ip != dest_ip + { + if !check_dest( + dest_ip, + current_device.virtual_netmask, + current_device.virtual_network, + ) { + if let Some(r_dest_ip) = ip_route.route(&dest_ip) { + //路由的目标不能是自己 + if r_dest_ip == src_ip { + return Ok(()); + } + //需要修改目的地址 + dest_ip = r_dest_ip; + net_packet.set_destination(r_dest_ip); + } else { + return Ok(()); + } + } + #[cfg(feature = "ip_proxy")] + if let Some(proxy_map) = proxy_map { + let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?; + proxy_map.send_handle(&mut ipv4_packet)?; + } + } + if dest_ip.is_multicast() { //当作广播处理 dest_ip = Ipv4Addr::BROADCAST; net_packet.set_destination(Ipv4Addr::BROADCAST); } + + let mut net_packet = if compressor.compress(&net_packet, &mut out)? { + out.set_default_version(); + out.set_protocol(protocol::Protocol::IpTurn); + out.set_transport_protocol(ip_turn_packet::Protocol::Ipv4.into()); + out.first_set_ttl(6); + out.set_source(src_ip); + out.set_destination(dest_ip); + out + } else { + net_packet + }; if dest_ip.is_broadcast() || current_device.broadcast_ip == dest_ip { // 广播 发送到直连目标 client_cipher.encrypt_ipv4(&mut net_packet)?; @@ -305,33 +327,13 @@ fn base_handle( )?; return Ok(()); } - if !check_dest( - dest_ip, - current_device.virtual_netmask, - current_device.virtual_network, - ) { - if let Some(r_dest_ip) = ip_route.route(&dest_ip) { - //路由的目标不能是自己 - if r_dest_ip == src_ip { - return Ok(()); - } - //需要修改目的地址 - dest_ip = r_dest_ip; - net_packet.set_destination(r_dest_ip); - } else { - return Ok(()); - } - } - #[cfg(feature = "ip_proxy")] - if let Some(proxy_map) = proxy_map { - let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?; - proxy_map.send_handle(&mut ipv4_packet)?; - } + client_cipher.encrypt_ipv4(&mut net_packet)?; context.send_ipv4_by_id( net_packet.buffer(), &dest_ip, current_device.connect_server, current_device.status.online(), - ) + )?; + Ok(()) } diff --git a/vnt/src/handle/tun_tap/unix.rs b/vnt/src/handle/tun_tap/unix.rs index a397623..c49323a 100644 --- a/vnt/src/handle/tun_tap/unix.rs +++ b/vnt/src/handle/tun_tap/unix.rs @@ -1,5 +1,7 @@ use crate::channel::context::ChannelContext; +use crate::channel::BUFFER_SIZE; use crate::cipher::Cipher; +use crate::compression::Compressor; use crate::external_route::ExternalRoute; use crate::handle::tun_tap::channel_group::GroupSyncSender; use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; @@ -30,7 +32,8 @@ pub(crate) fn start_simple( server_cipher: Cipher, up_counter: &mut SingleU64Adder, device_list: Arc)>>, -) -> io::Result<()> { + compressor: Compressor, +) -> anyhow::Result<()> { let poll = Poll::new()?; let waker = Arc::new(Waker::new(poll.registry(), STOP)?); let _waker = waker.clone(); @@ -49,6 +52,7 @@ pub(crate) fn start_simple( server_cipher, up_counter, device_list, + compressor, ) { log::error!("{:?}", e); }; @@ -68,8 +72,10 @@ fn start_simple0( server_cipher: Cipher, up_counter: &mut SingleU64Adder, device_list: Arc)>>, -) -> io::Result<()> { - let mut buf = [0; 1024 * 16]; + compressor: Compressor, +) -> anyhow::Result<()> { + let mut buf = [0; BUFFER_SIZE]; + let mut extend = [0; BUFFER_SIZE]; let fd = device.as_tun_fd(); fd.set_nonblock()?; SourceFd(&fd.as_raw_fd()).register(poll.registry(), FD, Interest::READABLE)?; @@ -102,6 +108,7 @@ fn start_simple0( context, &mut buf, len, + &mut extend, &device, current_device.load(), &ip_route, @@ -110,6 +117,7 @@ fn start_simple0( &client_cipher, &server_cipher, &device_list, + &compressor, ) { Ok(_) => {} Err(e) => { @@ -126,7 +134,7 @@ pub(crate) fn start_multi( device: Arc, group_sync_sender: GroupSyncSender<(Vec, usize)>, up_counter: &mut SingleU64Adder, -) -> io::Result<()> { +) -> anyhow::Result<()> { let poll = Poll::new()?; let waker = Arc::new(Waker::new(poll.registry(), STOP)?); let _waker = waker.clone(); @@ -146,7 +154,7 @@ fn start_multi0( device: Arc, mut group_sync_sender: GroupSyncSender<(Vec, usize)>, up_counter: &mut SingleU64Adder, -) -> io::Result<()> { +) -> anyhow::Result<()> { let fd = device.as_tun_fd(); fd.set_nonblock()?; SourceFd(&fd.as_raw_fd()).register(poll.registry(), FD, Interest::READABLE)?; diff --git a/vnt/src/handle/tun_tap/windows.rs b/vnt/src/handle/tun_tap/windows.rs index 01d317a..0d03b1e 100644 --- a/vnt/src/handle/tun_tap/windows.rs +++ b/vnt/src/handle/tun_tap/windows.rs @@ -1,5 +1,7 @@ use crate::channel::context::ChannelContext; +use crate::channel::BUFFER_SIZE; use crate::cipher::Cipher; +use crate::compression::Compressor; use crate::external_route::ExternalRoute; use crate::handle::tun_tap::channel_group::GroupSyncSender; use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; @@ -8,7 +10,6 @@ use crate::ip_proxy::IpProxyMap; use crate::util::{SingleU64Adder, StopManager}; use crossbeam_utils::atomic::AtomicCell; use parking_lot::Mutex; -use std::io; use std::sync::Arc; use tun::device::IFace; use tun::Device; @@ -24,7 +25,8 @@ pub(crate) fn start_simple( server_cipher: Cipher, up_counter: &mut SingleU64Adder, device_list: Arc)>>, -) -> io::Result<()> { + compressor: Compressor, +) -> anyhow::Result<()> { let worker = { let device = device.clone(); stop_manager.add_listener("tun_device".into(), move || { @@ -44,6 +46,7 @@ pub(crate) fn start_simple( server_cipher, up_counter, device_list, + compressor, ) { log::error!("{:?}", e); } @@ -60,8 +63,10 @@ fn start_simple0( server_cipher: Cipher, up_counter: &mut SingleU64Adder, device_list: Arc)>>, -) -> io::Result<()> { - let mut buf = [0; 1024 * 16]; + compressor: Compressor, +) -> anyhow::Result<()> { + let mut buf = [0; BUFFER_SIZE]; + let mut extend = [0; BUFFER_SIZE]; loop { let len = device.read(&mut buf[12..])? + 12; //单线程的 @@ -72,6 +77,7 @@ fn start_simple0( context, &mut buf, len, + &mut extend, &device, current_device.load(), &ip_route, @@ -80,6 +86,7 @@ fn start_simple0( &client_cipher, &server_cipher, &device_list, + &compressor, ) { Ok(_) => {} Err(e) => { @@ -93,7 +100,7 @@ pub(crate) fn start_multi( device: Arc, group_sync_sender: GroupSyncSender<(Vec, usize)>, up_counter: &mut SingleU64Adder, -) -> io::Result<()> { +) -> anyhow::Result<()> { let worker = { let device = device.clone(); stop_manager.add_listener("tun_device_multi".into(), move || { @@ -112,7 +119,7 @@ fn start_multi0( device: Arc, mut group_sync_sender: GroupSyncSender<(Vec, usize)>, up_counter: &mut SingleU64Adder, -) -> io::Result<()> { +) -> anyhow::Result<()> { loop { let mut buf = vec![0; 1024 * 16]; let len = device.read(&mut buf[12..])? + 12; diff --git a/vnt/src/ip_proxy/mod.rs b/vnt/src/ip_proxy/mod.rs index 2387aef..0361889 100644 --- a/vnt/src/ip_proxy/mod.rs +++ b/vnt/src/ip_proxy/mod.rs @@ -10,11 +10,13 @@ use packet::ip::ipv4::packet::IpV4Packet; use crate::channel::context::ChannelContext; use crate::cipher::Cipher; use crate::handle::CurrentDeviceInfo; +#[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] use crate::ip_proxy::icmp_proxy::IcmpProxy; use crate::ip_proxy::tcp_proxy::TcpProxy; use crate::ip_proxy::udp_proxy::UdpProxy; use crate::util::StopManager; +#[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] pub mod icmp_proxy; pub mod tcp_proxy; pub mod udp_proxy; @@ -31,6 +33,7 @@ pub trait ProxyHandler { #[derive(Clone)] pub struct IpProxyMap { + #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] icmp_proxy: IcmpProxy, tcp_proxy: TcpProxy, udp_proxy: UdpProxy, @@ -65,15 +68,17 @@ pub fn init_proxy( } async fn init_proxy0( - context: ChannelContext, - current_device: Arc>, - client_cipher: Cipher, + _context: ChannelContext, + _current_device: Arc>, + _client_cipher: Cipher, ) -> anyhow::Result { - let icmp_proxy = IcmpProxy::new(context, current_device, client_cipher).await?; + #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] + let icmp_proxy = IcmpProxy::new(_context, _current_device, _client_cipher).await?; let tcp_proxy = TcpProxy::new().await?; let udp_proxy = UdpProxy::new().await?; Ok(IpProxyMap { + #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] icmp_proxy, tcp_proxy, udp_proxy, @@ -90,6 +95,7 @@ impl ProxyHandler for IpProxyMap { match ipv4.protocol() { ipv4::protocol::Protocol::Tcp => self.tcp_proxy.recv_handle(ipv4, source, destination), ipv4::protocol::Protocol::Udp => self.udp_proxy.recv_handle(ipv4, source, destination), + #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] ipv4::protocol::Protocol::Icmp => { self.icmp_proxy.recv_handle(ipv4, source, destination) } @@ -110,6 +116,7 @@ impl ProxyHandler for IpProxyMap { match ipv4.protocol() { ipv4::protocol::Protocol::Tcp => self.tcp_proxy.send_handle(ipv4), ipv4::protocol::Protocol::Udp => self.udp_proxy.send_handle(ipv4), + #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] ipv4::protocol::Protocol::Icmp => self.icmp_proxy.send_handle(ipv4), _ => Ok(()), } diff --git a/vnt/src/lib.rs b/vnt/src/lib.rs index a159e14..0a07456 100644 --- a/vnt/src/lib.rs +++ b/vnt/src/lib.rs @@ -16,3 +16,4 @@ pub mod tun_tap_device; pub mod util; pub use handle::callback::*; +pub mod compression; diff --git a/vnt/src/nat/stun.rs b/vnt/src/nat/stun.rs index ef4a636..05b9a58 100644 --- a/vnt/src/nat/stun.rs +++ b/vnt/src/nat/stun.rs @@ -71,11 +71,15 @@ pub fn stun_test_nat0(stun_servers: Vec) -> io::Result<(NatType, Vec io::Result> { diff --git a/vnt/src/protocol/body.rs b/vnt/src/protocol/body.rs index 384175d..bfa5fd2 100644 --- a/vnt/src/protocol/body.rs +++ b/vnt/src/protocol/body.rs @@ -253,6 +253,85 @@ impl + AsMut<[u8]>> AesCbcSecretBody { &mut self.buffer.as_mut()[..end] } } +/* ChaCah20加密数据体 + 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 + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + | 数据体 | + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + | finger(32) | + | finger(32) | + | finger(32) | + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + + 注:finger用于快速校验数据是否被修改,上层可使用token、协议头参与计算finger, + 确保服务端和客户端都能感知修改(服务端不能解密也能校验指纹) +*/ +pub struct ChaCah20SecretBody { + buffer: B, + exist_finger: bool, +} + +impl> ChaCah20SecretBody { + pub fn new(buffer: B, exist_finger: bool) -> io::Result> { + let len = buffer.as_ref().len(); + let min_len = if exist_finger { 12 } else { 0 }; + // 不能大于udp最大载荷长度 + if len < min_len || len > 65535 - 20 - 8 - 12 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "ChaCah20SecretBody length overflow", + )); + } + Ok(ChaCah20SecretBody { + buffer, + exist_finger, + }) + } + pub fn en_body(&self) -> &[u8] { + let mut end = self.buffer.as_ref().len(); + if self.exist_finger { + end -= 12; + } + &self.buffer.as_ref()[..end] + } + pub fn finger(&self) -> &[u8] { + if self.exist_finger { + let end = self.buffer.as_ref().len(); + &self.buffer.as_ref()[end - 12..end] + } else { + &[] + } + } +} + +impl + AsMut<[u8]>> ChaCah20SecretBody { + pub fn set_finger(&mut self, finger: &[u8]) -> io::Result<()> { + if self.exist_finger { + 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", + )) + } + } + pub fn en_body_mut(&mut self) -> &mut [u8] { + let mut end = self.buffer.as_ref().len(); + if self.exist_finger { + end -= 12; + } + &mut self.buffer.as_mut()[..end] + } +} /* rsa加密数据体 0 15 31 diff --git a/vnt/src/protocol/extension.rs b/vnt/src/protocol/extension.rs new file mode 100644 index 0000000..1e43e06 --- /dev/null +++ b/vnt/src/protocol/extension.rs @@ -0,0 +1,141 @@ +/* 扩展协议 + 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 + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + | 扩展数据(n) | + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + | 扩展数据(n) | type(8) | + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + 注:扩展数据的长度由type决定 +*/ + +use anyhow::anyhow; +use std::io; + +use crate::protocol::NetPacket; + +#[derive(Eq, PartialEq, Copy, Clone, Debug)] +pub enum ExtensionTailType { + Compression, + Unknown(u8), +} + +impl From for ExtensionTailType { + fn from(value: u8) -> Self { + if value == 0 { + ExtensionTailType::Compression + } else { + ExtensionTailType::Unknown(value) + } + } +} + +pub enum ExtensionTailPacket { + Compression(CompressionExtensionTail), + Unknown, +} + +impl + AsMut<[u8]>> NetPacket { + /// 分离尾部数据 + pub fn split_tail_packet(&mut self) -> anyhow::Result> { + if self.is_extension() { + let payload = self.payload(); + if let Some(v) = payload.last() { + return match ExtensionTailType::from(*v) { + ExtensionTailType::Compression => { + let data_len = self.data_len - 4; + self.set_data_len(data_len)?; + self.set_extension_flag(false); + Ok(ExtensionTailPacket::Compression( + CompressionExtensionTail::new( + &self.raw_buffer()[data_len..data_len + 4], + ), + )) + } + ExtensionTailType::Unknown(e) => Err(anyhow!("unknown extension {}", e)), + }; + } + } + Err(anyhow!("not extension")) + } + /// 追加压缩扩展 + pub fn append_compression_extension_tail( + &mut self, + ) -> io::Result> { + let len = self.data_len; + //增加数据长度 + self.set_data_len(self.data_len + 4)?; + self.set_extension_flag(true); + let mut tail = CompressionExtensionTail::new(&mut self.buffer_mut()[len..]); + tail.init(); + return Ok(tail); + } +} + +/* 扩展协议 + 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 + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + | algorithm(8) | | type(8) | + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ + 注:扩展数据的长度由type决定 +*/ +/// 压缩扩展 +pub struct CompressionExtensionTail { + buffer: B, +} + +impl> CompressionExtensionTail { + pub fn new(buffer: B) -> CompressionExtensionTail { + assert_eq!(buffer.as_ref().len(), 4); + CompressionExtensionTail { buffer } + } +} + +impl> CompressionExtensionTail { + pub fn algorithm(&self) -> CompressionAlgorithm { + self.buffer.as_ref()[0].into() + } +} + +impl + AsMut<[u8]>> CompressionExtensionTail { + pub fn init(&mut self) { + self.buffer.as_mut().fill(0); + } + pub fn set_algorithm(&mut self, algorithm: CompressionAlgorithm) { + self.buffer.as_mut()[0] = algorithm.into() + } +} + +#[derive(Eq, PartialEq, Copy, Clone, Debug)] +pub enum CompressionAlgorithm { + #[cfg(feature = "lz4_compress")] + Lz4, + #[cfg(feature = "zstd_compress")] + Zstd, + Unknown(u8), +} + +impl From for CompressionAlgorithm { + fn from(value: u8) -> Self { + match value { + #[cfg(feature = "lz4_compress")] + 1 => CompressionAlgorithm::Lz4, + #[cfg(feature = "zstd_compress")] + 2 => CompressionAlgorithm::Zstd, + v => CompressionAlgorithm::Unknown(v), + } + } +} + +impl From for u8 { + fn from(value: CompressionAlgorithm) -> Self { + match value { + #[cfg(feature = "lz4_compress")] + CompressionAlgorithm::Lz4 => 1, + #[cfg(feature = "zstd_compress")] + CompressionAlgorithm::Zstd => 2, + CompressionAlgorithm::Unknown(val) => val, + } + } +} diff --git a/vnt/src/protocol/mod.rs b/vnt/src/protocol/mod.rs index 19bedca..0a96cc3 100644 --- a/vnt/src/protocol/mod.rs +++ b/vnt/src/protocol/mod.rs @@ -21,6 +21,7 @@ pub const HEAD_LEN: usize = 12; pub mod body; pub mod control_packet; pub mod error_packet; +pub mod extension; pub mod ip_turn_packet; pub mod other_turn_packet; pub mod service_packet; @@ -101,6 +102,10 @@ pub struct NetPacket { } impl> NetPacket { + pub fn unchecked(buffer: B) -> Self { + let data_len = buffer.as_ref().len(); + Self { data_len, buffer } + } pub fn new(buffer: B) -> io::Result> { let data_len = buffer.as_ref().len(); Self::new0(data_len, buffer) @@ -158,6 +163,10 @@ impl> NetPacket { pub fn is_gateway(&self) -> bool { self.buffer.as_ref()[0] & 0x40 == 0x40 } + /// 扩展协议 + pub fn is_extension(&self) -> bool { + self.buffer.as_ref()[0] & 0x20 == 0x20 + } pub fn version(&self) -> Version { Version::from(self.buffer.as_ref()[0] & 0x0F) } @@ -190,6 +199,9 @@ impl> NetPacket { } impl + AsMut<[u8]>> NetPacket { + pub fn head_mut(&mut self) -> &mut [u8] { + &mut self.buffer.as_mut()[..12] + } pub fn buffer_mut(&mut self) -> &mut [u8] { &mut self.buffer.as_mut()[..self.data_len] } @@ -208,6 +220,13 @@ impl + AsMut<[u8]>> NetPacket { self.buffer.as_mut()[0] = self.buffer.as_ref()[0] & 0xBF }; } + pub fn set_extension_flag(&mut self, is_extension: bool) { + if is_extension { + self.buffer.as_mut()[0] = self.buffer.as_ref()[0] | 0x20 + } else { + self.buffer.as_mut()[0] = self.buffer.as_ref()[0] & 0xDF + }; + } pub fn set_default_version(&mut self) { let v: u8 = Version::V2.into(); self.buffer.as_mut()[0] = (self.buffer.as_ref()[0] & 0xF0) | (0x0F & v); @@ -264,6 +283,10 @@ impl + AsMut<[u8]>> NetPacket { self.data_len = data_len; Ok(()) } + pub fn set_payload_len(&mut self, payload_len: usize) -> io::Result<()> { + let data_len = HEAD_LEN + payload_len; + self.set_data_len(data_len) + } pub fn set_data_len_max(&mut self) { self.data_len = self.buffer.as_ref().len(); } diff --git a/vnt/src/tun_tap_device/tun_create_helper.rs b/vnt/src/tun_tap_device/tun_create_helper.rs index 2f83423..67b7665 100644 --- a/vnt/src/tun_tap_device/tun_create_helper.rs +++ b/vnt/src/tun_tap_device/tun_create_helper.rs @@ -8,6 +8,7 @@ use tun::Device; use crate::channel::context::ChannelContext; use crate::cipher::Cipher; +use crate::compression::Compressor; use crate::external_route::ExternalRoute; use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; #[cfg(feature = "ip_proxy")] @@ -78,6 +79,7 @@ struct TunDeviceHelperInner { parallel: usize, up_counter: SingleU64Adder, device_list: Arc)>>, + compressor: Compressor, } impl TunDeviceHelper { @@ -92,6 +94,7 @@ impl TunDeviceHelper { parallel: usize, up_counter: SingleU64Adder, device_list: Arc)>>, + compressor: Compressor, ) -> Self { Self { inner: Arc::new(AtomicCell::new(Some(TunDeviceHelperInner { @@ -106,6 +109,7 @@ impl TunDeviceHelper { parallel, up_counter, device_list, + compressor, }))), } } @@ -124,6 +128,7 @@ impl TunDeviceHelper { inner.parallel, inner.up_counter, inner.device_list, + inner.compressor, )?; Ok(()) } else { diff --git a/vnt/src/util/dns_query.rs b/vnt/src/util/dns_query.rs index 3104e2b..4b7300c 100644 --- a/vnt/src/util/dns_query.rs +++ b/vnt/src/util/dns_query.rs @@ -44,6 +44,7 @@ fn address_choose0(addrs: Vec) -> anyhow::Result { let v4: Vec = addrs.iter().filter(|v| v.is_ipv4()).copied().collect(); let v6: Vec = addrs.iter().filter(|v| v.is_ipv6()).copied().collect(); let check_addr = |addrs: &Vec| -> anyhow::Result { + let mut err = Vec::new(); if !addrs.is_empty() { let udp = if addrs[0].is_ipv6() { UdpSocket::bind("[::]:0")? @@ -51,12 +52,14 @@ fn address_choose0(addrs: Vec) -> anyhow::Result { UdpSocket::bind("0.0.0.0:0")? }; for addr in addrs { - if udp.connect(addr).is_ok() { + if let Err(e) = udp.connect(addr) { + err.push((*addr, e)); + } else { return Ok(*addr); } } } - Err(anyhow::anyhow!("Unable to connect to address {:?}", addrs)) + Err(anyhow::anyhow!("Unable to connect to address {:?}", err)) }; if v6.is_empty() { return check_addr(&v4); diff --git a/vnt/src/util/notify.rs b/vnt/src/util/notify.rs index 5e8ec68..defa700 100644 --- a/vnt/src/util/notify.rs +++ b/vnt/src/util/notify.rs @@ -1,9 +1,10 @@ use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::Arc; +use std::thread; use std::thread::Thread; use std::time::Duration; -use std::{io, thread}; +use anyhow::anyhow; use parking_lot::Mutex; #[derive(Clone)] @@ -20,7 +21,7 @@ impl StopManager { inner: Arc::new(StopManagerInner::new(f)), } } - pub fn add_listener(&self, name: String, f: F) -> io::Result + pub fn add_listener(&self, name: String, f: F) -> anyhow::Result where F: FnOnce() + Send + 'static, { @@ -61,23 +62,20 @@ impl StopManagerInner { stop_call: Mutex::new(Some(Box::new(f))), } } - fn add_listener(self: &Arc, name: String, f: F) -> io::Result + fn add_listener(self: &Arc, name: String, f: F) -> anyhow::Result where F: FnOnce() + Send + 'static, { if name.is_empty() { - return Err(io::Error::new(io::ErrorKind::Other, "name cannot be empty")); + return Err(anyhow!("name cannot be empty")); } let mut guard = self.listeners.lock(); if guard.0 { - return Err(io::Error::new(io::ErrorKind::Other, "stopped")); + return Err(anyhow!("stopped")); } for (n, _) in &guard.1 { if &name == n { - return Err(io::Error::new( - io::ErrorKind::Other, - format!("stop add_listener {:?} name already exists", name), - )); + return Err(anyhow!("stop add_listener {:?} name already exists", name)); } } guard.1.push((name.clone(), Box::new(f))); diff --git a/vnt/src/util/scheduler.rs b/vnt/src/util/scheduler.rs index 3680ea8..5d5485e 100644 --- a/vnt/src/util/scheduler.rs +++ b/vnt/src/util/scheduler.rs @@ -2,7 +2,6 @@ use crate::util::StopManager; use std::collections::BinaryHeap; use std::{ cmp::Ordering, - io, sync::mpsc::{sync_channel, Receiver, SyncSender}, time::{Duration, Instant}, }; @@ -36,7 +35,7 @@ pub struct Scheduler { sender: SyncSender, } impl Scheduler { - pub fn new(stop_manager: StopManager) -> io::Result { + pub fn new(stop_manager: StopManager) -> anyhow::Result { let (sender, receiver) = sync_channel::(32); let s = Self { sender }; let s_inner = s.clone();