Compare commits

..
188 Commits
Author SHA1 Message Date
lubeilin b26c4b97b2 上报状态 2024-03-24 21:50:47 +08:00
lubeilin 5569c67ba4 增加日志 2024-03-24 21:50:39 +08:00
lubeilin 69da6de1ed 离线时才检测服务器地址 2024-03-24 09:18:12 +08:00
lubeilin ae983f014b 降低地址探测频率 2024-03-23 12:46:20 +08:00
lubeilin 499e3bbfdf 去除重复逻辑 2024-03-21 23:08:48 +08:00
lubeilin 01f6890fd3 减少离线时发包 2024-03-21 21:25:24 +08:00
lubeilin 9badbe180c 不转发来源和目的相同的数据 2024-03-20 12:21:19 +08:00
lubeilin 30b1e71aa1 fmt 2024-03-19 23:54:33 +08:00
lubeilin b36cc352d5 兼容纯ipv4 2024-03-19 23:53:16 +08:00
lubeilin 12d888cefc 忽略跃点设置失败的异常 2024-03-18 21:35:22 +08:00
lubeilin ad1df41029 调整打洞 2024-03-17 15:41:25 +08:00
lubeilin 67498dfc82 增加序列号 2024-03-17 14:28:03 +08:00
lubeilin cf52fdde57 修改线程名称 2024-03-17 14:27:49 +08:00
lubeilin a20082d40b 已支持ipv6 2024-03-14 23:31:40 +08:00
lubeilin 8e29b020e0 Merge branch 'main' into mio 2024-03-14 12:06:47 +08:00
lubeilin 20132861e9 完善日志,修复已知问题 2024-03-14 12:05:35 +08:00
lubeilin a1e9b3c133 增加参数说明 2024-03-13 23:24:36 +08:00
lubeilin 8a4d21849e [mio] 调整代码 2024-03-13 21:23:32 +08:00
lubeilin d2c5aac178 [mio] 完善提示信息 2024-03-13 21:17:01 +08:00
lu 066e655c30 修复macos上已知问题 2024-03-12 23:52:53 +08:00
lubeilin 8627046deb [mio] 兼容mac 2024-03-12 21:30:06 +08:00
lubeilin 76b86bbd0f [mio] 调整打洞策略 2024-03-11 22:12:41 +08:00
lubeilin e78994ceb4 [mio] 修复重连问题 2024-03-11 21:59:13 +08:00
lubeilin c63813e590 [mio] 减少中继探测数据包 2024-03-10 21:07:06 +08:00
lubeilin 39becec2e5 [mio] 调整重连 2024-03-10 16:08:17 +08:00
lubeilin 44e379df82 [mio] 增加日志 2024-03-10 16:07:59 +08:00
lubeilin b2b58a23c1 [mio] 主通道改为异步 2024-03-10 14:11:06 +08:00
lubeilin 03529bc339 [mio] 返回详细路由 2024-03-10 14:10:49 +08:00
lubeilin 996a80e710 [mio] 调整打洞处理 2024-03-10 14:09:47 +08:00
lubeilin ac837077f6 [mio] 简化代码 2024-03-10 14:09:00 +08:00
lubeilin b25f330b0c [mio] 调整打洞处理 2024-03-10 14:08:30 +08:00
lubeilin c783a96e13 [mio] 打印移除路由日志 2024-03-10 14:07:04 +08:00
lubeilin 999a5b9117 [mio] 减少日志 2024-03-10 14:06:37 +08:00
lubeilin 6850681352 [mio] 增加缓冲区长度 2024-03-10 14:05:51 +08:00
lubeilin f2b21c9e4a [mio] 增加队列长度 2024-03-10 14:05:28 +08:00
lubeilin d27f065507 [mio] 调整通道使用策略 2024-03-07 21:48:54 +08:00
lubeilin 5e913568a3 优化通道选择 2024-03-07 13:11:45 +08:00
lubeilin 97b6db3485 输出增加颜色 2024-03-07 12:25:52 +08:00
lubeilin ae63ce80b5 [mio] 启动成功输出 2024-03-06 22:53:56 +08:00
lubeilin c0e930a3ef [mio] 简化代理 2024-03-06 22:46:17 +08:00
lubeilin f8ee8e242f [mio] fmt 2024-03-06 22:40:11 +08:00
lubeilin 20e3ae4d58 完善代理 2024-03-06 12:55:26 +08:00
lubeilin a910d7a673 修复打洞端口为0的问题 2024-03-06 12:55:07 +08:00
lubeilin a4b5b2a028 [mio] fmt 2024-03-05 21:36:21 +08:00
lubeilin a8a1346c98 [mio] 使用实际长度判断 2024-03-05 21:36:12 +08:00
lubeilin 734e613452 [mio] 完善错误处理 2024-03-05 21:35:49 +08:00
lubeilin e2003d8dea [mio] 优化代理 2024-03-05 21:34:51 +08:00
lubeilin f0a2edadfa [mio] 输出ip冲突错误 2024-03-05 21:33:25 +08:00
lubeilin 4c174c3aac [mio] 增加参数说明 2024-03-05 21:32:39 +08:00
lubeilin 147156d96d [mio] 增加模拟弱网参数 2024-03-05 21:32:20 +08:00
lubeilin 64d62272a6 [mio] 完善说明 2024-03-04 21:43:07 +08:00
lubeilin 91713f0fed [mio] fmt 2024-03-04 21:14:05 +08:00
lubeilin e7f2571469 [mio] device_id获取不到时默认值取空 2024-03-04 21:13:05 +08:00
lubeilin 48b41c4b22 [mio] 修复ip代理相关问题 2024-03-04 21:11:41 +08:00
lubeilin 042cc7c797 [mio] 增加开源许可 2024-03-04 21:10:49 +08:00
lubeilin 24c1edff8c [mio] 完善说明 2024-03-03 22:05:58 +08:00
lubeilin ec2a74d06f [mio] 完善说明 2024-03-03 21:35:49 +08:00
lubeilin 63316ee23e [mio] 优化显示 2024-03-03 19:23:43 +08:00
lubeilin f0139c68f7 [mio] 路由为空则回收数据 2024-03-03 19:23:34 +08:00
lubeilin 590721a4da [mio] 默认网卡名称改为vnt-tun/tap 2024-03-03 18:31:19 +08:00
lubeilin 521280537a [mio] 兼容padavan路由器找不到程序当前路径的问题 2024-03-03 09:45:52 +08:00
lubeilin 1eb2a0f5cf [mio] 减少读更新 2024-03-03 00:35:02 +08:00
lubeilin f0ac70c995 [mio] 修复relay参数问题 2024-03-02 23:47:56 +08:00
lubeilin 2570a9ead2 [mio] 修复ip tcp代理问题 2024-03-02 23:39:51 +08:00
lubeilin 54da8ac71f [mio] 判断停止状态 2024-03-02 23:14:16 +08:00
lubeilin bd73ed213e [mio] stop阻塞的问题 2024-03-02 23:13:59 +08:00
lubeilin 20fd060821 [mio] 改为阻塞连接 2024-03-02 15:50:04 +08:00
lubeilin e8df40925a [mio] 增加连接超时时间 2024-03-02 15:49:23 +08:00
lubeilin 2c0a5d459a [mio] 修改连接超时时间 2024-03-02 15:49:06 +08:00
lubeilin b67814f125 [mio] 去除临时输出 2024-03-02 14:15:18 +08:00
lubeilin e3f328cc9b [mio] 修复指定通道类型无效的问题 2024-03-02 14:03:19 +08:00
lubeilin b6aba36510 [mio] fmt 2024-03-02 14:02:20 +08:00
lubeilin 22720a8f8f [mio] fmt 2024-03-02 14:02:01 +08:00
lubeilin cc717a4060 [mio] 去除获取目录失败的unwrap 2024-03-02 11:57:00 +08:00
lubeilin 5264455eba [mio] 回收计数器 2024-03-02 11:50:18 +08:00
lubeilin b0d6b1f884 [mio] 简化流量统计,兼容32位系统 2024-03-01 23:40:14 +08:00
lubeilin 7cd8647c0c [mio] fmt 2024-03-01 23:34:37 +08:00
lubeilin 19187aff47 [mio] 设置广播掩码 2024-03-01 23:34:14 +08:00
lubeilin 23b536f771 [mio] 修复停止命令失效问题 2024-03-01 23:33:47 +08:00
lubeilin 119f719a9f [mio] 支持仅使用p2p模式 2024-03-01 21:48:07 +08:00
lubeilin 9673eaab05 [mio] 添加组播和广播路由 2024-03-01 21:47:49 +08:00
lubeilin 390b8dd242 [mio] 确认离线才调用disconnect 2024-03-01 21:47:14 +08:00
lubeilin 1b1663c87e [mio] 简化校验 2024-03-01 21:46:48 +08:00
lubeilin 7bc4203deb [mio] 网卡名称使用String 2024-03-01 21:44:27 +08:00
lubeilin c798c36914 [mio] 删除多余文件 2024-03-01 21:36:46 +08:00
lubeilin 14e7b518b5 [mio] 更改版本为 1.2.9 2024-02-29 22:40:01 +08:00
lubeilin 88eb3f6514 [mio] fmt 2024-02-29 22:39:35 +08:00
lubeilin d78e48211d [mio] 处理server_encrypt代码逻辑 2024-02-29 22:39:20 +08:00
lubeilin 55a510aa34 [mio] 去除多余代码 2024-02-29 22:31:08 +08:00
lubeilin 8d7eb5ba1b [mio] 适配新版vnt,增加jni java端代码 2024-02-29 22:29:39 +08:00
lubeilin ceadadee68 [mio] 适配新版vnt 2024-02-29 22:29:14 +08:00
lubeilin e3e2a262e4 [mio] 增加依赖 2024-02-29 22:28:52 +08:00
lubeilin 5850a59ec9 [mio] 去除多余代码 2024-02-29 22:28:20 +08:00
lubeilin 318c267c34 [mio] 简化vnt创建 2024-02-29 22:27:47 +08:00
lubeilin 98bd713a91 [mio] 整理数据处理逻辑 2024-02-29 22:27:15 +08:00
lubeilin fe499d0476 [mio] 使用mio改写ip代理 2024-02-29 22:26:53 +08:00
lubeilin ff665d7dcf [mio] 更新protobuf 2024-02-29 22:25:30 +08:00
lubeilin 573537fbfb [mio] 使用io结果简化异常处理 2024-02-29 22:25:11 +08:00
lubeilin 9e1bd28e2d [mio] 创建tun/tap设备 2024-02-29 22:22:25 +08:00
lubeilin 27a6c686a7 [mio] 合并各平台的tun/tap处理 2024-02-29 22:20:30 +08:00
lubeilin 2f6d743931 [mio] 加入端口组,去除tokio 2024-02-29 22:19:41 +08:00
lubeilin ec3aa01b9c [mio] 增加统计、监听器、定时器 2024-02-29 22:17:27 +08:00
lubeilin 02b988e898 [mio] 去除Arc 2024-02-29 22:15:06 +08:00
lubeilin b57cd2ed45 [mio] 简化feature 2024-02-29 22:12:40 +08:00
lubeilin af5d284470 [mio] 支持多通道传输,使用mio代替tokio 2024-02-29 22:06:20 +08:00
lubeilin e7f1165287 [mio] 增加打洞端口组 2024-02-29 22:05:17 +08:00
lubeilin e64e17267d 增加tcp通道处理逻辑 2024-01-07 21:09:26 +08:00
lubeilin 7098111ad1 增加tcp通道,去除并行逻辑 2024-01-06 12:06:21 +08:00
lubeilin 57b903ea29 命令默认使用39271端口 2024-01-06 11:30:05 +08:00
lubeilin d06082fc7a 默认输出到当前路径 2024-01-06 11:26:09 +08:00
lubeilin 77a78daac1 版本号使用环境变量 2024-01-06 11:22:14 +08:00
lubeilin d72fab04a2 mips平台使用1.71.1版本rust 2024-01-06 11:21:59 +08:00
lubeilin 76eaead0b9 修改对称网络处理逻辑 2023-12-31 13:42:32 +08:00
lubeilin 858ca9bbe7 去除udp通道的arc包装 2023-12-30 22:15:48 +08:00
lubeilin cbc4a7378c 优化tcp通道 2023-12-30 21:39:53 +08:00
lubeilin ee34f525e6 nat地址判断 2023-12-30 13:07:54 +08:00
lubeilin 698e2531e8 增加参数校验 2023-12-30 12:00:47 +08:00
lubeilin 16f833ec72 优化延迟优先参数 2023-12-28 23:18:57 +08:00
lubeilin 94d6caef7e 修复连接tcp地址的问题 2023-12-26 22:27:16 +08:00
lubeilin 364012f9dd fmt 2023-12-26 21:47:01 +08:00
lubeilin cf4b1f418f 忽略地址校验 2023-12-26 21:46:41 +08:00
lubeilin c6465977ef 将ipv4转换成ipv6 2023-12-26 21:46:23 +08:00
lubeilin c577e6381f 将ipv4转换成ipv6 2023-12-25 23:06:47 +08:00
lubeilin 134e31f563 fmt 2023-12-24 13:42:41 +08:00
lubeilin 0580b89f48 升级版本号 2023-12-24 12:22:34 +08:00
lubeilin 37080af275 支持ipv6服务端 2023-12-24 12:00:44 +08:00
lubeilin 26d68ac059 版本改为1.2.7 2023-10-31 20:49:25 +08:00
lubeilin 6292c1c381 去除溢出检查 2023-10-31 20:49:16 +08:00
lubeilin b0c3f25a29 使用读写锁简化操作,增加延迟优先选项 2023-10-31 20:48:38 +08:00
lubeilin d2e09d3da5 增加配置文件的参数说明 2023-10-11 20:46:39 +08:00
lbl8603 293c5b90a4 Merge pull request #22 from taotieren/contrib
Update README.md
2023-10-11 09:19:37 +08:00
taotieren e2323361f9 Update README.md 2023-10-10 23:23:25 +08:00
lbl8603 d05cff99ee Merge pull request #21 from taotieren/aur
Update README.md
2023-10-10 22:52:30 +08:00
taotieren e9b1b2ef3b Update README.md 2023-10-10 22:31:38 +08:00
lbl8603 a8ea2c14fc Merge pull request #20 from taotieren/aur
Add AUR vnt-git
2023-10-10 22:22:10 +08:00
taotieren 0688cb4515 Add AUR vnt-git 2023-10-10 22:18:37 +08:00
lubeilin c3261d7a57 修改版本为1.2.6 2023-10-09 21:20:10 +08:00
lubeilin ef8d13f61b 编译features增加ring 2023-10-09 20:56:58 +08:00
lubeilin fc104d5dee 增加参数关闭ip代理 2023-10-09 20:56:04 +08:00
lubeilin 9ad0525216 减少复制 2023-10-08 20:52:12 +08:00
lubeilin 3ea1250b53 完善编译和参数说明 2023-10-08 20:03:14 +08:00
lubeilin 9df34207f7 解决不同features编译时告警的问题 2023-10-08 17:46:31 +08:00
lubeilin dacce892ff 可选ip代理,关闭后可使用外部命令来进行ip转发 2023-10-08 17:16:02 +08:00
lubeilin a9495f1d30 升级tap相关依赖 2023-10-06 22:44:02 +08:00
lubeilin 0892c6eee9 修复进程异常退出的问题 2023-10-06 22:25:30 +08:00
lubeilin f746eadcf1 fmt 2023-10-06 22:22:35 +08:00
lubeilin da2371541c 修改依赖版本 2023-10-06 22:22:22 +08:00
lubeilin cca91d4331 tcp改为使用同步方法 2023-10-06 22:21:29 +08:00
lubeilin 11aa3b1d2c 修改features说明 2023-09-28 10:35:50 +08:00
lubeilin d0b0c61a03 Merge remote-tracking branch 'origin/main' 2023-09-28 10:27:10 +08:00
lubeilin b0501db4e5 修改features说明 2023-09-28 10:26:38 +08:00
lubeilin 465a5c75ae 修改默认features 2023-09-28 10:05:52 +08:00
lbl8603 978460aca8 更新 README.md 2023-09-27 23:51:52 +08:00
lubeilin 4b86752c79 去除不必要的features 2023-09-27 19:59:50 +08:00
lubeilin bd61cad7b6 加密算法可选 2023-09-27 19:51:23 +08:00
lubeilin 40cbd2e26b 更新参数说明 2023-09-26 21:56:04 +08:00
lubeilin 2cab580b4e 更新版本 2023-09-26 21:53:55 +08:00
lubeilin 3d243fb01d 支持读取配置文件和自定义端口 2023-09-26 21:53:41 +08:00
lubeilin ca4e8d14f0 支持sm4-cbc加密 2023-09-26 21:53:07 +08:00
lubeilin 9c098c55c9 调整长度判断 2023-09-26 21:51:56 +08:00
lubeilin 707b07b8d3 ipv6改为完整地址 2023-09-24 15:42:19 +08:00
lubeilin de5a6971f0 更新读取时间不需要再插入 2023-09-24 14:58:03 +08:00
lubeilin 056036c4d2 减少注册和探测nat的频率 2023-09-24 14:53:37 +08:00
lubeilin 1cfb188845 去除多余依赖 2023-09-24 13:45:47 +08:00
lubeilin 73a2c31854 连接通过关闭同时关闭tap 2023-09-24 13:42:09 +08:00
lubeilin 92eea536f8 删除多余依赖 2023-09-24 12:57:49 +08:00
lubeilin 00936a923e 避免直接关闭网卡 2023-09-24 12:55:42 +08:00
lubeilin ba69ba78af 增加线程名称 2023-09-23 23:09:37 +08:00
lubeilin 17b206bace 1.2.4.3 2023-09-23 21:38:31 +08:00
lubeilin 58d5a4f5da 增加wintun日志 2023-09-23 21:38:20 +08:00
lubeilin 2438d14175 避免短时间重复上传服务端密钥 2023-09-23 21:33:00 +08:00
lubeilin 99f8526799 去除tap广播路由 2023-09-22 22:49:05 +08:00
lubeilin 56fcbd64ed 增加日志 2023-09-22 22:43:23 +08:00
lubeilin 3766b2b7c1 修改命令超时时间 2023-09-22 22:43:02 +08:00
lubeilin baf0698fe4 去除广播路由 2023-09-22 22:17:13 +08:00
lubeilin 301938b9fc 增加小版本 2023-09-22 18:19:13 +08:00
lubeilin 16a37c713a 增加日志 2023-09-22 18:18:29 +08:00
lubeilin d412a769dd fmt 2023-09-22 18:18:10 +08:00
lubeilin c8eecc87fd 调整心跳间隔,服务端和客户端心跳分离 2023-09-22 18:17:12 +08:00
lubeilin 6a11db70c8 调整代理超时时间 2023-09-22 18:16:06 +08:00
lubeilin d7c121a756 commit:
1.去除缓冲池
2.数据处理改为同步方法
3.fmt
2023-09-20 19:54:49 +08:00
lubeilin 3429ee8bd6 增加提示 2023-09-20 16:03:25 +08:00
lubeilin 4422f9f8b7 Merge remote-tracking branch 'origin/main'
# Conflicts:
#	vnt/src/ip_proxy/tcp_proxy.rs
2023-09-20 15:39:54 +08:00
lubeilin 57ed454c93 修复内网ip断线问题 2023-09-20 11:11:31 +08:00
lubeilin 236205c0f3 修复内网ip断线问题 2023-09-19 18:25:14 +08:00
lubeilin 99b4bf0041 Merge remote-tracking branch 'origin/main' 2023-09-18 21:51:17 +08:00
lubeilin 9495e39700 优化nat校验 2023-09-18 21:51:08 +08:00
lbl8603 75e244e3a8 Update README.md 2023-09-18 11:17:21 +08:00
161 changed files with 11869 additions and 9577 deletions
+9 -3
View File
@@ -47,13 +47,13 @@ jobs:
FEATURES: ring-cipher,openssl-vendored
- TARGET: aarch64-unknown-linux-musl # tested on aws t4g.nano in alpine container
OS: ubuntu-latest
FEATURES: default
FEATURES: ring-cipher,openssl-vendored
- TARGET: armv7-unknown-linux-musleabihf # raspberry pi 2-3-4, not tested
OS: ubuntu-latest
FEATURES: openssl-vendored
- TARGET: arm-unknown-linux-musleabihf # raspberry pi 0-1, not tested
OS: ubuntu-latest
FEATURES: openssl-vendored
FEATURES: ring-cipher,openssl-vendored
- TARGET: x86_64-apple-darwin # tested on a mac, is not properly signed so there are security warnings
OS: macos-latest
FEATURES: ring-cipher,openssl-vendored
@@ -68,7 +68,7 @@ jobs:
FEATURES: ring-cipher,openssl-vendored
- TARGET: mipsel-unknown-linux-musl # openwrt
OS: ubuntu-latest
FEATURES: openssl-vendored
FEATURES: openssl-vendored,ring-cipher
- TARGET: mips-unknown-linux-musl # openwrt
OS: ubuntu-latest
FEATURES: openssl-vendored
@@ -120,6 +120,12 @@ jobs:
;;
esac
if [[ $TARGET =~ ^mips.*$ ]]; then
# mips平台使用1.71.1版本
rustup install 1.71.1
rustup default 1.71.1
fi
if [ -n "$MUSL_URI" ]; then
mkdir -p ./musl_gcc
wget -c https://musl.cc/$MUSL_URI.tgz -P ./musl_gcc/
-1
View File
@@ -6,7 +6,6 @@ opt-level = 'z'
debug = 0
debug-assertions = false
strip= "debuginfo"
overflow-checks = true
lto = true
panic = 'abort'
incremental = false
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+123 -8
View File
@@ -3,7 +3,9 @@
A virtual network tool (VPN)
将不同网络下的多个设备虚拟到一个局域网下
### vnt-cli参数详解 [参数说明](https://github.com/lbl8603/vnt/blob/main/vnt-cli/README.md)
### 快速使用:
1. 指定一个token,在多台设备上运行该程序,例如:
@@ -61,14 +63,105 @@ A virtual network tool (VPN)
前提条件:安装rust编译环境([install rust](https://www.rust-lang.org/zh-CN/tools/install))
```
到项目根目录下执行 cargo build -p vnt-cli
也可按需编译,将得到更小的二进制文件,使用--no-default-features排除默认features
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代理 | 是 |
### ip转发/代理
如果编译时去除了内置的ip代理(或使用--no-proxy关闭了代理),则可以使用网卡NAT转发来实现点对网,
一般来说使用网卡NAT转发会比内置的ip代理性能更好
<details> <summary>NAT配置可参考如下示例,点击展开</summary>
### 在出口一端做如下配置
注意原有的-i(入口)和-o(出口)的参数不能少
### windows
参考 https://learn.microsoft.com/zh-cn/virtualization/hyper-v-on-windows/user-guide/setup-nat-network
```shell
#设置nat,名字可以自己取,网段是vnt的网段
New-NetNat -Name vntnat -InternalIPInterfaceAddressPrefix 10.26.0.0/24
#查看设置
Get-NetNat
```
### linux
```shell
# 开启ip转发
sudo sysctl -w net.ipv4.ip_forward=1
# 开启nat转发 表示来源10.26.0.0/24的数据通过nat映射后再从vnt-tun以外的其他网卡发出去
sudo iptables -t nat -A POSTROUTING ! -o vnt-tun -s 10.26.0.0/24 -j MASQUERADE
# 或者这样 表示来源10.26.0.0/24的数据通过nat映射后再从eth0网卡发出去
sudo iptables -t nat -A POSTROUTING -o eth0 -s 10.26.0.0/24 -j MASQUERADE
# 查看设置
iptables -vnL -t nat
```
### Arch Linux
[![Packaging status](https://repology.org/badge/vertical-allrepos/vnt.svg)](https://repology.org/project/vnt/versions)
- 通过 AUR 安装 [vnt-git](https://aur.archlinux.org/packages/vnt-git)
```bash
yay -Syu vnt
```
- 通过 `systemd` 设置开机自启及配置
```bash
sudo systemctl enable --now vnt-cli@
sudo systemctl status vnt-cli@
```
- 启用内置 `IPv4` 转发规则
```bash
sudo sysctl --system
```
- 通过内置防火墙文件配置防火墙转发规则
```bash
sudo cat /etc/vnt/iptables-vnt.rules >> /etc/iptables/iptables.rules
sudo iptables-restore iptables.rules
```
### macos
```shell
# 开启ip转发
sudo sysctl -w net.ipv4.ip_forward=1
# 配置NAT转发规则
# 在/etc/pf.conf文件中添加以下规则,en0是出口网卡,10.26.0.0/24是来源网段
nat on en0 from 10.26.0.0/24 to any -> (en0)
# 加载规则
sudo pfctl -f /etc/pf.conf -e
```
</details>
### 支持平台
- Mac
- Linux
- Arch Linux `yay -Syu vnt`
- Windows
- 使用tun网卡 依赖wintun.dll([win-tun](https://www.wintun.net/))(将dll放到同目录下,建议使用版本0.14.1)
- 默认使用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)
@@ -86,9 +179,11 @@ A virtual network tool (VPN)
- p2p组播/广播
- 客户端数据加密
- 服务端数据加密
### 结构
<details> <summary>展开</summary>
<pre>
0 15 31
@@ -118,11 +213,11 @@ A virtual network tool (VPN)
### Todo
- 桌面UI(测试中)
- 支持Ipv6(1.2.2已支持客户端之间的ipv6,待支持客户端和服务端之间的ipv6通信)
### 常见问题
<details> <summary>展开</summary>
#### 问题1: 设置网络地址失败
##### 可能原因:
@@ -138,26 +233,46 @@ vnt默认使用10.26.0.0/24网段,和本地网络适配器的ip冲突
#### 问题2: windows系统上wintun.dll加载失败
##### 可能原因:
没有下载wintun.dll 或者使用的wintun.dll有问题
##### 解决方法:
1. 下载最新版的wintun.dll [下载链接](https://www.wintun.net/builds/wintun-0.14.1.zip)
2. 解压后找到对应架构的目录,通常是amd64
3. 将对应的wintun.dll放到和vnt-cli同目录下(或者放到C盘Windows目录下)
4. 再次启动vnt-cli
#### 问题3: 丢包严重,或是不能正常组网通信
##### 可能原因:
某些宽带下(比如广电宽带)UDP丢包严重
##### 解决方法:
1. 使用TCP模式中继转发(vnt-cli增加--tcp参数)
2. 如果p2p后效果很差,可以选择禁用p2p(vnt-cli增加--relay参数)
2. 如果p2p后效果很差,可以选择禁用p2p(vnt-cli增加--use-channel relay 参数)
#### 问题4:重启后虚拟IP发生变化,或指定了IP不能启动
##### 可能原因:
设备重启后程序自动获取的id值改变,导致注册时重新分配了新的IP,或是IP冲突
##### 解决方法:
1. 命令行启动增加-d参数(使用配置文件启动则在配置文件中增加device_id参数),要保证每个设备的值都不一样,取值可以任意64位以内字符串
</details>
### 交流群
QQ:1034868233
QQ: 1034868233
### 其他
可使用社区小伙伴搭建的中继服务器
1. -s vnt.8443.eu.org:29871
### 参与贡献
<a href="https://github.com/lbl8603/vnt/graphs/contributors">
<img src="https://contrib.rocks/image?repo=lbl8603/vnt" />
</a>
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "common"
version = "1.2.3"
version = "1.2.9"
edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
+13 -9
View File
@@ -1,20 +1,19 @@
[package]
name = "vnt-cli"
version = "1.2.3"
version = "1.2.9"
edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
[dependencies]
vnt = { path = "../vnt", package = "vnt", optional = true }
vnt = { path = "../vnt", package = "vnt",default-features = false }
common = { path = "../common" }
tokio = { version = "1.32.0", features = ["full"] }
getopts = "0.2.21"
console = "0.15.2"
os_info = "3.7.0"
dirs = "4.0.0"
serde = "1.0"
serde_json = "1.0.94"
#serde_json = "1.0.94"
serde_yaml = "0.9.32"
log = "0.4.17"
log4rs = "1.2.0"
[dependencies.uuid]
@@ -30,11 +29,16 @@ sudo = "0.6.0"
winapi = { version = "0.3.9", features = ["handleapi", "processthreadsapi", "winnt", "securitybaseapi", "impl-default"] }
[features]
default = ["vnt"]
default = ["server_encrypt","aes_gcm","aes_cbc","aes_ecb","sm4_cbc","ip_proxy"]
openssl = ["vnt/openssl"]
openssl-vendored = ["vnt/openssl-vendored"]
ring-cipher = ["vnt/ring-cipher"]
aes_cbc=["vnt/aes_cbc"]
aes_ecb=["vnt/aes_ecb"]
sm4_cbc=["vnt/sm4_cbc"]
aes_gcm=["vnt/aes_gcm"]
server_encrypt=["vnt/server_encrypt"]
ip_proxy=["vnt/ip_proxy"]
[build-dependencies]
embed-manifest = "1.4.0"
embed-manifest = "1.4.0"
rand = "0.9.0-alpha.0"
+77 -12
View File
@@ -15,6 +15,8 @@
使用stun服务探测客户端NAT类型,不同类型有不同的打洞策略
### -a
加了此参数表示使用tap网卡,默认使用tun网卡,tun网卡效率更高
### --nic `<tun0>`
指定虚拟网卡名称,默认tun模式使用vnt-tuntap模式使用vnt-tap
### -i `<in-ip>`、-o `<out-ip>`
配置点对网(IP代理)时使用,例如A(虚拟ip:10.26.0.2)通过B(虚拟ip:10.26.0.3,本地出口ip:192.168.0.10)访问C(目标网段192.168.0.0/24)
@@ -29,19 +31,17 @@
提升通信安全性,使用该密码生成的密钥对客户端数据进行加密,并且服务端无法解密(包括中继数据)。使用相同密码的客户端才能通信
| 密码位数 | 加密算法 |
|---------|-------|
| 小于8 | AES128-GCM
| 大于等于8 | AES256-GCM |
| 密码位数 | 加密算法 |
|-------|------------|
| 小于8 | AES128-GCM |
| 大于等于8 | AES256-GCM |
### -W
开启和服务端通信的数据加密,采用rsa+aes256gcm加密客户端和服务端之间通信的数据,可以避免token泄漏、中间人攻击
### -m
模拟组播,高频使用组播通信时,可以尝试开启此参数,默认情况下会把组播当作广播发给所有节点
默认情况(组播当广播发送):稳定性好,使用组播频率低时更省流量
模拟组播:高频使用组播时防止广播泛洪,客户端和中继服务器会维护组播成员等信息,注意使用此选项时,虚拟网内所有成员都需要开启此选项
注意:
1. -w `<password>`是用于客户端-客户端之间的加密,password不会传递到服务端,只添加这个参数不会加密客户端-服务端通信的数据
2. -W 用于开启客户端-服务端之间的加密
### -u `<mtu>`
@@ -54,7 +54,8 @@
### --par `<parallel>`
任务并行度(必须为正整数),默认值为1,该值表示处理网卡读写的任务数,组网设备数较多、处理延迟较大时可适当调大此值
### --model `<model>`
加密模式,可选值 aes_gcm/aes_cbc/aes_ecb,默认使用aes_gcm,通常情况aes_gcm安全性高、aes_ecb性能更好
加密模式,可选值 aes_gcm/aes_cbc/aes_ecb/sm4_cbc,默认使用aes_gcm,通常情况aes_gcm安全性高、aes_ecb性能更好,但是在低性能设备上sm4_cbc也许速度会更快;
| 密码位数 | model | 加密算法 |
|-------|---------|------------|
@@ -64,11 +65,75 @@
| `>=`8 | aes_cbc | AES256-CBC |
| 1~8位 | aes_ecb | AES128-ECB |
| `>=`8 | aes_ecb | AES256-ECB |
| `>0` | sm4_cbc | SM4-CBC |
### --finger
开启数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启,开启会损耗一部分性能
### --relay
禁用p2p,在网络环境很差时,只使用服务器中转效果可能更好(可以配合--tcp参数一起使用)
注意:默认情况下服务端不会对中转的数据做校验,如果要对中转的数据做校验,则需要客户端、服务端都开启此参数
### --punch `<punch>`
取值ipv4/ipv6,选择只使用ipv4打洞或者只使用ipv6打洞,默认两则都会使用
### --ports `<port1,port2>`
指定本地监听的端口组,多个端口使用逗号分隔,多个端口可以分摊流量,增加并发,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)
### -f `<conf>`
指定配置文件
配置文件采用yaml格式,可参考:
```yaml
# 全部参数
tap: false #是否使用tap
token: xxx #组网token
device_id: xxx #当前设备id
name: windows 11 #当前设备名称
server_address: ip:port #注册和中继服务器
stun_server: #stun服务器
- stun1.l.google.com:19302
- stun2.l.google.com:19302
in_ips: #代理ip入站
- 192.168.1.0/24,10.26.0.3
out_ips: #代理ip出站
- 0.0.0.0/0
password: xxx #密码
mtu: 1420 #mtu
tcp: false #tcp模式
ip: 10.26.0.2 #指定虚拟ip
use_channel: relay #relay:仅中继模式.p2p:仅直连模式
server_encrypt: true #服务端加密
parallel: 1 #任务并行度
cipher_model: aes_gcm #客户端加密算法
finger: false #关闭数据指纹
punch_model: ipv4 #打洞模式
ports:
- 0 #使用随机端口,tcp监听此端口
- 0
cmd: false #关闭控制台输入
no_proxy: false #是否关闭内置代理,true为关闭
first_latency: false #是否优先低延迟通道,默认为false,表示优先使用p2p通道
device_name: vnt-tun #网卡名称
packet_loss: 0 #指定丢包率 取值0~1之间的数 用于模拟弱网
packet_delay: 0 #指定延迟 单位毫秒 用于模拟弱网
```
或者需要哪个配置就加哪个,当然token是必须的
```yaml
# 部分参数
token: xxx #组网token
```
### --use-channel `<relay/p2p>`
- relay:仅中继模式,会禁止打洞/p2p直连,只使用服务器转发
- p2p:仅直连模式,会禁止网络数据从服务器/客户端转发,只会使用服务器转发控制包
### --packet-loss `<0>`
模拟丢包,取值0~1之间的小数,程序会按设定的概率主动丢包。在模拟弱网环境会有帮助。
### --list
在后台运行时,查看其他设备列表
### --all
+14 -7
View File
@@ -1,10 +1,17 @@
// use embed_manifest::{embed_manifest, new_manifest};
// use embed_manifest::manifest::ExecutionLevel;
use rand::Rng;
use std::fs::File;
use std::io::Write;
fn main() {
////强制用管理员运行貌似体验更差了
// if std::env::var_os("CARGO_CFG_WINDOWS").is_some() {
// embed_manifest(new_manifest("vnt")
// .requested_execution_level(ExecutionLevel::RequireAdministrator)).expect("unable to embed manifest file");
// }
// 生成随机序列号
let serial_number = format!(
"{}-{}-{}",
rand::thread_rng().gen_range(100..1000),
rand::thread_rng().gen_range(100..1000),
rand::thread_rng().gen_range(100..1000)
);
let generated_code = format!(r#"pub const SERIAL_NUMBER: &str = "{}";"#, serial_number);
let dest_path = "src/generated_serial_number.rs";
let mut file = File::create(&dest_path).unwrap();
file.write_all(generated_code.as_bytes()).unwrap();
}
+52
View File
@@ -0,0 +1,52 @@
use std::process;
use console::style;
use vnt::handle::callback::{ConnectInfo, ErrorType};
use vnt::{DeviceInfo, ErrorInfo, HandshakeInfo, RegisterInfo, VntCallback};
#[derive(Clone)]
pub struct VntHandler {}
impl VntCallback for VntHandler {
fn success(&self) {
println!(" {} ", style("====== Connect Successfully ======").green())
}
fn create_tun(&self, info: DeviceInfo) {
println!("create_tun {}", info)
}
fn connect(&self, info: ConnectInfo) {
println!("connect {}", info)
}
fn handshake(&self, info: HandshakeInfo) -> bool {
println!("handshake {}", info);
true
}
fn register(&self, info: RegisterInfo) -> bool {
println!("register {}", style(info).green());
true
}
fn error(&self, info: ErrorInfo) {
log::error!("error {:?}", info);
println!("{}", style(format!("error {}", info)).red());
match info.code {
ErrorType::TokenError
| ErrorType::AddressExhausted
| ErrorType::IpAlreadyExists
| ErrorType::InvalidIp
| ErrorType::LocalIpExists => {
self.stop();
}
_ => {}
}
}
fn stop(&self) {
println!("stopped");
process::exit(0)
}
}
+36 -44
View File
@@ -1,3 +1,4 @@
use serde::Deserialize;
use std::io;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4, UdpSocket};
use std::str::FromStr;
@@ -6,68 +7,59 @@ use std::time::Duration;
use crate::command::entity::{DeviceItem, Info, RouteItem};
pub struct CommandClient {
buf: [u8; 10240],
udp: UdpSocket,
}
impl CommandClient {
pub fn new() -> io::Result<Self> {
let path_buf = crate::app_home()?.join("command-port");
if !path_buf.exists() {
return Err(io::Error::new(io::ErrorKind::Other, "not started"));
}
let port = std::fs::read_to_string(path_buf)?;
let port = match u16::from_str(&port) {
Ok(port) => port,
Err(_) => {
return Err(io::Error::new(
io::ErrorKind::Other,
"'command-port' file error",
));
}
};
let port = read_command_port().unwrap_or_else(|e| {
log::warn!("read_command_port:{:?}", e);
39271
});
let udp = UdpSocket::bind("127.0.0.1:0")?;
udp.set_read_timeout(Some(Duration::from_secs(2)))?;
udp.set_read_timeout(Some(Duration::from_secs(5)))?;
udp.connect(SocketAddr::V4(SocketAddrV4::new(
Ipv4Addr::new(127, 0, 0, 1),
port,
)))?;
Ok(Self { udp })
Ok(Self {
udp,
buf: [0; 10240],
})
}
}
fn read_command_port() -> io::Result<u16> {
let path_buf = crate::app_home()?.join("command-port");
let port = std::fs::read_to_string(path_buf)?;
match u16::from_str(&port) {
Ok(port) => Ok(port),
Err(_) => {
return Err(io::Error::new(
io::ErrorKind::Other,
"'command-port' file error",
));
}
}
}
impl CommandClient {
pub fn list(&self) -> io::Result<Vec<DeviceItem>> {
self.udp.send(b"list")?;
let mut buf = [0; 10240];
let len = self.udp.recv(&mut buf)?;
match serde_json::from_slice::<Vec<DeviceItem>>(&buf[..len]) {
Ok(val) => Ok(val),
Err(e) => {
log::error!("{:?}", e);
Err(io::Error::new(io::ErrorKind::Other, "data error"))
}
}
pub fn list(&mut self) -> io::Result<Vec<DeviceItem>> {
self.send_cmd(b"list")
}
pub fn route(&self) -> io::Result<Vec<RouteItem>> {
self.udp.send(b"route")?;
let mut buf = [0; 10240];
let len = self.udp.recv(&mut buf)?;
match serde_json::from_slice::<Vec<RouteItem>>(&buf[..len]) {
Ok(val) => Ok(val),
Err(e) => {
log::error!("{:?}", e);
Err(io::Error::new(io::ErrorKind::Other, "data error"))
}
}
pub fn route(&mut self) -> io::Result<Vec<RouteItem>> {
self.send_cmd(b"route")
}
pub fn info(&self) -> io::Result<Info> {
self.udp.send(b"info")?;
let mut buf = [0; 10240];
let len = self.udp.recv(&mut buf)?;
match serde_json::from_slice::<Info>(&buf[..len]) {
pub fn info(&mut self) -> io::Result<Info> {
self.send_cmd(b"info")
}
fn send_cmd<'a, V: Deserialize<'a>>(&'a mut self, cmd: &[u8]) -> io::Result<V> {
self.udp.send(cmd)?;
let len = self.udp.recv(&mut self.buf)?;
match serde_yaml::from_slice::<V>(&self.buf[..len]) {
Ok(val) => Ok(val),
Err(e) => {
log::error!("{:?},{:?}", &buf[..len], e);
log::error!("{:?},{:?}", &self.buf[..len], e);
Err(io::Error::new(io::ErrorKind::Other, "data error"))
}
}
+2
View File
@@ -11,6 +11,8 @@ pub struct Info {
pub public_ips: String,
pub local_addr: String,
pub ipv6_addr: String,
pub up: u64,
pub down: u64,
}
#[derive(Serialize, Deserialize, Debug)]
+65 -35
View File
@@ -1,8 +1,9 @@
use crate::command::entity::{DeviceItem, Info, RouteItem};
use crate::console_out;
use std::io;
use vnt::core::Vnt;
use crate::command::entity::{DeviceItem, Info, RouteItem};
use crate::console_out;
pub mod client;
pub mod entity;
pub mod server;
@@ -17,12 +18,12 @@ pub enum CommandEnum {
pub fn command(cmd: CommandEnum) {
if let Err(e) = command_(cmd) {
println!("cmd: {}", e);
println!("cmd: {:?}", e);
}
}
fn command_(cmd: CommandEnum) -> io::Result<()> {
let command_client = client::CommandClient::new()?;
let mut command_client = client::CommandClient::new()?;
match cmd {
CommandEnum::Route => {
let list = command_client.route()?;
@@ -50,25 +51,31 @@ fn command_(cmd: CommandEnum) -> io::Result<()> {
pub fn command_route(vnt: &Vnt) -> Vec<RouteItem> {
let route_table = vnt.route_table();
let mut route_list = Vec::with_capacity(route_table.len());
for (destination, route) in route_table {
let next_hop = vnt
.route_key(&route.route_key())
.map_or(String::new(), |v| v.to_string());
let metric = route.metric.to_string();
let rt = if route.rt < 0 {
"".to_string()
} else {
route.rt.to_string()
};
let interface = route.addr.to_string();
let item = RouteItem {
destination: destination.to_string(),
next_hop,
metric,
rt,
interface,
};
route_list.push(item);
for (destination, routes) in route_table {
for route in routes {
let next_hop = vnt
.route_key(&route.route_key())
.map_or(String::new(), |v| v.to_string());
let metric = route.metric.to_string();
let rt = if route.rt < 0 {
"".to_string()
} else {
route.rt.to_string()
};
let interface = if route.is_tcp {
format!("tcp@{}", route.addr)
} else {
route.addr.to_string()
};
let item = RouteItem {
destination: destination.to_string(),
next_hop,
metric,
rt,
interface,
};
route_list.push(item);
}
}
route_list
}
@@ -87,8 +94,14 @@ pub fn command_list(vnt: &Vnt) -> Vec<DeviceItem> {
let public_ips: Vec<String> =
nat_info.public_ips.iter().map(|v| v.to_string()).collect();
let public_ips = public_ips.join(",");
let local_ip = nat_info.local_ipv4_addr.ip().to_string();
let ipv6 = nat_info.ipv6_addr.ip().to_string();
let local_ip = nat_info
.local_ipv4()
.map(|v| v.to_string())
.unwrap_or("None".to_string());
let ipv6 = nat_info
.ipv6()
.map(|v| v.to_string())
.unwrap_or("None".to_string());
(nat_type, public_ips, local_ip, ipv6)
} else {
(
@@ -100,11 +113,22 @@ pub fn command_list(vnt: &Vnt) -> Vec<DeviceItem> {
};
let (nat_traversal_type, rt) = if let Some(route) = vnt.route(&peer.virtual_ip) {
let nat_traversal_type = if route.metric == 1 {
"p2p"
} else if route.addr == info.connect_server {
"server-relay"
if route.is_tcp {
"tcp-p2p"
} else {
"p2p"
}
} else {
"client-relay"
let next_hop = vnt.route_key(&route.route_key());
if let Some(next_hop) = next_hop {
if info.is_gateway(&next_hop) {
"server-relay"
} else {
"client-relay"
}
} else {
"server-relay"
}
}
.to_string();
let rt = if route.rt < 0 {
@@ -148,12 +172,16 @@ pub fn command_info(vnt: &Vnt) -> Info {
let nat_type = format!("{:?}", nat_info.nat_type);
let public_ips: Vec<String> = nat_info.public_ips.iter().map(|v| v.to_string()).collect();
let public_ips = public_ips.join(",");
let local_addr = nat_info.local_ipv4_addr.to_string();
let ipv6_addr = if nat_info.ipv6_addr.ip().is_unspecified() {
"None".to_string()
} else {
nat_info.ipv6_addr.ip().to_string()
};
let local_addr = nat_info
.local_ipv4()
.map(|v| v.to_string())
.unwrap_or("None".to_string());
let ipv6_addr = nat_info
.ipv6()
.map(|v| v.to_string())
.unwrap_or("None".to_string());
let up = vnt.up_stream();
let down = vnt.down_stream();
Info {
name,
virtual_ip,
@@ -165,5 +193,7 @@ pub fn command_info(vnt: &Vnt) -> Info {
public_ips,
local_addr,
ipv6_addr,
up,
down,
}
}
+35 -29
View File
@@ -1,6 +1,6 @@
use std::io;
use std::io::Write;
use tokio::net::UdpSocket;
use std::net::UdpSocket;
use vnt::core::Vnt;
@@ -13,19 +13,27 @@ impl CommandServer {
}
impl CommandServer {
pub async fn start(self, vnt: Vnt) -> io::Result<()> {
let udp = UdpSocket::bind("127.0.0.1:0").await?;
let path_buf = crate::app_home()?.join("command-port");
let mut file = std::fs::File::create(path_buf)?;
file.write_all(udp.local_addr()?.port().to_string().as_bytes())?;
file.sync_all()?;
pub fn start(self, vnt: Vnt) -> io::Result<()> {
let udp = if let Ok(udp) = UdpSocket::bind("127.0.0.1:39271") {
udp
} else {
UdpSocket::bind("127.0.0.1:0")?
};
let addr = udp.local_addr()?;
log::info!("启动后台cmd:{:?}", addr);
if let Err(e) = save_port(addr.port()) {
log::warn!("保存后台命令端口失败:{:?}", e);
}
let mut buf = [0u8; 64];
loop {
let (len, addr) = udp.recv_from(&mut buf).await?;
let (len, addr) = udp.recv_from(&mut buf)?;
match std::str::from_utf8(&buf[..len]) {
Ok(cmd) => {
if let Ok(out) = command(cmd, &vnt) {
let _ = udp.send_to(out.as_bytes(), addr).await;
if let Err(e) = udp.send_to(out.as_bytes(), addr) {
log::warn!("cmd={},err={:?}", cmd, e);
}
if "stopped" == &out {
break;
}
@@ -39,33 +47,31 @@ impl CommandServer {
Ok(())
}
}
fn save_port(port: u16) -> io::Result<()> {
let path_buf = crate::app_home()?.join("command-port");
let mut file = std::fs::File::create(path_buf)?;
file.write_all(port.to_string().as_bytes())?;
file.sync_all()
}
fn command(cmd: &str, vnt: &Vnt) -> io::Result<String> {
let cmd = cmd.trim();
let out_str = match cmd {
"route" => match serde_json::to_string(&crate::command::command_route(vnt)) {
Ok(str) => str,
Err(e) => {
format!("{:?}", e)
}
},
"list" => match serde_json::to_string(&crate::command::command_list(vnt)) {
Ok(str) => str,
Err(e) => {
format!("{:?}", e)
}
},
"info" => match serde_json::to_string(&crate::command::command_info(vnt)) {
Ok(str) => str,
Err(e) => {
format!("{:?}", e)
}
},
"route" => serde_yaml::to_string(&crate::command::command_route(vnt))
.unwrap_or_else(|e| format!("error {:?}", e)),
"list" => serde_yaml::to_string(&crate::command::command_list(vnt))
.unwrap_or_else(|e| format!("error {:?}", e)),
"info" => serde_yaml::to_string(&crate::command::command_info(vnt))
.unwrap_or_else(|e| format!("error {:?}", e)),
"stop" => {
vnt.stop()?;
vnt.stop();
"stopped".to_string()
}
_ => {
format!("command '{}' not found. \n Try to enter: 'help'\n", cmd)
format!(
"command '{}' not found. Try to enter: 'route'/'list'/'stop' \n",
cmd
)
}
};
Ok(out_str)
+199
View File
@@ -0,0 +1,199 @@
use std::io;
use std::net::{Ipv4Addr, ToSocketAddrs};
use std::str::FromStr;
use serde::{Deserialize, Serialize};
use vnt::channel::punch::PunchModel;
use vnt::channel::UseChannelType;
use vnt::cipher::CipherModel;
use vnt::core::Config;
#[derive(Serialize, Deserialize, Debug)]
#[serde(default)]
pub struct FileConfig {
#[cfg(any(target_os = "windows", target_os = "linux"))]
pub tap: bool,
pub token: String,
pub device_id: String,
pub name: String,
pub server_address: String,
pub stun_server: Vec<String>,
pub in_ips: Vec<String>,
pub out_ips: Vec<String>,
pub password: Option<String>,
pub mtu: Option<u32>,
pub tcp: bool,
pub ip: Option<String>,
pub use_channel: String,
#[cfg(feature = "ip_proxy")]
pub no_proxy: bool,
pub server_encrypt: bool,
pub parallel: usize,
pub cipher_model: String,
pub finger: bool,
pub punch_model: String,
pub ports: Option<Vec<u16>>,
pub cmd: bool,
pub first_latency: bool,
pub device_name: Option<String>,
pub packet_loss: Option<f64>,
pub packet_delay: u32,
}
impl Default for FileConfig {
fn default() -> Self {
Self {
#[cfg(any(target_os = "windows", target_os = "linux"))]
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.qq.com:3478".to_string(),
],
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: "aes_gcm".to_string(),
finger: false,
punch_model: "all".to_string(),
ports: None,
cmd: false,
first_latency: false,
device_name: None,
packet_loss: None,
packet_delay: 0,
}
}
}
pub fn read_config(file_path: &str) -> io::Result<(Config, bool)> {
let conf = std::fs::read_to_string(file_path)?;
let file_conf = match serde_yaml::from_str::<FileConfig>(&conf) {
Ok(val) => val,
Err(e) => {
log::error!("{:?}", e);
return Err(io::Error::new(io::ErrorKind::Other, format!("{}", e)));
}
};
if file_conf.token.is_empty() {
return Err(io::Error::new(io::ErrorKind::Other, "token is_empty"));
}
let server_address = match file_conf.server_address.to_socket_addrs() {
Ok(mut addr) => {
if let Some(addr) = addr.next() {
addr
} else {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("server_address {:?} error", &file_conf.server_address),
));
}
}
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("server_address {:?} error:{}", &file_conf.server_address, e),
));
}
};
let in_ips = match common::args_parse::ips_parse(&file_conf.in_ips) {
Ok(in_ips) => in_ips,
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("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(io::Error::new(
io::ErrorKind::Other,
format!("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| {
io::Error::new(
io::ErrorKind::Other,
format!("ip {:?} error:{}", &file_conf.ip, e),
)
})?),
};
let cipher_model = CipherModel::from_str(&file_conf.cipher_model)
.map_err(|e| io::Error::new(io::ErrorKind::Other, e))?;
let punch_model = PunchModel::from_str(&file_conf.punch_model)
.map_err(|e| io::Error::new(io::ErrorKind::Other, e))?;
let use_channel_type = UseChannelType::from_str(&file_conf.use_channel)
.map_err(|e| io::Error::new(io::ErrorKind::Other, e))?;
let config = Config::new(
#[cfg(any(target_os = "windows", target_os = "linux"))]
file_conf.tap,
file_conf.token,
file_conf.device_id,
file_conf.name,
server_address,
file_conf.server_address,
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,
)
.unwrap();
Ok((config, file_conf.cmd))
}
pub fn get_device_id() -> String {
if let Some(id) = common::identifier::get_unique_identifier() {
id
} else {
let path_buf = match crate::app_home() {
Ok(path_buf) => path_buf.join("device-id"),
Err(e) => {
log::warn!("{:?}", e);
return String::new();
}
};
if let Ok(id) = std::fs::read_to_string(path_buf.as_path()) {
id
} else {
let id = uuid::Uuid::new_v4().to_string();
let _ = std::fs::write(path_buf, &id);
id
}
}
}
+25 -1
View File
@@ -18,6 +18,30 @@ pub fn console_info(status: Info) {
println!("Public ips: {}", style(status.public_ips).green());
println!("Local addr: {}", style(status.local_addr).green());
println!("IPv6: {}", style(status.ipv6_addr).green());
println!("Up: {}", style(convert(status.up)).green());
println!("Down: {}", style(convert(status.down)).green());
}
fn convert(num: u64) -> String {
let gigabytes = num / (1024 * 1024 * 1024);
let remaining_bytes = num % (1024 * 1024 * 1024);
let megabytes = remaining_bytes / (1024 * 1024);
let remaining_bytes = remaining_bytes % (1024 * 1024);
let kilobytes = remaining_bytes / 1024;
let remaining_bytes = remaining_bytes % 1024;
let mut s = String::new();
if gigabytes > 0 {
s.push_str(&format!("{} GB ", gigabytes));
}
if megabytes > 0 {
s.push_str(&format!("{} MB ", megabytes));
}
if kilobytes > 0 {
s.push_str(&format!("{} KB ", kilobytes));
}
if remaining_bytes > 0 {
s.push_str(&format!("{} bytes", remaining_bytes));
}
s
}
pub fn console_route_table(mut list: Vec<RouteItem>) {
@@ -76,7 +100,7 @@ pub fn console_device_list(mut list: Vec<DeviceItem>) {
("".to_string(), Style::new().red()),
]);
} else {
if &item.nat_traversal_type == "p2p" {
if item.nat_traversal_type.contains("p2p") {
out_list.push(vec![
(item.name, Style::new().green()),
(item.virtual_ip, Style::new().green()),
+316 -290
View File
@@ -1,28 +1,39 @@
use std::io;
use std::net::{Ipv4Addr, ToSocketAddrs};
use std::path::PathBuf;
use std::str::FromStr;
use std::{io, thread};
use console::style;
use getopts::Options;
use tokio::io::{AsyncBufReadExt, BufReader};
use tokio::signal;
use common::args_parse::{ips_parse, out_ips_parse};
use vnt::channel::punch::PunchModel;
use vnt::channel::UseChannelType;
use vnt::cipher::CipherModel;
use vnt::core::{Config, Vnt, VntUtil};
use vnt::handle::handshake_handler::HandshakeEnum;
use vnt::handle::registration_handler::ReqEnum;
use vnt::core::{Config, Vnt};
mod command;
mod config;
mod console_out;
mod generated_serial_number;
mod root_check;
pub fn app_home() -> io::Result<PathBuf> {
let path = dirs::home_dir()
.ok_or(io::Error::new(io::ErrorKind::Other, "not home"))?
.join(".vnt-cli");
let root_path = match std::env::current_exe() {
Ok(path) => {
if let Some(v) = path.as_path().parent() {
v.to_path_buf()
} else {
log::warn!("current_exe parent none:{:?}", path);
PathBuf::new()
}
}
Err(e) => {
log::warn!("current_exe err:{:?}", e);
PathBuf::new()
}
};
let path = root_path.join("env");
if !path.exists() {
std::fs::create_dir_all(&path)?;
}
@@ -41,25 +52,27 @@ fn main() {
opts.optopt("s", "", "注册和中继服务器地址", "<server>");
opts.optmulti("e", "", "stun服务器", "<stun-server>");
opts.optflag("a", "", "使用tap模式");
opts.optopt("", "nic", "虚拟网卡名称,windows下使用tap则必填", "<tun0>");
opts.optmulti("i", "", "配置点对网(IP代理)入站时使用", "<in-ip>");
opts.optmulti("o", "", "配置点对网出站时使用", "<out-ip>");
opts.optopt("w", "", "客户端加密", "<password>");
opts.optflag("W", "", "服务端加密");
opts.optflag("m", "", "模拟组播");
opts.optopt("u", "", "自定义mtu(默认为1430)", "<mtu>");
opts.optflag("", "tcp", "tcp");
opts.optopt("", "ip", "指定虚拟ip", "<ip>");
opts.optflag("", "relay", "仅使用服务器转发");
opts.optopt("", "par", "任务并行度(必须为正整数)", "<parallel>");
opts.optopt("", "thread", "线程数(必须为正整数)", "<thread>");
opts.optopt("", "model", "加密模式", "<model>");
opts.optflag("", "finger", "指纹校验");
opts.optopt(
"",
"punch",
"取值ipv4/ipv6,表示仅使用ipv4或ipv6打洞",
"<punch>",
);
opts.optopt("", "punch", "取值ipv4/ipv6", "<punch>");
opts.optopt("", "ports", "监听的端口", "<port,port>");
opts.optflag("", "cmd", "开启窗口输入");
opts.optflag("", "no-proxy", "关闭内置代理");
opts.optflag("", "first-latency", "优先延迟");
opts.optopt("", "use-channel", "使用通道 relay/p2p", "<use-channel>");
opts.optopt("", "packet-loss", "丢包率", "<packet-loss>");
opts.optopt("", "packet-delay", "延迟", "<packet-delay>");
opts.optopt("f", "", "配置文件", "<conf>");
//"后台运行时,查看其他设备列表"
opts.optflag("", "list", "后台运行时,查看其他设备列表");
opts.optflag("", "all", "后台运行时,查看其他设备完整信息");
@@ -101,299 +114,266 @@ fn main() {
command::command(command::CommandEnum::All);
return;
}
if !matches.opt_present("k") {
print_usage(&program, opts);
println!("parameter -k not found .");
return;
}
let tap = matches.opt_present("a");
let token: String = matches.opt_get("k").unwrap().unwrap();
let device_id = matches.opt_get_default("d", String::new()).unwrap();
let device_id = if device_id.is_empty() {
if let Some(id) = common::identifier::get_unique_identifier() {
id
} else {
let path_buf = app_home().unwrap().join("device-id");
if let Ok(id) = std::fs::read_to_string(path_buf.as_path()) {
id
} else {
let id = uuid::Uuid::new_v4().to_string();
let _ = std::fs::write(path_buf, &id);
id
}
}
} else {
device_id
};
if device_id.is_empty() {
print_usage(&program, opts);
println!("parameter -d not found .");
return;
}
let name = matches
.opt_get_default("n", os_info::get().to_string())
.unwrap();
let server_address_str = matches
.opt_get_default("s", "nat1.wherewego.top:29872".to_string())
.unwrap();
let server_address = match server_address_str.to_socket_addrs() {
Ok(mut addr) => {
if let Some(addr) = addr.next() {
addr
} else {
println!("parameter '-s {}' error .", server_address_str);
let conf = matches.opt_str("f");
let (config, cmd) = if conf.is_some() {
match config::read_config(&conf.unwrap()) {
Ok(c) => c,
Err(e) => {
println!("conf err {}", e);
return;
}
}
Err(e) => {
println!("parameter '-s {}' error {}.", server_address_str, e);
} else {
if !matches.opt_present("k") {
print_usage(&program, opts);
println!("parameter -k not found .");
return;
}
};
let mut stun_server = matches.opt_strs("e");
if stun_server.is_empty() {
stun_server.push("stun1.l.google.com:19302".to_string());
stun_server.push("stun2.l.google.com:19302".to_string());
stun_server.push("stun.qq.com:3478".to_string());
}
#[cfg(any(target_os = "windows", target_os = "linux"))]
let tap = matches.opt_present("a");
let device_name = matches.opt_str("nic");
let token: String = matches.opt_get("k").unwrap().unwrap();
let device_id = matches.opt_get_default("d", String::new()).unwrap();
let device_id = if device_id.is_empty() {
config::get_device_id()
} else {
device_id
};
if device_id.is_empty() {
print_usage(&program, opts);
println!("parameter -d not found .");
return;
}
let name = matches
.opt_get_default("n", os_info::get().to_string())
.unwrap();
let server_address_str = matches
.opt_get_default("s", "nat1.wherewego.top:29872".to_string())
.unwrap();
let server_address = match server_address_str.to_socket_addrs() {
Ok(mut addr) => {
if let Some(addr) = addr.next() {
addr
} else {
println!("parameter '-s {}' error .", server_address_str);
return;
}
}
Err(e) => {
println!("parameter '-s {}' error {}.", server_address_str, e);
return;
}
};
let mut stun_server = matches.opt_strs("e");
if stun_server.is_empty() {
stun_server.push("stun1.l.google.com:19302".to_string());
stun_server.push("stun2.l.google.com:19302".to_string());
stun_server.push("stun.qq.com:3478".to_string());
}
let in_ip = matches.opt_strs("i");
let in_ip = match ips_parse(&in_ip) {
Ok(in_ip) => in_ip,
Err(e) => {
print_usage(&program, opts);
println!();
println!("-i: {:?} {}", in_ip, e);
println!("example: -i 192.168.0.0/24,10.26.0.3");
return;
}
};
let out_ip = matches.opt_strs("o");
let out_ip = match out_ips_parse(&out_ip) {
Ok(out_ip) => out_ip,
Err(e) => {
print_usage(&program, opts);
println!();
println!("-o: {:?} {}", out_ip, e);
println!("example: -o 0.0.0.0/0");
return;
}
};
let password: Option<String> = matches.opt_get("w").unwrap();
let server_encrypt = matches.opt_present("W");
let simulate_multicast = matches.opt_present("m");
let unused_cmd = matches.opt_present("c");
let mtu: Option<String> = matches.opt_get("u").unwrap();
let mtu = if let Some(mtu) = mtu {
match u16::from_str(&mtu) {
Ok(mtu) => Some(mtu),
let in_ip = matches.opt_strs("i");
let in_ip = match ips_parse(&in_ip) {
Ok(in_ip) => in_ip,
Err(e) => {
print_usage(&program, opts);
println!();
println!("'-u {}' {}", mtu, e);
println!("-i: {:?} {}", in_ip, e);
println!("example: -i 192.168.0.0/24,10.26.0.3");
return;
}
};
let out_ip = matches.opt_strs("o");
let out_ip = match out_ips_parse(&out_ip) {
Ok(out_ip) => out_ip,
Err(e) => {
print_usage(&program, opts);
println!();
println!("-o: {:?} {}", out_ip, e);
println!("example: -o 0.0.0.0/0");
return;
}
};
let password: Option<String> = matches.opt_get("w").unwrap();
let server_encrypt = matches.opt_present("W");
#[cfg(not(feature = "server_encrypt"))]
{
if server_encrypt {
println!("Server encryption not supported");
return;
}
}
} else {
None
};
let virtual_ip: Option<String> = matches.opt_get("ip").unwrap();
let virtual_ip =
virtual_ip.map(|v| Ipv4Addr::from_str(&v).expect(&format!("'--ip {}' error", v)));
if let Some(virtual_ip) = virtual_ip {
if virtual_ip.is_unspecified() || virtual_ip.is_broadcast() || virtual_ip.is_multicast() {
println!("'--ip {}' invalid", virtual_ip);
let mtu: Option<String> = matches.opt_get("u").unwrap();
let mtu = if let Some(mtu) = mtu {
match u32::from_str(&mtu) {
Ok(mtu) => Some(mtu),
Err(e) => {
print_usage(&program, opts);
println!();
println!("'-u {}' {}", mtu, e);
return;
}
}
} else {
None
};
let virtual_ip: Option<String> = matches.opt_get("ip").unwrap();
let virtual_ip =
virtual_ip.map(|v| Ipv4Addr::from_str(&v).expect(&format!("'--ip {}' error", v)));
if let Some(virtual_ip) = virtual_ip {
if virtual_ip.is_unspecified() || virtual_ip.is_broadcast() || virtual_ip.is_multicast()
{
println!("'--ip {}' invalid", virtual_ip);
return;
}
}
let tcp_channel = matches.opt_present("tcp");
let relay = matches.opt_present("relay");
let parallel = matches.opt_get::<usize>("par").unwrap().unwrap_or(1);
if parallel == 0 {
println!("'--par {}' invalid", parallel);
return;
}
}
let tcp_channel = matches.opt_present("tcp");
let relay = matches.opt_present("relay");
let parallel = matches.opt_get::<usize>("par").unwrap().unwrap_or(1);
if parallel == 0 {
println!("'--par {}' invalid", parallel);
return;
}
let cipher_model = matches
.opt_get::<CipherModel>("model")
.unwrap()
.unwrap_or(CipherModel::AesGcm);
let cipher_model = match matches.opt_get::<CipherModel>("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() {
println!("'--model ' undefined");
return;
}
model.unwrap_or(CipherModel::None)
}
#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))]
model.unwrap_or(CipherModel::AesGcm)
}
Err(e) => {
println!("'--model ' invalid,{}", e);
return;
}
};
let finger = matches.opt_present("finger");
let punch_model = matches
.opt_get::<PunchModel>("punch")
.unwrap()
.unwrap_or(PunchModel::All);
let finger = matches.opt_present("finger");
let punch_model = matches
.opt_get::<PunchModel>("punch")
.unwrap()
.unwrap_or(PunchModel::All);
let use_channel_type = matches
.opt_get::<UseChannelType>("use-channel")
.unwrap()
.unwrap_or_else(|| {
if relay {
UseChannelType::Relay
} else {
UseChannelType::All
}
});
let ports = matches
.opt_get::<String>("ports")
.unwrap_or(None)
.map(|v| v.split(",").map(|x| x.parse().unwrap_or(0)).collect());
let cmd = matches.opt_present("cmd");
#[cfg(feature = "ip_proxy")]
let no_proxy = matches.opt_present("no-proxy");
let first_latency = matches.opt_present("first-latency");
let packet_loss = matches
.opt_get::<f64>("packet-loss")
.expect("--packet-loss");
let packet_delay = matches
.opt_get::<u32>("packet-delay")
.expect("--packet-delay")
.unwrap_or(0);
let config = match Config::new(
#[cfg(any(target_os = "windows", target_os = "linux"))]
tap,
token,
device_id,
name,
server_address,
server_address_str,
stun_server,
in_ip,
out_ip,
password,
mtu,
tcp_channel,
virtual_ip,
#[cfg(feature = "ip_proxy")]
no_proxy,
server_encrypt,
parallel,
cipher_model,
finger,
punch_model,
ports,
first_latency,
device_name,
use_channel_type,
packet_loss,
packet_delay,
) {
Ok(config) => config,
Err(e) => {
println!("config error: {}", e);
return;
}
};
(config, cmd)
};
println!("version {}", vnt::VNT_VERSION);
let config = Config::new(
tap,
token,
device_id,
name,
server_address,
server_address_str,
stun_server,
in_ip,
out_ip,
password,
simulate_multicast,
mtu,
tcp_channel,
virtual_ip,
relay,
server_encrypt,
parallel,
cipher_model,
finger,
punch_model,
);
main0(config, !unused_cmd);
println!("Serial:{}", generated_serial_number::SERIAL_NUMBER);
main0(config, cmd);
std::process::exit(0);
}
#[tokio::main]
async fn main0(config: Config, show_cmd: bool) {
let server_encrypt = config.server_encrypt;
let mut vnt_util = VntUtil::new(config).await.unwrap();
let mut conn_count = 0;
let response = loop {
if conn_count > 0 {
tokio::time::sleep(std::time::Duration::from_secs(2)).await;
}
conn_count += 1;
if let Err(e) = vnt_util.connect().await {
println!("connect server failed {}", e);
return;
}
match vnt_util.handshake().await {
Ok(response) => {
if server_encrypt {
let finger = response.unwrap().finger().unwrap();
println!("{}{}", green("server fingerprint:".to_string()), finger);
match vnt_util.secret_handshake().await {
Ok(_) => {}
Err(e) => {
match e {
HandshakeEnum::NotSecret => {}
HandshakeEnum::KeyError => {}
HandshakeEnum::Timeout => {
println!("handshake timeout")
}
HandshakeEnum::ServerError(str) => {
println!("error:{}", str);
}
HandshakeEnum::Other(str) => {
println!("error:{}", str);
}
}
continue;
}
}
}
match vnt_util.register().await {
Ok(response) => {
break response;
}
Err(e) => match e {
ReqEnum::TokenError => {
println!("token error");
return;
}
ReqEnum::AddressExhausted => {
println!("address exhausted");
return;
}
ReqEnum::Timeout => {
println!("timeout...");
}
ReqEnum::ServerError(str) => {
println!("error:{}", str);
}
ReqEnum::Other(str) => {
println!("error:{}", str);
}
ReqEnum::IpAlreadyExists => {
println!("ip already exists");
return;
}
ReqEnum::InvalidIp => {
println!("invalid ip");
return;
}
},
}
mod callback;
fn main0(config: Config, show_cmd: bool) {
let vnt_util = Vnt::new(config, callback::VntHandler {}).unwrap();
let vnt_c = vnt_util.clone();
thread::Builder::new()
.name("CommandServer".into())
.spawn(move || {
if let Err(e) = command::server::CommandServer::new().start(vnt_c) {
log::warn!("cmd:{:?}", e);
}
Err(e) => match e {
HandshakeEnum::NotSecret => {
println!("The server does not support encryption");
return;
}
HandshakeEnum::KeyError => {}
HandshakeEnum::Timeout => {
println!("handshake timeout")
}
HandshakeEnum::ServerError(str) => {
println!("error:{}", str);
}
HandshakeEnum::Other(str) => {
println!("error:{}", str);
}
},
}
};
println!(" ====== Connect Successfully ====== ");
println!("virtual_gateway:{}", response.virtual_gateway);
println!("virtual_ip:{}", green(response.virtual_ip.to_string()));
let driver_info = vnt_util.create_iface().unwrap();
println!(" ====== Create Network Interface Successfully ====== ");
println!("name:{}", driver_info.name);
println!("version:{}", driver_info.version);
let mut vnt = match vnt_util.build().await {
Ok(vnt) => vnt,
Err(e) => {
println!("error:{}", e);
return;
}
};
println!(" ====== Start Successfully ====== ");
let vnt_c = vnt.clone();
tokio::spawn(async {
if let Err(e) = command::server::CommandServer::new().start(vnt_c).await {
println!("command error :{}", e);
}
});
})
.expect("CommandServer");
if show_cmd {
let stdin = tokio::io::stdin();
let mut cmd = String::new();
let mut reader = BufReader::new(stdin);
loop {
cmd.clear();
println!("input:list,info,route,all,stop");
tokio::select! {
_ = vnt.wait_stop()=>{
return;
}
_ = signal::ctrl_c()=>{
let _ = vnt.stop();
vnt.wait_stop_ms(std::time::Duration::from_secs(3)).await;
std::process::exit(0);
}
rs = reader.read_line(&mut cmd)=>{
match rs {
Ok(len) => {
if !command(&cmd[..len],&vnt){
break;
}
}
Err(e) => {
println!("input err:{}",e);
break;
}
println!("======== input:list,info,route,all,stop ========");
match io::stdin().read_line(&mut cmd) {
Ok(len) => {
if !command(&cmd[..len], &vnt_util) {
break;
}
}
Err(e) => {
println!("input err:{}", e);
break;
}
}
}
}
vnt.wait_stop().await;
vnt_util.wait()
}
fn command(cmd: &str, vnt: &Vnt) -> bool {
@@ -430,31 +410,77 @@ fn command(cmd: &str, vnt: &Vnt) -> bool {
fn print_usage(program: &str, _opts: Options) {
println!("Usage: {} [options]", program);
println!("version:{}", vnt::VNT_VERSION);
println!("Serial:{}", generated_serial_number::SERIAL_NUMBER);
println!("Options:");
println!(
" -k <token> {}",
green("必选,使用相同的token,就能组建一个局域网络".to_string())
green("使用相同的token,就能组建一个局域网络".to_string())
);
println!(" -n <name> 给设备一个名字,便于区分不同设备,默认使用系统版本");
println!(" -d <id> 设备唯一标识符,不使用--ip参数时,服务端凭此参数分配虚拟ip");
println!(" -c 关闭交互式命令,使用此参数禁用控制台输入");
println!(" -d <id> 设备唯一标识符,不使用--ip参数时,服务端凭此参数分配虚拟ip,注意不能重复");
println!(" -s <server> 注册和中继服务器地址");
println!(" -e <stun-server> stun服务器,用于探测NAT类型,可多次指定,如-e addr1 -e addr2");
println!(" -a 使用tap模式,默认使用tun模式");
println!(" -i <in-ip> 配置点对网(IP代理)时使用,-i 192.168.0.0/24,10.26.0.3表示允许接收网段192.168.0.0/24的数据");
println!(" 并转发到10.26.0.3,可指定多个网段");
#[cfg(feature = "ip_proxy")]
println!(" -o <out-ip> 配置点对网时使用,-o 192.168.0.0/24表示允许将数据转发到192.168.0.0/24,可指定多个网段");
println!(" -w <password> 使用该密码生成的密钥对客户端数据进行加密,并且服务端无法解密,使用相同密码的客户端才能通信");
#[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 <password> 使用该密码生成的密钥对客户端数据进行加密,并且服务端无法解密,使用相同密码的客户端才能通信");
}
#[cfg(feature = "server_encrypt")]
println!(" -W 加密当前客户端和服务端通信的数据,请留意服务端指纹是否正确");
println!(" -m 模拟组播,默认情况下组播数据会被当作广播发送,开启后会模拟真实组播的数据发送");
println!(" -u <mtu> 自定义mtu(不加密默认为1450,加密默认为1410)");
println!(" -f <conf_file> 读取配置文件中的配置");
println!(" --tcp 和服务端使用tcp通信,默认使用udp,遇到udp qos时可指定使用tcp");
println!(" --ip <ip> 指定虚拟ip,指定的ip不能和其他设备重复,必须有效并且在服务端所属网段下,默认情况由服务端分配");
println!(" --relay 仅使用服务器转发,不使用p2p,默认情况允许使用p2p");
println!(" --par <parallel> 任务并行度(必须为正整数),默认值为1");
println!(" --model <model> 加密模式(默认aes_gcm),可选值aes_gcm/aes_cbc/aes_ecb,一般来说性能:aes_ecb>aes_cbc>aes_gcm");
println!(" --finger 增加数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启");
println!(" --punch <punch> 取值ipv4/ipv6ipv4表示仅使用ipv4打洞");
if !enums.is_empty() {
println!(
" --model <model> 加密模式(默认aes_gcm),可选值{}",
&enums[1..]
);
}
if !enums.is_empty() {
println!(" --finger 增加数据指纹校验,可增加安全性,如果服务端开启指纹校验,则客户端也必须开启");
}
println!(" --punch <punch> 取值ipv4/ipv6/all,ipv4表示仅使用ipv4打洞");
println!(" --ports <port,port> 取值0~65535,指定本地监听的一组端口,默认监听两个随机端口,使用过多端口会增加网络负担");
println!(" --cmd 开启交互式命令,使用此参数开启控制台输入");
#[cfg(feature = "ip_proxy")]
println!(" --no-proxy 关闭内置代理,如需点对网则需要配置网卡NAT转发");
println!(" --first-latency 优先低延迟的通道,默认情况优先使用p2p通道");
println!(" --use-channel <p2p> 使用通道 relay/p2p/all,默认两者都使用");
println!(" --nic <tun0> 指定虚拟网卡名称");
println!(" --packet-loss <0> 模拟丢包,取值0~1之间的小数,程序会按设定的概率主动丢包,可用于模拟弱网");
println!(
" --packet-delay <0> 模拟延迟,整数,单位毫秒(ms),程序会按设定的值延迟发包,可用于模拟弱网"
);
println!();
println!(
+5 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "vnt-jni"
version = "1.2.3"
version = "1.2.9"
edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
@@ -8,7 +8,11 @@ edition = "2021"
[dependencies]
common = { path = "../common" }
vnt = {path="../vnt"}
parking_lot = "0.12.1"
jni = { version = "0.21.1", default-features = false }
log = "0.4.20"
spki = { version = "0.7.2", features = ["fingerprint", "alloc","base64","pem"]}
[lib]
crate-type = ["staticlib", "cdylib"]
@@ -0,0 +1,53 @@
package top.wherewego.vnt.jni;
import top.wherewego.vnt.jni.param.*;
/**
* 回调
*
* @author https://github.com/lbl8603/vnt
*/
public interface CallBack {
/**
* 创建虚拟网卡成功的回调方法
*
* @param info 网卡信息
*/
void createTun(DeviceInfo info);
/**
* 连接服务端
*
* @param info 将要连接的服务端信息
*/
void connect(ConnectInfo info);
/**
* 和服务端握手
*
* @param info 握手信息
* @return 是否确认握手
*/
boolean handshake(HandshakeInfo info);
/**
* 注册成功回调
*
* @param info 注册信息
* @return 是否确认注册信息
*/
boolean register(RegisterInfo info);
/**
* 异常回调
*
* @param info 错误信息
*/
void error(ErrorInfo info);
/**
* 服务停止
*/
void stop();
}
@@ -0,0 +1,264 @@
package top.wherewego.vnt.jni;
/**
* 启动配置
*
* @author https://github.com/lbl8603/vnt
*/
public class Config {
/**
* 是否是tap模式,仅支持windows和linux
*/
private boolean tap;
/**
* 组网标识
*/
private String token;
/**
* 设备名称
*/
private String name;
/**
* 客户端间加密的密码
*/
private String password;
/**
* 客户端间加密模式 aes_gcm/aes_cbc/aes_ecb/sm4_cbc
*/
private String cipherModel;
/**
* 打洞模式 ipv4/ipv6/all
*/
private String punchModel;
/**
* mtu 默认自动计算
*/
private Integer mtu;
/**
* 是否开启服务端加密
*/
private boolean serverEncrypt;
/**
* 仅使用中继转发
*/
private boolean relay;
/**
* 设备id,请使用唯一值
*/
private String deviceId;
/**
* 服务端地址
*/
private String server;
/**
* stun服务地址
*/
private String[] stunServer;
/**
* 和服务端使用tcp通信,默认使用udp
*/
private boolean tcp;
/**
* 指定组网IP
*/
private String ip;
/**
* 开启加密指纹校验
*/
private boolean finger;
/**
* 延迟优先,默认p2p优先
*/
private boolean firstLatency;
/**
* 点对网入口 格式 192.168.0.0/26,10.26.0.2
*/
private String[] inIps;
/**
* 点对网出口 格式 192.168.0.0/26
*/
private String[] outIps;
/**
* 端口组,udp会监听一组端口,tcp监听ports[0]端口
*/
private int[] ports;
/**
* 虚拟网卡名称 仅在linux、windows、macos上支持
*/
private String deviceName;
/**
* 虚拟网卡fd 仅在android上支持
*/
private int deviceFd;
public Config() {
}
public boolean isTap() {
return tap;
}
public void setTap(boolean tap) {
this.tap = tap;
}
public String getToken() {
return token;
}
public void setToken(String token) {
this.token = token;
}
public String getName() {
return name;
}
public void setName(String name) {
this.name = name;
}
public String getPassword() {
return password;
}
public void setPassword(String password) {
this.password = password;
}
public String getCipherModel() {
return cipherModel;
}
public void setCipherModel(String cipherModel) {
this.cipherModel = cipherModel;
}
public String getPunchModel() {
return punchModel;
}
public void setPunchModel(String punchModel) {
this.punchModel = punchModel;
}
public Integer getMtu() {
return mtu;
}
public void setMtu(Integer mtu) {
this.mtu = mtu;
}
public boolean isServerEncrypt() {
return serverEncrypt;
}
public void setServerEncrypt(boolean serverEncrypt) {
this.serverEncrypt = serverEncrypt;
}
public boolean isRelay() {
return relay;
}
public void setRelay(boolean relay) {
this.relay = relay;
}
public String getDeviceId() {
return deviceId;
}
public void setDeviceId(String deviceId) {
this.deviceId = deviceId;
}
public String getServer() {
return server;
}
public void setServer(String server) {
this.server = server;
}
public String[] getStunServer() {
return stunServer;
}
public void setStunServer(String[] stunServer) {
this.stunServer = stunServer;
}
public boolean isTcp() {
return tcp;
}
public void setTcp(boolean tcp) {
this.tcp = tcp;
}
public String getIp() {
return ip;
}
public void setIp(String ip) {
this.ip = ip;
}
public boolean isFinger() {
return finger;
}
public void setFinger(boolean finger) {
this.finger = finger;
}
public boolean isFirstLatency() {
return firstLatency;
}
public void setFirstLatency(boolean firstLatency) {
this.firstLatency = firstLatency;
}
public String[] getInIps() {
return inIps;
}
public void setInIps(String[] inIps) {
this.inIps = inIps;
}
public String[] getOutIps() {
return outIps;
}
public void setOutIps(String[] outIps) {
this.outIps = outIps;
}
public int[] getPorts() {
return ports;
}
public void setPorts(int[] ports) {
this.ports = ports;
}
public String getDeviceName() {
return deviceName;
}
public void setDeviceName(String deviceName) {
this.deviceName = deviceName;
}
public int getDeviceFd() {
return deviceFd;
}
public void setDeviceFd(int deviceFd) {
this.deviceFd = deviceFd;
}
}
@@ -0,0 +1,29 @@
package top.wherewego.vnt.jni;
/**
* @author lubeilin
* @date: 2024/02/27 18:31
*/
public class IpUtils {
public static String intToIpAddress(int ipAddress) {
return ((ipAddress & 0xFF000000) >>> 24) + "." +
((ipAddress & 0x00FF0000) >>> 16) + "." +
((ipAddress & 0x0000FF00) >>> 8) + "." +
(ipAddress & 0x000000FF);
}
public static int subnetMaskToPrefixLength(int subnetMask) {
int prefixLength = 0;
int bit = 1 << 31;
while (subnetMask != 0) {
if ((subnetMask & bit) != bit) {
break;
}
prefixLength++;
subnetMask <<= 1;
}
return prefixLength;
}
}
@@ -0,0 +1,46 @@
package top.wherewego.vnt.jni;
/**
* 对端设备信息
*
* @author https://github.com/lbl8603/vnt
*/
public class PeerDeviceInfo {
private final int virtualIp;
private final String name;
private final String status;
private final Route route;
public PeerDeviceInfo(int virtualIp, String name, String status, Route route) {
this.virtualIp = virtualIp;
this.name = name;
this.status = status;
this.route = route;
}
public int getVirtualIp() {
return virtualIp;
}
public String getName() {
return name;
}
public String getStatus() {
return status;
}
public Route getRoute() {
return route;
}
@Override
public String toString() {
return "PeerDeviceInfo{" +
"virtualIp=" + IpUtils.intToIpAddress(virtualIp) +
", name='" + name + '\'' +
", status='" + status + '\'' +
", route=" + route +
'}';
}
}
@@ -0,0 +1,39 @@
package top.wherewego.vnt.jni;
/**
* 路由信息
*
* @author https://github.com/lbl8603/vnt
*/
public class Route {
private final String address;
private final byte metric;
private final int rt;
public Route(String address, byte metric, int rt) {
this.address = address;
this.metric = metric;
this.rt = rt;
}
public String getAddress() {
return address;
}
public byte getMetric() {
return metric;
}
public int getRt() {
return rt;
}
@Override
public String toString() {
return "Route{" +
"address='" + address + '\'' +
", metric=" + metric +
", rt=" + rt +
'}';
}
}
@@ -0,0 +1,47 @@
package top.wherewego.vnt.jni;
import java.io.Closeable;
import java.io.IOException;
/**
* vnt的Java映射
*
* @author https://github.com/lbl8603/vnt
*/
public class Vnt implements Closeable {
private final long raw;
public Vnt(Config config, CallBack callBack) {
this.raw = new0(config, callBack);
if(this.raw == 0){
throw new RuntimeException();
}
}
public void stop() {
stop0(raw);
}
public void await() {
wait0(raw);
}
public PeerDeviceInfo[] list() {
return list0(raw);
}
private native long new0(Config config, CallBack callBack);
private native void stop0(long raw);
private native void wait0(long raw);
private native void drop0(long raw);
private native PeerDeviceInfo[] list0(long raw);
@Override
public void close() throws IOException {
drop0(raw);
}
}
@@ -0,0 +1,32 @@
package top.wherewego.vnt.jni.param;
/**
* 连接信息
*
* @author https://github.com/lbl8603/vnt
*/
public class ConnectInfo {
private final long count;
private final String address;
public ConnectInfo(long count, String address) {
this.count = count;
this.address = address;
}
public long getCount() {
return count;
}
public String getAddress() {
return address;
}
@Override
public String toString() {
return "ConnectInfo{" +
"count=" + count +
", address='" + address + '\'' +
'}';
}
}
@@ -0,0 +1,38 @@
package top.wherewego.vnt.jni.param;
/**
* 网卡信息
*
* @author https://github.com/lbl8603/vnt
*/
public class DeviceInfo {
/**
* 虚拟网卡名称
*/
private final String name;
/**
* 虚拟网卡版本
*/
private final String version;
public DeviceInfo(String name, String version) {
this.name = name;
this.version = version;
}
public String getName() {
return name;
}
public String getVersion() {
return version;
}
@Override
public String toString() {
return "DeviceInfo{" +
"name='" + name + '\'' +
", version='" + version + '\'' +
'}';
}
}
@@ -0,0 +1,55 @@
package top.wherewego.vnt.jni.param;
/**
* 异常回调信息
*
* @author https://github.com/lbl8603/vnt
*/
public class ErrorInfo {
/**
* 错误码
*/
public final ErrorCodeEnum code;
/**
* 错误信息,可能为空
*/
public final String msg;
public ErrorInfo(int code, String msg) {
this.code = switch (code) {
case 1 -> ErrorCodeEnum.TokenError;
case 2 -> ErrorCodeEnum.Disconnect;
case 3 -> ErrorCodeEnum.AddressExhausted;
case 4 -> ErrorCodeEnum.IpAlreadyExists;
case 5 -> ErrorCodeEnum.InvalidIp;
case 6 -> ErrorCodeEnum.Unknown;
default -> null;
};
this.msg = msg;
}
public ErrorCodeEnum getCode() {
return code;
}
public String getMsg() {
return msg;
}
public enum ErrorCodeEnum {
TokenError,
Disconnect,
AddressExhausted,
IpAlreadyExists,
InvalidIp,
Unknown,
}
@Override
public String toString() {
return "ErrorInfo{" +
"code=" + code +
", msg='" + msg + '\'' +
'}';
}
}
@@ -0,0 +1,54 @@
package top.wherewego.vnt.jni.param;
/**
* 握手回调信息
*
* @author https://github.com/lbl8603/vnt
*/
public class HandshakeInfo {
/**
* 公钥 pem格式 CRLF分隔,不加密时为空
*/
private final String publicKey;
/**
* 公钥签名,不加密时为空
*/
private final String finger;
/**
* 服务端版本
*/
private final String version;
public HandshakeInfo() {
this.publicKey = "publicKey";
this.finger = "finger";
this.version = "version";
}
public HandshakeInfo(String publicKey, String finger, String version) {
this.publicKey = publicKey;
this.finger = finger;
this.version = version;
}
public String getPublicKey() {
return publicKey;
}
public String getFinger() {
return finger;
}
public String getVersion() {
return version;
}
@Override
public String toString() {
return "HandshakeInfo{" +
"publicKey='" + publicKey + '\'' +
", finger='" + finger + '\'' +
", version='" + version + '\'' +
'}';
}
}
@@ -0,0 +1,48 @@
package top.wherewego.vnt.jni.param;
/**
* 注册回调信息
*
* @author https://github.com/lbl8603/vnt
*/
public class RegisterInfo {
/**
* 虚拟IP
*/
public final String virtualIp;
/**
* 掩码
*/
public final String virtualNetmask;
/**
* 网关
*/
public final String virtualGateway;
public RegisterInfo(String virtualIp, String virtualNetmask, String virtualGateway) {
this.virtualIp = virtualIp;
this.virtualNetmask = virtualNetmask;
this.virtualGateway = virtualGateway;
}
public String getVirtualIp() {
return virtualIp;
}
public String getVirtualNetmask() {
return virtualNetmask;
}
public String getVirtualGateway() {
return virtualGateway;
}
@Override
public String toString() {
return "RegisterInfo{" +
"virtualIp='" + virtualIp + '\'' +
", virtualNetmask='" + virtualNetmask + '\'' +
", virtualGateway='" + virtualGateway + '\'' +
'}';
}
}
+187
View File
@@ -0,0 +1,187 @@
use std::sync::Arc;
use jni::objects::{GlobalRef, JString, JValue};
use jni::{JNIEnv, JavaVM};
use spki::der::pem::LineEnding;
use spki::EncodePublicKey;
use vnt::handle::callback::ConnectInfo;
use vnt::{DeviceInfo, ErrorInfo, HandshakeInfo, RegisterInfo, VntCallback};
#[derive(Clone)]
pub struct CallBack {
jvm: Arc<JavaVM>,
this: GlobalRef,
}
unsafe impl Send for CallBack {}
impl CallBack {
pub fn new(jvm: JavaVM, this: GlobalRef) -> Self {
Self {
jvm: Arc::new(jvm),
this,
}
}
}
impl CallBack {
fn create_tun0(&self, info: DeviceInfo) -> jni::errors::Result<()> {
let env = &mut self.jvm.attach_current_thread()? as &mut JNIEnv;
let param = env.new_object(
"top/wherewego/vnt/jni/param/DeviceInfo",
"(Ljava/lang/String;Ljava/lang/String;)V",
&[
JValue::Object(&env.new_string(info.name)?.into()),
JValue::Object(&env.new_string(info.version)?.into()),
],
)?;
env.call_method(
&self.this,
"createTun",
"(Ltop/wherewego/vnt/jni/param/DeviceInfo;)V",
&[JValue::Object(&param)],
)?;
Ok(())
}
fn connect0(&self, info: ConnectInfo) -> jni::errors::Result<()> {
let env = &mut self.jvm.attach_current_thread()? as &mut JNIEnv;
let param = env.new_object(
"top/wherewego/vnt/jni/param/ConnectInfo",
"(JLjava/lang/String;)V",
&[
JValue::Long(info.count as _),
JValue::Object(&env.new_string(info.address.to_string())?.into()),
],
)?;
env.call_method(
&self.this,
"connect",
"(Ltop/wherewego/vnt/jni/param/ConnectInfo;)V",
&[JValue::Object(&param)],
)?;
Ok(())
}
fn handshake0(&self, info: HandshakeInfo) -> jni::errors::Result<bool> {
let env = &mut self.jvm.attach_current_thread()? as &mut JNIEnv;
let public_key = if let Some(public_key) = info.public_key {
match public_key.to_public_key_pem(LineEnding::CRLF) {
Ok(public_key) => env.new_string(public_key)?,
Err(e) => {
log::warn!("{:?}", e);
JString::default()
}
}
} else {
JString::default()
};
let finger = if let Some(finger) = info.finger {
env.new_string(finger)?
} else {
JString::default()
};
let param = env.new_object(
"top/wherewego/vnt/jni/param/HandshakeInfo",
"(Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;)V",
&[
JValue::Object(&public_key),
JValue::Object(&finger),
JValue::Object(&env.new_string(info.version)?.into()),
],
)?;
let rs = env.call_method(
&self.this,
"handshake",
"(Ltop/wherewego/vnt/jni/param/HandshakeInfo;)Z",
&[JValue::Object(&param)],
)?;
rs.z()
}
fn register0(&self, info: RegisterInfo) -> jni::errors::Result<bool> {
let env = &mut self.jvm.attach_current_thread()? as &mut JNIEnv;
let param = env.new_object(
"top/wherewego/vnt/jni/param/RegisterInfo",
"(Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;)V",
&[
JValue::Object(&env.new_string(info.virtual_ip.to_string())?.into()),
JValue::Object(&env.new_string(info.virtual_netmask.to_string())?.into()),
JValue::Object(&env.new_string(info.virtual_gateway.to_string())?.into()),
],
)?;
let rs = env.call_method(
&self.this,
"register",
"(Ltop/wherewego/vnt/jni/param/RegisterInfo;)Z",
&[JValue::Object(&param)],
)?;
rs.z()
}
fn error0(&self, info: ErrorInfo) -> jni::errors::Result<()> {
let code: u8 = info.code.into();
let env = &mut self.jvm.attach_current_thread()? as &mut JNIEnv;
let msg = if let Some(msg) = info.msg {
env.new_string(msg)?
} else {
JString::default()
};
let param = env.new_object(
"top/wherewego/vnt/jni/param/ErrorInfo",
"(ILjava/lang/String;)V",
&[JValue::Int(code as _), JValue::Object(&msg.into())],
)?;
env.call_method(
&self.this,
"error",
"(Ltop/wherewego/vnt/jni/param/ErrorInfo;)V",
&[JValue::Object(&param)],
)?;
Ok(())
}
fn stop0(&self) -> jni::errors::Result<()> {
let env = &mut self.jvm.attach_current_thread()? as &mut JNIEnv;
env.call_method(&self.this, "error", "()V", &[])?;
Ok(())
}
}
impl VntCallback for CallBack {
fn success(&self) {
}
fn create_tun(&self, info: DeviceInfo) {
if let Err(e) = self.create_tun0(info) {
log::warn!("create_tun {:?}", e);
}
}
fn connect(&self, info: ConnectInfo) {
if let Err(e) = self.connect0(info) {
log::warn!("connect {:?}", e);
}
}
fn handshake(&self, info: HandshakeInfo) -> bool {
self.handshake0(info).unwrap_or_else(|e| {
log::warn!("handshake {:?}", e);
false
})
}
fn register(&self, info: RegisterInfo) -> bool {
self.register0(info).unwrap_or_else(|e| {
log::warn!("register {:?}", e);
false
})
}
fn error(&self, info: ErrorInfo) {
if let Err(e) = self.error0(info) {
log::warn!("error {:?}", e);
}
}
fn stop(&self) {
if let Err(e) = self.stop0() {
log::warn!("stop {:?}", e);
}
}
}
+149
View File
@@ -0,0 +1,149 @@
use std::net::ToSocketAddrs;
use std::str::FromStr;
use jni::errors::Error;
use jni::objects::JObject;
use jni::JNIEnv;
use vnt::channel::punch::PunchModel;
use vnt::channel::UseChannelType;
use vnt::cipher::CipherModel;
use vnt::core::Config;
use crate::utils::*;
pub fn new_config(env: &mut JNIEnv, config: JObject) -> Result<Config, Error> {
#[cfg(any(target_os = "windows", target_os = "linux"))]
let tap = env.get_field(&config, "tap", "Z")?.z()?;
let token = to_string_not_null(env, &config, "token")?;
let name = to_string_not_null(env, &config, "name")?;
let device_id = to_string_not_null(env, &config, "deviceId")?;
let password = to_string(env, &config, "password")?;
let server_address_str = to_string_not_null(env, &config, "server")?;
let stun_server = to_string_array_not_null(env, &config, "stunServer")?;
let cipher_model = to_string_not_null(env, &config, "cipherModel")?;
let punch_model = to_string(env, &config, "punchModel")?;
let mtu = to_integer(env, &config, "mtu")?.map(|v| v as u32);
let tcp = env.get_field(&config, "tcp", "Z")?.z()?;
let server_encrypt = env.get_field(&config, "serverEncrypt", "Z")?.z()?;
let use_channel = to_string(env, &config, "useChannel")?;
let finger = env.get_field(&config, "finger", "Z")?.z()?;
let first_latency = env.get_field(&config, "firstLatency", "Z")?.z()?;
let in_ips = to_string_array(env, &config, "inIps")?;
let out_ips = to_string_array(env, &config, "outIps")?;
let ports =
to_i32_array(env, &config, "ports")?.map(|v| v.into_iter().map(|v| v as u16).collect());
let ip = if let Some(ip) = to_string(env, &config, "ip")? {
match ip.parse() {
Ok(ip) => Some(ip),
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("ip {} err: {}", ip, e),
)
.expect("throw");
return Err(Error::JavaException);
}
}
} else {
None
};
let in_ips = if let Some(in_ips) = in_ips {
match common::args_parse::ips_parse(&in_ips) {
Ok(in_ips) => in_ips,
Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("in_ips {}", e))
.expect("throw");
return Err(Error::JavaException);
}
}
} else {
vec![]
};
let out_ips = if let Some(out_ips) = out_ips {
match common::args_parse::out_ips_parse(&out_ips) {
Ok(out_ips) => out_ips,
Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("out_ips {}", e))
.expect("throw");
return Err(Error::JavaException);
}
}
} else {
vec![]
};
let server_address = match server_address_str.to_socket_addrs() {
Ok(mut rs) => {
if let Some(addr) = rs.next() {
addr
} else {
env.throw_new("java/lang/RuntimeException", "server address err")
.expect("throw");
return Err(Error::JavaException);
}
}
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("server address {}", e),
)
.expect("throw");
return Err(Error::JavaException);
}
};
let cipher_model = match CipherModel::from_str(&cipher_model) {
Ok(cipher_model) => cipher_model,
Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("cipher_model {}", e))
.expect("throw");
return Err(Error::JavaException);
}
};
#[cfg(not(target_os = "android"))]
let device_name = to_string(env, &config, "deviceName")?;
#[cfg(target_os = "android")]
let device_fd = env.get_field(&config, "deviceFd", "I")?.i()? as i32;
let config = match Config::new(
#[cfg(any(target_os = "windows", target_os = "linux"))]
tap,
token,
device_id,
name,
server_address,
server_address_str,
stun_server,
in_ips,
out_ips,
password,
mtu,
tcp,
ip,
false,
server_encrypt,
1,
cipher_model,
finger,
PunchModel::from_str(&punch_model.unwrap_or_default()).unwrap_or_default(),
ports,
first_latency,
#[cfg(not(target_os = "android"))]
device_name,
#[cfg(target_os = "android")]
device_fd,
UseChannelType::from_str(&use_channel.unwrap_or_default()).unwrap_or_default(),
None,
0,
) {
Ok(config) => config,
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt start error {}", e),
)
.expect("throw");
return Err(Error::JavaException);
}
};
Ok(config)
}
+3 -1
View File
@@ -1,2 +1,4 @@
pub mod callback;
pub mod config;
pub mod utils;
pub mod vnt;
pub mod vnt_util;
+121
View File
@@ -0,0 +1,121 @@
use jni::errors::Error;
use jni::objects::{JIntArray, JObject, JObjectArray, JString};
use jni::JNIEnv;
pub fn to_string_not_null(
env: &mut JNIEnv,
config: &JObject,
name: &'static str,
) -> Result<String, Error> {
let value = env.get_field(config, name, "Ljava/lang/String;")?.l()?;
if value.is_null() {
env.throw_new("java/lang/NullPointerException", name)
.expect("throw");
return Err(Error::NullPtr(name));
}
let binding = JString::from(value);
let value = env.get_string(binding.as_ref())?;
match value.to_str() {
Ok(value) => Ok(value.to_string()),
Err(_) => {
env.throw_new("java/lang/RuntimeException", "not utf-8")
.expect("throw");
return Err(Error::JavaException);
}
}
}
pub fn to_string(env: &mut JNIEnv, config: &JObject, name: &str) -> Result<Option<String>, Error> {
let value = env.get_field(config, name, "Ljava/lang/String;")?.l()?;
if value.is_null() {
return Ok(None);
}
let tmp = JString::from(value);
let value = env.get_string(tmp.as_ref())?;
match value.to_str() {
Ok(value) => Ok(Some(value.to_string())),
Err(_) => {
env.throw_new("java/lang/RuntimeException", "not utf-8")
.expect("throw");
return Err(Error::JavaException);
}
}
}
pub fn to_string_array_not_null(
env: &mut JNIEnv,
config: &JObject,
name: &str,
) -> Result<Vec<String>, Error> {
match to_string_array(env, config, name)? {
None => {
env.throw_new("java/lang/NullPointerException", name)
.expect("throw");
return Err(Error::JavaException);
}
Some(rs) => Ok(rs),
}
}
pub fn to_string_array(
env: &mut JNIEnv,
config: &JObject,
name: &str,
) -> Result<Option<Vec<String>>, Error> {
let value = env.get_field(config, name, "[Ljava/lang/String;")?.l()?;
if value.is_null() {
return Ok(None);
}
let arr = JObjectArray::from(value);
let len = env.get_array_length(&arr)?;
let mut rs = Vec::with_capacity(len as usize);
for index in 0..len {
let object = env.get_object_array_element(&arr, index)?;
if object.is_null() {
env.throw_new(
"java/lang/NullPointerException",
format!("{},index={}", name, index),
)
.expect("throw");
return Err(Error::JavaException);
}
match env.get_string(JString::from(object).as_ref())?.to_str() {
Ok(value) => {
rs.push(value.to_string());
}
Err(_) => {
env.throw_new("java/lang/RuntimeException", "not utf-8")
.expect("throw");
return Err(Error::JavaException);
}
}
}
Ok(Some(rs))
}
pub fn to_i32_array(
env: &mut JNIEnv,
config: &JObject,
name: &str,
) -> Result<Option<Vec<i32>>, Error> {
let obj = env.get_field(&config, name, "[I")?.l()?;
if obj.is_null() {
Ok(None)
} else {
let j_arr = JIntArray::from(obj);
let len = env.get_array_length(&j_arr)?;
let mut arr = vec![0i32; len as usize];
env.get_int_array_region(j_arr, 0, &mut arr)?;
Ok(Some(arr))
}
}
pub fn to_integer(env: &mut JNIEnv, config: &JObject, name: &str) -> Result<Option<i32>, Error> {
let value = env.get_field(config, name, "Ljava/lang/Integer;")?.l()?;
if value.is_null() {
return Ok(None);
}
// 调用 intValue
return Ok(Some(
env.call_method(value, "intValue", "()I", &[])?.i()? as _
));
}
+51 -25
View File
@@ -1,45 +1,71 @@
use std::ptr;
use jni::errors::Error;
use jni::objects::{JClass, JObject, JValue};
use jni::sys::{jboolean, jbyte, jint, jlong, jobject, jobjectArray, jsize};
use jni::sys::{jbyte, jint, jlong, jobject, jobjectArray, jsize};
use jni::JNIEnv;
use std::ptr;
use vnt::channel::Route;
use vnt::core::sync::VntSync;
use vnt::core::Vnt;
use vnt::handle::PeerDeviceInfo;
use crate::callback::CallBack;
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_new0(
mut env: JNIEnv<'static>,
_class: JClass,
config: JObject,
call_back: JObject<'static>,
) -> jlong {
let jvm = if let Ok(jvm) = env.get_java_vm() {
jvm
} else {
return 0;
};
match crate::config::new_config(&mut env, config) {
Ok(config) => {
let call_back = if let Ok(call_back) = env.new_global_ref(call_back) {
call_back
} else {
return 0;
};
let vnt_util = match Vnt::new(config, CallBack::new(jvm, call_back)) {
Ok(vnt_util) => vnt_util,
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt start error {}", e),
)
.expect("throw");
return 0;
}
};
let ptr = Box::into_raw(Box::new(vnt_util));
return ptr as jlong;
}
Err(_) => {}
}
return 0;
}
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_stop0(
_env: JNIEnv,
_class: JClass,
raw_vnt: jlong,
) {
let vnt = raw_vnt as *mut VntSync;
let vnt = raw_vnt as *mut Vnt;
let _ = (&*vnt).stop();
}
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_waitStop0(
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_wait0(
_env: JNIEnv,
_class: JClass,
raw_vnt: jlong,
) {
let vnt = raw_vnt as *mut VntSync;
let _ = (&mut *vnt).wait_stop();
}
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_waitStopMs0(
_env: JNIEnv,
_class: JClass,
raw_vnt: jlong,
ms: jlong,
) -> jboolean {
let vnt = raw_vnt as *mut VntSync;
if (&mut *vnt).wait_stop_ms(ms as _) {
jni::sys::JNI_TRUE
} else {
jni::sys::JNI_FALSE
}
let vnt = raw_vnt as *mut Vnt;
let _ = (&*vnt).wait();
}
#[no_mangle]
@@ -48,7 +74,7 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_drop0(
_class: JClass,
raw_vnt: jlong,
) {
let vnt = raw_vnt as *mut VntSync;
let vnt = raw_vnt as *mut Vnt;
let _ = Box::from_raw(vnt).stop();
}
@@ -58,7 +84,7 @@ pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_Vnt_list0(
_class: JClass,
raw_vnt: jlong,
) -> jobjectArray {
let vnt = raw_vnt as *mut VntSync;
let vnt = raw_vnt as *mut Vnt;
let vnt = &mut *vnt;
let list = vnt.device_list();
-371
View File
@@ -1,371 +0,0 @@
use std::net::ToSocketAddrs;
use std::ptr;
use std::str::FromStr;
use jni::errors::Error;
use jni::objects::{JClass, JObject, JString, JValue};
#[cfg(not(target_os = "android"))]
use jni::sys::jboolean;
use jni::sys::{jint, jlong, jobject};
use jni::JNIEnv;
use vnt::channel::punch::PunchModel;
use vnt::cipher::CipherModel;
use vnt::core::sync::VntUtilSync;
use vnt::core::Config;
use vnt::handle::registration_handler::{RegResponse, ReqEnum};
#[cfg(not(target_os = "android"))]
use vnt::tun_tap_device::DriverInfo;
fn to_string_not_null(
env: &mut JNIEnv,
config: &JObject,
name: &'static str,
) -> Result<String, Error> {
let value = env.get_field(config, name, "Ljava/lang/String;")?.l()?;
if value.is_null() {
env.throw_new("java/lang/NullPointerException", name)
.expect("throw");
return Err(Error::NullPtr(name));
}
let binding = JString::from(value);
let value = env.get_string(binding.as_ref())?;
match value.to_str() {
Ok(value) => Ok(value.to_string()),
Err(_) => {
env.throw_new("java/lang/RuntimeException", "not utf-8")
.expect("throw");
return Err(Error::JavaException);
}
}
}
fn to_string(env: &mut JNIEnv, config: &JObject, name: &str) -> Result<Option<String>, Error> {
let value = env.get_field(config, name, "Ljava/lang/String;")?.l()?;
if value.is_null() {
return Ok(None);
}
let tmp = JString::from(value);
let value = env.get_string(tmp.as_ref())?;
match value.to_str() {
Ok(value) => Ok(Some(value.to_string())),
Err(_) => {
env.throw_new("java/lang/RuntimeException", "not utf-8")
.expect("throw");
return Err(Error::JavaException);
}
}
}
fn new_sync(env: &mut JNIEnv, config: JObject) -> Result<VntUtilSync, Error> {
let token = to_string_not_null(env, &config, "token")?;
let name = to_string_not_null(env, &config, "name")?;
let device_id = to_string_not_null(env, &config, "deviceId")?;
let password = to_string(env, &config, "password")?;
let server_address_str = to_string_not_null(env, &config, "server")?;
let stun_server_str = to_string_not_null(env, &config, "stunServer")?;
let cipher_model = to_string_not_null(env, &config, "cipherModel")?;
let tcp = env.get_field(&config, "tcp", "Z")?.z()?;
let finger = env.get_field(&config, "finger", "Z")?.z()?;
let in_ips = to_string(env, &config, "inIps")?;
let out_ips = to_string(env, &config, "outIps")?;
let in_ips = if let Some(in_ips) = in_ips {
let in_ips: Vec<&str> = in_ips.split("\n").collect();
let in_ips = in_ips.iter().map(|v| v.to_string()).collect();
match common::args_parse::ips_parse(&in_ips) {
Ok(in_ips) => in_ips,
Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("in_ips {}", e))
.expect("throw");
return Err(Error::JavaException);
}
}
} else {
vec![]
};
let out_ips = if let Some(out_ips) = out_ips {
let out_ips: Vec<&str> = out_ips.split("\n").collect();
let out_ips = out_ips.iter().map(|v| v.to_string()).collect();
match common::args_parse::out_ips_parse(&out_ips) {
Ok(out_ips) => out_ips,
Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("out_ips {}", e))
.expect("throw");
return Err(Error::JavaException);
}
}
} else {
vec![]
};
let server_address = match server_address_str.to_socket_addrs() {
Ok(mut rs) => {
if let Some(addr) = rs.next() {
addr
} else {
env.throw_new("java/lang/RuntimeException", "server address err")
.expect("throw");
return Err(Error::JavaException);
}
}
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("server address {}", e),
)
.expect("throw");
return Err(Error::JavaException);
}
};
let cipher_model = match CipherModel::from_str(&cipher_model) {
Ok(cipher_model) => cipher_model,
Err(e) => {
env.throw_new("java/lang/RuntimeException", format!("cipher_model {}", e))
.expect("throw");
return Err(Error::JavaException);
}
};
let mut stun_server = Vec::new();
for addr in stun_server_str.split(",") {
stun_server.push(addr.trim().to_string());
}
let config = Config::new(
false,
token,
device_id,
name,
server_address,
server_address_str,
stun_server,
in_ips,
out_ips,
password,
false,
None,
tcp,
None,
false,
false,
1,
cipher_model,
finger,
PunchModel::All,
);
match VntUtilSync::new(config) {
Ok(vnt_util) => Ok(vnt_util),
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt start error {}", e),
)
.expect("throw");
return Err(Error::JavaException);
}
}
}
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_new0(
mut env: JNIEnv,
_class: JClass,
config: JObject,
) -> jlong {
match new_sync(&mut env, config) {
Ok(vnt_util) => {
let ptr = Box::into_raw(Box::new(vnt_util));
return ptr as jlong;
}
Err(_) => {}
}
return 0;
}
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_connect0(
mut env: JNIEnv,
_class: JClass,
raw_vnt_util: jlong,
) {
let raw_vnt_util = raw_vnt_util as *mut VntUtilSync;
match (&mut *raw_vnt_util).connect() {
Ok(_) => {}
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt connect error {}", e),
)
.expect("throw");
}
}
}
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_register0(
mut env: JNIEnv,
_class: JClass,
raw_vnt_util: jlong,
) -> jobject {
let raw_vnt_util = raw_vnt_util as *mut VntUtilSync;
match (&mut *raw_vnt_util).register() {
Ok(response) => match reg_response(&mut env, response) {
Ok(res) => {
return res;
}
Err(e) => {
env.throw(format!("vnt register error {}", e))
.expect("throw");
}
},
Err(e) => match e {
ReqEnum::TokenError => {
env.throw_new(
"top/wherewego/vnt/jni/exception/TokenErrorException",
"TokenError",
)
.expect("throw");
}
ReqEnum::AddressExhausted => {
env.throw_new(
"top/wherewego/vnt/jni/exception/AddressExhaustedException",
"AddressExhausted",
)
.expect("throw");
}
ReqEnum::Timeout => {
env.throw_new(
"top/wherewego/vnt/jni/exception/TimeoutException",
"Timeout",
)
.expect("throw");
}
ReqEnum::ServerError(str) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt register error {}", str),
)
.expect("throw");
}
ReqEnum::Other(str) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt register error {}", str),
)
.expect("throw");
}
ReqEnum::IpAlreadyExists => {
env.throw_new(
"top/wherewego/vnt/jni/exception/IpAlreadyExistsException",
"IpAlreadyExists",
)
.expect("throw");
}
ReqEnum::InvalidIp => {
env.throw_new(
"top/wherewego/vnt/jni/exception/InvalidIpException",
"InvalidIp",
)
.expect("throw");
}
},
}
return ptr::null_mut();
}
#[cfg(target_os = "android")]
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_createIface0(
_env: JNIEnv,
_class: JClass,
raw_vnt_util: jlong,
fd: jint,
) {
let raw_vnt_util = raw_vnt_util as *mut VntUtilSync;
(&mut *raw_vnt_util).create_iface(fd as i32);
}
#[cfg(not(target_os = "android"))]
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_createIface0(
mut env: JNIEnv,
_class: JClass,
raw_vnt_util: jlong,
) -> jobject {
let raw_vnt_util = raw_vnt_util as *mut VntUtilSync;
let rs = (&mut *raw_vnt_util).create_iface();
match rs {
Ok(driver_info) => match driver_info_e(&mut env, driver_info) {
Ok(res) => {
return res;
}
Err(e) => {
env.throw(format!("vnt create iface error {}", e))
.expect("throw");
}
},
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt create iface error {}", e),
)
.expect("throw");
}
}
return ptr::null_mut();
}
#[no_mangle]
pub unsafe extern "C" fn Java_top_wherewego_vnt_jni_VntUtil_build0(
mut env: JNIEnv,
_class: JClass,
raw_vnt_util: jlong,
) -> jlong {
let raw_vnt_util = Box::from_raw(raw_vnt_util as *mut VntUtilSync);
match raw_vnt_util.build() {
Ok(rs) => {
return Box::into_raw(Box::new(rs)) as jlong;
}
Err(e) => {
env.throw_new(
"java/lang/RuntimeException",
format!("vnt start error:{:?}", e),
)
.expect("throw");
}
}
return 0;
}
fn reg_response(env: &mut JNIEnv, response: RegResponse) -> Result<jobject, Error> {
let virtual_ip = u32::from(response.virtual_ip);
let virtual_gateway = u32::from(response.virtual_gateway);
let virtual_netmask = u32::from(response.virtual_netmask);
let response = env.new_object(
"top/wherewego/vnt/jni/RegResponse",
"(III)V",
&[
JValue::Int(virtual_ip as jint),
JValue::Int(virtual_gateway as jint),
JValue::Int(virtual_netmask as jint),
],
)?;
Ok(response.into_raw())
}
#[cfg(not(target_os = "android"))]
fn driver_info_e(env: &mut JNIEnv, driver_info: DriverInfo) -> Result<jobject, Error> {
let is_tun = driver_info.device_type.is_tun();
let name = driver_info.name;
let version = driver_info.version;
let mac = driver_info.mac.unwrap_or(String::new());
let response = env.new_object(
"top/wherewego/vnt/jni/DriverInfo",
"(ZLjava/lang/String;Ljava/lang/String;Ljava/lang/String;)V",
&[
JValue::Bool(is_tun as jboolean),
JValue::Object(&env.new_string(name)?.into()),
JValue::Object(&env.new_string(version)?.into()),
JValue::Object(&env.new_string(mac)?.into()),
],
)?;
Ok(response.into_raw())
}
+19 -20
View File
@@ -1,44 +1,39 @@
[package]
name = "vnt"
version = "1.2.3"
version = "1.2.9"
edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
[dependencies]
tun= {path = "tun"}
packet = { path = "./packet" }
bytes = "1.3.0"
bytes = "1.5.0"
log = "0.4.17"
libc = "0.2.137"
crossbeam-utils = "0.8"
crossbeam-epoch = "0.9.15"
dashmap = "5.5.1"
parking_lot = "0.12.1"
byte-pool = "0.2.4"
lazy_static = "1.4.0"
rand = "0.8.5"
sha2 = { version = "0.10.6", features = ["oid"] }
thiserror = "1.0.37"
protobuf = "3.2.0"
socket2 = { version = "0.5.2", features = ["all"] }
tokio = { version = "1.32.0", features = ["full"] }
aes-gcm = { version = "0.10.2" }
ring = { version = "0.16.20", optional = true }
cbc = "0.1.2"
ecb = "0.1.2"
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}
aes = "0.8.3"
stun-format = { version = "1.0.1", features = ["fmt", "rfc3489"] }
rsa = { version = "0.7.2", features = [] }
spki = { version = "0.6.0", features = ["fingerprint", "alloc"] }
rsa = { version = "0.9.2", features = [] ,optional = true}
spki = { version = "0.7.2", features = ["fingerprint", "alloc","base64"] ,optional = true}
openssl-sys = { git = "https://github.com/lbl8603/rust-openssl" ,optional = true}
libsm = {git="https://github.com/lbl8603/libsm" ,optional = true}
[target.'cfg(any(target_os = "linux",target_os = "macos"))'.dependencies]
tun = { path = "./rust-tun" }
mio = {version = "0.8.10",features = ["os-poll","net"]}
[target.'cfg(target_os = "windows")'.dependencies]
win-tun-tap = { path = "./win-tun-tap" }
libloading = "0.7.4"
libloading = "0.8.0"
[build-dependencies]
@@ -46,10 +41,14 @@ protobuf-codegen = "3.2.0"
protoc-bin-vendored = "3.0.0"
[features]
default = []
default = ["server_encrypt","aes_gcm","aes_cbc","aes_ecb","sm4_cbc","ip_proxy"]
openssl = ["openssl-sys"]
# 从源码编译
openssl-vendored = ["openssl-sys/vendored"]
ring-cipher = ["ring"]
aes_cbc=["cbc"]
aes_ecb=["ecb"]
sm4_cbc=["libsm"]
aes_gcm=["aes-gcm"]
server_encrypt =["aes-gcm","rsa","spki"]
ip_proxy=[]
+8 -1
View File
@@ -76,7 +76,14 @@ impl<B: AsRef<[u8]>> TcpPacket<B> {
Ok(packet)
}
}
impl<B: AsRef<[u8]> + AsMut<[u8]>> TcpPacket<B> {
pub fn set_source_ip(&mut self, value: Ipv4Addr) {
self.source_ip = value;
}
pub fn set_destination_ip(&mut self, value: Ipv4Addr) {
self.destination_ip = value;
}
}
impl<B: AsRef<[u8]> + AsMut<[u8]>> TcpPacket<B> {
fn set_checksum(&mut self, value: u16) {
self.buffer.as_mut()[16..18].copy_from_slice(&value.to_be_bytes())
+65 -50
View File
@@ -1,63 +1,78 @@
syntax = "proto3";
message HandshakeRequest{
string version = 1;
bool secret = 2;
message HandshakeRequest {
string version = 1;
bool secret = 2;
}
message HandshakeResponse{
string version = 1;
bool secret = 2;
bytes public_key = 3;
string key_finger = 4;
message HandshakeResponse {
string version = 1;
bool secret = 2;
bytes public_key = 3;
string key_finger = 4;
}
message SecretHandshakeRequest{
string token = 1;
bytes key = 2;
message SecretHandshakeRequest {
string token = 1;
bytes key = 2;
}
message RegistrationRequest{
string token = 1;
string device_id = 2;
string name = 3;
bool is_fast = 4;
string version = 5;
fixed32 virtual_ip = 6;
bool allow_ip_change = 7;
bool client_secret = 8;
message RegistrationRequest {
string token = 1;
string device_id = 2;
string name = 3;
bool is_fast = 4;
string version = 5;
fixed32 virtual_ip = 6;
bool allow_ip_change = 7;
bool client_secret = 8;
}
message RegistrationResponse{
fixed32 virtual_ip = 1;
fixed32 virtual_gateway = 2;
fixed32 virtual_netmask = 3;
uint32 epoch = 4;
repeated DeviceInfo device_info_list = 5;
fixed32 public_ip = 6;
uint32 public_port = 7;
bytes public_ipv6 = 8;
message RegistrationResponse {
fixed32 virtual_ip = 1;
fixed32 virtual_gateway = 2;
fixed32 virtual_netmask = 3;
uint32 epoch = 4;
repeated DeviceInfo device_info_list = 5;
fixed32 public_ip = 6;
uint32 public_port = 7;
bytes public_ipv6 = 8;
}
message DeviceInfo{
string name = 1;
fixed32 virtual_ip = 2;
uint32 device_status = 3;
bool client_secret = 4;
message DeviceInfo {
string name = 1;
fixed32 virtual_ip = 2;
uint32 device_status = 3;
bool client_secret = 4;
}
message DeviceList{
uint32 epoch = 1;
repeated DeviceInfo device_info_list = 2;
message DeviceList {
uint32 epoch = 1;
repeated DeviceInfo device_info_list = 2;
}
message PunchInfo{
repeated fixed32 public_ip_list = 2;
uint32 public_port = 3;
uint32 public_port_range = 4;
PunchNatType nat_type = 5;
bool reply = 6;
fixed32 local_ip = 7;
uint32 local_port = 8;
bytes ipv6 = 9;
uint32 ipv6_port = 10;
message PunchInfo {
repeated fixed32 public_ip_list = 2;
uint32 public_port = 3;
uint32 public_port_range = 4;
PunchNatType nat_type = 5;
bool reply = 6;
fixed32 local_ip = 7;
uint32 local_port = 8;
bytes ipv6 = 9;
uint32 ipv6_port = 10;
uint32 tcp_port = 11;
repeated uint32 udp_ports = 12;
repeated uint32 public_ports = 13;
}
enum PunchNatType{
Symmetric = 0;
Cone = 1;
enum PunchNatType {
Symmetric = 0;
Cone = 1;
}
/// 向服务器上报客户端状态信息
message ClientStatusInfo {
fixed32 source = 1;
repeated RouteItem p2p_list = 2;
uint64 up_stream = 3;
uint64 down_stream = 4;
PunchNatType nat_type = 5;
}
message RouteItem {
fixed32 next_ip = 1;
}
-25
View File
@@ -1,25 +0,0 @@
[package]
name = "tun"
version = "0.5.4"
edition = "2018"
authors = ["meh. <[email protected]>"]
license = "WTFPL"
description = "TUN device creation and handling."
repository = "https://github.com/meh/rust-tun"
keywords = ["tun", "network", "tunnel", "bindings"]
[dependencies]
libc = "0.2"
thiserror = "1"
[target.'cfg(any(target_os = "linux", target_os = "macos", target_os = "ios", target_os = "android"))'.dependencies]
bytes = { version = "1", optional = true }
byteorder = { version = "1", optional = true }
[target.'cfg(any(target_os = "linux", target_os = "macos"))'.dependencies]
ioctl = { version = "0.6", package = "ioctl-sys" }
-128
View File
@@ -1,128 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::net::{IpAddr, Ipv4Addr};
use std::net::{SocketAddr, SocketAddrV4};
use crate::error::*;
/// Helper trait to convert things into IPv4 addresses.
#[allow(clippy::wrong_self_convention)]
pub trait IntoAddress {
/// Convert the type to an `Ipv4Addr`.
fn into_address(&self) -> Result<Ipv4Addr>;
}
impl IntoAddress for u32 {
fn into_address(&self) -> Result<Ipv4Addr> {
Ok(Ipv4Addr::new(
((*self) & 0xff) as u8,
((*self >> 8) & 0xff) as u8,
((*self >> 16) & 0xff) as u8,
((*self >> 24) & 0xff) as u8,
))
}
}
impl IntoAddress for i32 {
fn into_address(&self) -> Result<Ipv4Addr> {
(*self as u32).into_address()
}
}
impl IntoAddress for (u8, u8, u8, u8) {
fn into_address(&self) -> Result<Ipv4Addr> {
Ok(Ipv4Addr::new(self.0, self.1, self.2, self.3))
}
}
impl IntoAddress for str {
fn into_address(&self) -> Result<Ipv4Addr> {
self.parse().map_err(|_| Error::InvalidAddress)
}
}
impl<'a> IntoAddress for &'a str {
fn into_address(&self) -> Result<Ipv4Addr> {
(*self).into_address()
}
}
impl IntoAddress for String {
fn into_address(&self) -> Result<Ipv4Addr> {
(&**self).into_address()
}
}
impl<'a> IntoAddress for &'a String {
fn into_address(&self) -> Result<Ipv4Addr> {
(&**self).into_address()
}
}
impl IntoAddress for Ipv4Addr {
fn into_address(&self) -> Result<Ipv4Addr> {
Ok(*self)
}
}
impl<'a> IntoAddress for &'a Ipv4Addr {
fn into_address(&self) -> Result<Ipv4Addr> {
(&**self).into_address()
}
}
impl IntoAddress for IpAddr {
fn into_address(&self) -> Result<Ipv4Addr> {
match *self {
IpAddr::V4(ref value) => Ok(*value),
IpAddr::V6(_) => Err(Error::InvalidAddress),
}
}
}
impl<'a> IntoAddress for &'a IpAddr {
fn into_address(&self) -> Result<Ipv4Addr> {
(&**self).into_address()
}
}
impl IntoAddress for SocketAddrV4 {
fn into_address(&self) -> Result<Ipv4Addr> {
Ok(*self.ip())
}
}
impl<'a> IntoAddress for &'a SocketAddrV4 {
fn into_address(&self) -> Result<Ipv4Addr> {
(&**self).into_address()
}
}
impl IntoAddress for SocketAddr {
fn into_address(&self) -> Result<Ipv4Addr> {
match *self {
SocketAddr::V4(ref value) => Ok(*value.ip()),
SocketAddr::V6(_) => Err(Error::InvalidAddress),
}
}
}
impl<'a> IntoAddress for &'a SocketAddr {
fn into_address(&self) -> Result<Ipv4Addr> {
(&**self).into_address()
}
}
-126
View File
@@ -1,126 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::net::Ipv4Addr;
use std::os::unix::io::RawFd;
use crate::address::IntoAddress;
use crate::platform;
/// TUN interface OSI layer of operation.
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Layer {
L2,
L3,
}
impl Default for Layer {
fn default() -> Self {
Layer::L3
}
}
/// Configuration builder for a TUN interface.
#[derive(Clone, Default, Debug)]
pub struct Configuration {
pub(crate) name: Option<String>,
pub(crate) platform: platform::Configuration,
pub(crate) address: Option<Ipv4Addr>,
pub(crate) destination: Option<Ipv4Addr>,
pub(crate) broadcast: Option<Ipv4Addr>,
pub(crate) netmask: Option<Ipv4Addr>,
pub(crate) mtu: Option<i32>,
pub(crate) enabled: Option<bool>,
pub(crate) layer: Option<Layer>,
pub(crate) queues: Option<usize>,
pub(crate) raw_fd: Option<RawFd>,
}
impl Configuration {
/// Access the platform dependant configuration.
pub fn platform<F>(&mut self, f: F) -> &mut Self
where
F: FnOnce(&mut platform::Configuration),
{
f(&mut self.platform);
self
}
/// Set the name.
pub fn name<S: AsRef<str>>(&mut self, name: S) -> &mut Self {
self.name = Some(name.as_ref().into());
self
}
/// Set the address.
pub fn address<A: IntoAddress>(&mut self, value: A) -> &mut Self {
self.address = Some(value.into_address().unwrap());
self
}
/// Set the destination address.
pub fn destination<A: IntoAddress>(&mut self, value: A) -> &mut Self {
self.destination = Some(value.into_address().unwrap());
self
}
/// Set the broadcast address.
pub fn broadcast<A: IntoAddress>(&mut self, value: A) -> &mut Self {
self.broadcast = Some(value.into_address().unwrap());
self
}
/// Set the netmask.
pub fn netmask<A: IntoAddress>(&mut self, value: A) -> &mut Self {
self.netmask = Some(value.into_address().unwrap());
self
}
/// Set the MTU.
pub fn mtu(&mut self, value: i32) -> &mut Self {
self.mtu = Some(value);
self
}
/// Set the interface to be enabled once created.
pub fn up(&mut self) -> &mut Self {
self.enabled = Some(true);
self
}
/// Set the interface to be disabled once created.
pub fn down(&mut self) -> &mut Self {
self.enabled = Some(false);
self
}
/// Set the OSI layer of operation.
pub fn layer(&mut self, value: Layer) -> &mut Self {
self.layer = Some(value);
self
}
/// Set the number of queues.
pub fn queues(&mut self, value: usize) -> &mut Self {
self.queues = Some(value);
self
}
/// Set the raw fd.
pub fn raw_fd(&mut self, fd: RawFd) -> &mut Self {
self.raw_fd = Some(fd);
self
}
}
-94
View File
@@ -1,94 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::net::Ipv4Addr;
use crate::configuration::Configuration;
use crate::error::*;
/// A TUN device.
pub trait Device {
type Queue;
/// Reconfigure the device.
fn configure(&mut self, config: &Configuration) -> Result<()> {
if let Some(ip) = config.address {
self.set_address(ip)?;
}
if let Some(ip) = config.destination {
self.set_destination(ip)?;
}
if let Some(ip) = config.broadcast {
self.set_broadcast(ip)?;
}
if let Some(ip) = config.netmask {
self.set_netmask(ip)?;
}
if let Some(mtu) = config.mtu {
self.set_mtu(mtu)?;
}
if let Some(enabled) = config.enabled {
self.enabled(enabled)?;
}
Ok(())
}
/// Get the device name.
fn name(&self) -> &str;
/// Set the device name.
fn set_name(&mut self, name: &str) -> Result<()>;
/// Turn on or off the interface.
fn enabled(&mut self, value: bool) -> Result<()>;
/// Get the address.
fn address(&self) -> Result<Ipv4Addr>;
/// Set the address.
fn set_address(&mut self, value: Ipv4Addr) -> Result<()>;
/// Get the destination address.
fn destination(&self) -> Result<Ipv4Addr>;
/// Set the destination address.
fn set_destination(&mut self, value: Ipv4Addr) -> Result<()>;
/// Get the broadcast address.
fn broadcast(&self) -> Result<Ipv4Addr>;
/// Set the broadcast address.
fn set_broadcast(&mut self, value: Ipv4Addr) -> Result<()>;
/// Get the netmask.
fn netmask(&self) -> Result<Ipv4Addr>;
/// Set the netmask.
fn set_netmask(&mut self, value: Ipv4Addr) -> Result<()>;
/// Get the MTU.
fn mtu(&self) -> Result<i32>;
/// Set the MTU.
fn set_mtu(&mut self, value: i32) -> Result<()>;
/// Get a device queue.
fn queue(&self, index: usize) -> Option<&Self::Queue>;
}
-54
View File
@@ -1,54 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::{ffi, io, num};
use thiserror::Error;
#[derive(Error, Debug)]
pub enum Error {
#[error("invalid configuration")]
InvalidConfig,
#[error("not implementated")]
NotImplemented,
#[error("device name too long")]
NameTooLong,
#[error("invalid device name")]
InvalidName,
#[error("invalid address")]
InvalidAddress,
#[error("invalid file descriptor")]
InvalidDescriptor,
#[error("unsuported network layer of operation")]
UnsupportedLayer,
#[error("invalid queues number")]
InvalidQueuesNumber,
#[error(transparent)]
Io(#[from] io::Error),
#[error(transparent)]
Nul(#[from] ffi::NulError),
#[error(transparent)]
ParseNum(#[from] num::ParseIntError),
}
pub type Result<T> = ::std::result::Result<T, Error>;
-32
View File
@@ -1,32 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
#![cfg(unix)]
mod error;
pub use crate::error::*;
mod address;
pub use crate::address::IntoAddress;
mod device;
pub use crate::device::Device;
mod configuration;
pub use crate::configuration::{Configuration, Layer};
pub mod platform;
pub use crate::platform::create;
pub fn configure() -> Configuration {
Configuration::default()
}
-384
View File
@@ -1,384 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::ffi::{CStr, CString};
use std::io;
use std::mem;
use std::net::Ipv4Addr;
use std::os::unix::io::AsRawFd;
use std::ptr;
use std::sync::Arc;
use std::vec::Vec;
use libc;
use libc::{c_char, c_short};
use libc::{AF_INET, O_RDWR, SOCK_DGRAM};
use crate::configuration::{Configuration, Layer};
use crate::device::Device as D;
use crate::error::*;
use crate::platform::linux::sys::*;
use crate::platform::posix::{self, Fd, SockAddr};
/// A TUN device using the TUN/TAP Linux driver.
pub struct Device {
name: String,
queues: Vec<Queue>,
ctl: Fd,
}
impl Device {
/// Create a new `Device` for the given `Configuration`.
pub fn new(config: &Configuration) -> Result<Self> {
let mut device = unsafe {
let dev = match config.name.as_ref() {
Some(name) => {
let name = CString::new(name.clone())?;
if name.as_bytes_with_nul().len() > IFNAMSIZ {
return Err(Error::NameTooLong);
}
Some(name)
}
None => None,
};
let mut queues = Vec::new();
let mut req: ifreq = mem::zeroed();
if let Some(dev) = dev.as_ref() {
ptr::copy_nonoverlapping(
dev.as_ptr() as *const c_char,
req.ifrn.name.as_mut_ptr(),
dev.as_bytes().len(),
);
}
let device_type: c_short = config.layer.unwrap_or(Layer::L3).into();
let queues_num = config.queues.unwrap_or(1);
if queues_num < 1 {
return Err(Error::InvalidQueuesNumber);
}
req.ifru.flags = device_type
| if config.platform.packet_information {
0
} else {
IFF_NO_PI
}
| if queues_num > 1 { IFF_MULTI_QUEUE } else { 0 };
for _ in 0..queues_num {
let tun = Fd::new(libc::open(b"/dev/net/tun\0".as_ptr() as *const _, O_RDWR))
.map_err(|_| io::Error::last_os_error())?;
if tunsetiff(tun.0, &mut req as *mut _ as *mut _) < 0 {
return Err(io::Error::last_os_error().into());
}
queues.push(Queue {
tun: Arc::new(tun),
pi_enabled: config.platform.packet_information,
});
}
let ctl = Fd::new(libc::socket(AF_INET, SOCK_DGRAM, 0))
.map_err(|_| io::Error::last_os_error())?;
Device {
name: CStr::from_ptr(req.ifrn.name.as_ptr())
.to_string_lossy()
.into(),
queues,
ctl,
}
};
device.configure(config)?;
Ok(device)
}
/// Prepare a new request.
unsafe fn request(&self) -> ifreq {
let mut req: ifreq = mem::zeroed();
ptr::copy_nonoverlapping(
self.name.as_ptr() as *const c_char,
req.ifrn.name.as_mut_ptr(),
self.name.len(),
);
req
}
// /// Make the device persistent.
// pub fn persist(&mut self) -> Result<()> {
// unsafe {
// if tunsetpersist(self.as_raw_fd(), &1) < 0 {
// Err(io::Error::last_os_error().into())
// } else {
// Ok(())
// }
// }
// }
// /// Set the owner of the device.
// pub fn user(&mut self, value: i32) -> Result<()> {
// unsafe {
// if tunsetowner(self.as_raw_fd(), &value) < 0 {
// Err(io::Error::last_os_error().into())
// } else {
// Ok(())
// }
// }
// }
//
// /// Set the group of the device.
// pub fn group(&mut self, value: i32) -> Result<()> {
// unsafe {
// if tunsetgroup(self.as_raw_fd(), &value) < 0 {
// Err(io::Error::last_os_error().into())
// } else {
// Ok(())
// }
// }
// }
/// Return whether the device has packet information
pub fn has_packet_information(&self) -> bool {
self.queues[0].has_packet_information()
}
/// Set non-blocking mode
pub fn set_nonblock(&self) -> io::Result<()> {
self.queues[0].set_nonblock()
}
}
impl D for Device {
type Queue = Queue;
fn name(&self) -> &str {
&self.name
}
fn set_name(&mut self, value: &str) -> Result<()> {
unsafe {
let name = CString::new(value)?;
if name.as_bytes_with_nul().len() > IFNAMSIZ {
return Err(Error::NameTooLong);
}
let mut req = self.request();
ptr::copy_nonoverlapping(
name.as_ptr() as *const c_char,
req.ifru.newname.as_mut_ptr(),
value.len(),
);
if siocsifname(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
self.name = value.into();
Ok(())
}
}
fn enabled(&mut self, value: bool) -> Result<()> {
unsafe {
let mut req = self.request();
if siocgifflags(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
if value {
req.ifru.flags |= IFF_UP | IFF_RUNNING;
} else {
req.ifru.flags &= !IFF_UP;
}
if siocsifflags(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn address(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifru.addr).map(Into::into)
}
}
fn set_address(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.addr = SockAddr::from(value).into();
if siocsifaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn destination(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifdstaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifru.dstaddr).map(Into::into)
}
}
fn set_destination(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.dstaddr = SockAddr::from(value).into();
if siocsifdstaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn broadcast(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifbrdaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifru.broadaddr).map(Into::into)
}
}
fn set_broadcast(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.broadaddr = SockAddr::from(value).into();
if siocsifbrdaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn netmask(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifnetmask(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifru.netmask).map(Into::into)
}
}
fn set_netmask(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.netmask = SockAddr::from(value).into();
if siocsifnetmask(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn mtu(&self) -> Result<i32> {
unsafe {
let mut req = self.request();
if siocgifmtu(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(req.ifru.mtu)
}
}
fn set_mtu(&mut self, value: i32) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.mtu = value;
if siocsifmtu(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn queue(&self, index: usize) -> Option<&Self::Queue> {
self.queues.get(index)
}
}
pub struct Queue {
tun: Arc<Fd>,
pi_enabled: bool,
}
impl Queue {
pub fn has_packet_information(&self) -> bool {
self.pi_enabled
}
pub fn set_nonblock(&self) -> io::Result<()> {
self.tun.set_nonblock()
}
pub fn reader(&self) -> posix::Reader {
posix::Reader(self.tun.clone())
}
pub fn writer(&self) -> posix::Writer {
posix::Writer(self.tun.clone())
}
}
impl From<Layer> for c_short {
fn from(layer: Layer) -> Self {
match layer {
Layer::L2 => IFF_TAP,
Layer::L3 => IFF_TUN,
}
}
}
-43
View File
@@ -1,43 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
//! Linux specific functionality.
pub mod sys;
mod device;
pub use self::device::{Device, Queue};
use crate::configuration::Configuration as C;
use crate::error::*;
/// Linux-only interface configuration.
#[derive(Copy, Clone, Default, Debug)]
pub struct Configuration {
pub(crate) packet_information: bool,
}
impl Configuration {
/// Enable or disable packet information, when enabled the first 4 bytes of
/// each packet is a header with flags and protocol type.
pub fn packet_information(&mut self, value: bool) -> &mut Self {
self.packet_information = value;
self
}
}
/// Create a TUN device with the given name.
pub fn create(configuration: &C) -> Result<Device> {
Device::new(configuration)
}
-111
View File
@@ -1,111 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
//! Bindings to internal Linux stuff.
use ioctl::*;
use libc::sockaddr;
use libc::{c_char, c_int, c_short, c_uchar, c_uint, c_ulong, c_ushort, c_void};
pub const IFNAMSIZ: usize = 16;
pub const IFF_UP: c_short = 0x1;
pub const IFF_RUNNING: c_short = 0x40;
pub const IFF_TUN: c_short = 0x0001;
pub const IFF_TAP: c_short = 0x0002;
pub const IFF_NO_PI: c_short = 0x1000;
pub const IFF_MULTI_QUEUE: c_short = 0x0100;
#[repr(C)]
#[derive(Copy, Clone)]
pub struct ifmap {
pub mem_start: c_ulong,
pub mem_end: c_ulong,
pub base_addr: c_ushort,
pub irq: c_uchar,
pub dma: c_uchar,
pub port: c_uchar,
}
#[repr(C)]
#[derive(Copy, Clone)]
pub union ifsu {
pub raw_hdlc_proto: *mut c_void,
pub cisco: *mut c_void,
pub fr: *mut c_void,
pub fr_pvc: *mut c_void,
pub fr_pvc_info: *mut c_void,
pub sync: *mut c_void,
pub te1: *mut c_void,
}
#[repr(C)]
#[derive(Copy, Clone)]
pub struct if_settings {
pub type_: c_uint,
pub size: c_uint,
pub ifsu: ifsu,
}
#[repr(C)]
#[derive(Copy, Clone)]
pub union ifrn {
pub name: [c_char; IFNAMSIZ],
}
#[repr(C)]
#[derive(Copy, Clone)]
pub union ifru {
pub addr: sockaddr,
pub dstaddr: sockaddr,
pub broadaddr: sockaddr,
pub netmask: sockaddr,
pub hwaddr: sockaddr,
pub flags: c_short,
pub ivalue: c_int,
pub mtu: c_int,
pub map: ifmap,
pub slave: [c_char; IFNAMSIZ],
pub newname: [c_char; IFNAMSIZ],
pub data: *mut c_void,
pub settings: if_settings,
}
#[repr(C)]
#[derive(Copy, Clone)]
pub struct ifreq {
pub ifrn: ifrn,
pub ifru: ifru,
}
ioctl!(bad read siocgifflags with 0x8913; ifreq);
ioctl!(bad write siocsifflags with 0x8914; ifreq);
ioctl!(bad read siocgifaddr with 0x8915; ifreq);
ioctl!(bad write siocsifaddr with 0x8916; ifreq);
ioctl!(bad read siocgifdstaddr with 0x8917; ifreq);
ioctl!(bad write siocsifdstaddr with 0x8918; ifreq);
ioctl!(bad read siocgifbrdaddr with 0x8919; ifreq);
ioctl!(bad write siocsifbrdaddr with 0x891a; ifreq);
ioctl!(bad read siocgifnetmask with 0x891b; ifreq);
ioctl!(bad write siocsifnetmask with 0x891c; ifreq);
ioctl!(bad read siocgifmtu with 0x8921; ifreq);
ioctl!(bad write siocsifmtu with 0x8922; ifreq);
ioctl!(bad write siocsifname with 0x8923; ifreq);
ioctl!(write tunsetiff with b'T', 202; c_int);
ioctl!(write tunsetpersist with b'T', 203; c_int);
ioctl!(write tunsetowner with b'T', 204; c_int);
ioctl!(write tunsetgroup with b'T', 206; c_int);
-445
View File
@@ -1,445 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
#![allow(unused_variables)]
use std::ffi::CStr;
use std::io;
use std::mem;
use std::net::Ipv4Addr;
use std::os::unix::io::AsRawFd;
use std::ptr;
use std::sync::Arc;
use libc;
use libc::{c_char, c_uint, c_void, sockaddr, socklen_t, AF_INET, SOCK_DGRAM};
use crate::configuration::{Configuration, Layer};
use crate::device::Device as D;
use crate::error::*;
use crate::platform::macos::sys::*;
use crate::platform::posix::{self, Fd, SockAddr};
/// A TUN device using the TUN macOS driver.
pub struct Device {
name: String,
queue: Queue,
ctl: Fd,
}
impl Device {
/// Create a new `Device` for the given `Configuration`.
pub fn new(config: &Configuration) -> Result<Self> {
let id = if let Some(name) = config.name.as_ref() {
if name.len() > IFNAMSIZ {
return Err(Error::NameTooLong);
}
if !name.starts_with("utun") {
return Err(Error::InvalidName);
}
name[4..].parse()?
} else {
0
};
if config.layer.filter(|l| *l != Layer::L3).is_some() {
return Err(Error::UnsupportedLayer);
}
let queues_number = config.queues.unwrap_or(1);
if queues_number != 1 {
return Err(Error::InvalidQueuesNumber);
}
let mut device = unsafe {
let tun = Fd::new(libc::socket(PF_SYSTEM, SOCK_DGRAM, SYSPROTO_CONTROL))
.map_err(|_| io::Error::last_os_error())?;
let mut info = ctl_info {
ctl_id: 0,
ctl_name: {
let mut buffer = [0; 96];
for (i, o) in UTUN_CONTROL_NAME.as_bytes().iter().zip(buffer.iter_mut()) {
*o = *i as _;
}
buffer
},
};
if ctliocginfo(tun.0, &mut info as *mut _ as *mut _) < 0 {
return Err(io::Error::last_os_error().into());
}
let addr = sockaddr_ctl {
sc_id: info.ctl_id,
sc_len: mem::size_of::<sockaddr_ctl>() as _,
sc_family: AF_SYSTEM,
ss_sysaddr: AF_SYS_CONTROL,
sc_unit: id as c_uint,
sc_reserved: [0; 5],
};
if libc::connect(
tun.0,
&addr as *const sockaddr_ctl as *const sockaddr,
mem::size_of_val(&addr) as socklen_t,
) < 0
{
return Err(io::Error::last_os_error().into());
}
let mut name = [0u8; 64];
let mut name_len: socklen_t = 64;
if libc::getsockopt(
tun.0,
SYSPROTO_CONTROL,
UTUN_OPT_IFNAME,
&mut name as *mut _ as *mut c_void,
&mut name_len as *mut socklen_t,
) < 0
{
return Err(io::Error::last_os_error().into());
}
let ctl = Fd::new(libc::socket(AF_INET, SOCK_DGRAM, 0))
.map_err(|_| io::Error::last_os_error())?;
Device {
name: CStr::from_ptr(name.as_ptr() as *const c_char)
.to_string_lossy()
.into(),
queue: Queue { tun: Arc::new(tun) },
ctl: ctl,
}
};
device.configure(&config)?;
Ok(device)
}
/// Prepare a new request.
pub unsafe fn request(&self) -> ifreq {
let mut req: ifreq = mem::zeroed();
ptr::copy_nonoverlapping(
self.name.as_ptr() as *const c_char,
req.ifrn.name.as_mut_ptr(),
self.name.len(),
);
req
}
/// Set the IPv4 alias of the device.
pub fn set_alias(&mut self, addr: Ipv4Addr, broadaddr: Ipv4Addr, mask: Ipv4Addr) -> Result<()> {
unsafe {
let mut req: ifaliasreq = mem::zeroed();
ptr::copy_nonoverlapping(
self.name.as_ptr() as *const c_char,
req.ifran.as_mut_ptr(),
self.name.len(),
);
req.addr = SockAddr::from(addr).into();
req.broadaddr = SockAddr::from(broadaddr).into();
req.mask = SockAddr::from(mask).into();
if siocaifaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
// /// Split the interface into a `Reader` and `Writer`.
// pub fn split(self) -> (posix::Reader, posix::Writer) {
// let fd = Arc::new(self.queue.tun);
// (posix::Reader(fd.clone()), posix::Writer(fd.clone()))
// }
/// Return whether the device has packet information
pub fn has_packet_information(&self) -> bool {
self.queue.has_packet_information()
}
/// Set non-blocking mode
pub fn set_nonblock(&self) -> io::Result<()> {
self.queue.set_nonblock()
}
}
// impl Read for Device {
// fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
// self.queue.tun.read(buf)
// }
//
// fn read_vectored(&mut self, bufs: &mut [io::IoSliceMut<'_>]) -> io::Result<usize> {
// self.queue.tun.read_vectored(bufs)
// }
// }
//
// impl Write for Device {
// fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
// self.queue.tun.write(buf)
// }
//
// fn flush(&mut self) -> io::Result<()> {
// self.queue.tun.flush()
// }
//
// fn write_vectored(&mut self, bufs: &[io::IoSlice<'_>]) -> io::Result<usize> {
// self.queue.tun.write_vectored(bufs)
// }
// }
impl D for Device {
type Queue = Queue;
fn name(&self) -> &str {
&self.name
}
// XXX: Cannot set interface name on Darwin.
fn set_name(&mut self, value: &str) -> Result<()> {
Err(Error::InvalidName)
}
fn enabled(&mut self, value: bool) -> Result<()> {
unsafe {
let mut req = self.request();
if siocgifflags(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
if value {
req.ifru.flags |= IFF_UP | IFF_RUNNING;
} else {
req.ifru.flags &= !IFF_UP;
}
if siocsifflags(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn address(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifru.addr).map(Into::into)
}
}
fn set_address(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.addr = SockAddr::from(value).into();
if siocsifaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn destination(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifdstaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifru.dstaddr).map(Into::into)
}
}
fn set_destination(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.dstaddr = SockAddr::from(value).into();
if siocsifdstaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn broadcast(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifbrdaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifru.broadaddr).map(Into::into)
}
}
fn set_broadcast(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.broadaddr = SockAddr::from(value).into();
if siocsifbrdaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn netmask(&self) -> Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifnetmask(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::unchecked(&req.ifru.addr).map(Into::into)
}
}
fn set_netmask(&mut self, value: Ipv4Addr) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.addr = SockAddr::from(value).into();
if siocsifnetmask(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn mtu(&self) -> Result<i32> {
unsafe {
let mut req = self.request();
if siocgifmtu(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(req.ifru.mtu)
}
}
fn set_mtu(&mut self, value: i32) -> Result<()> {
unsafe {
let mut req = self.request();
req.ifru.mtu = value;
if siocsifmtu(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error().into());
}
Ok(())
}
}
fn queue(&self, index: usize) -> Option<&Self::Queue> {
if index > 0 {
return None;
}
Some(&self.queue)
}
}
// impl AsRawFd for Device {
// fn as_raw_fd(&self) -> RawFd {
// self.queue.as_raw_fd()
// }
// }
//
// impl IntoRawFd for Device {
// fn into_raw_fd(self) -> RawFd {
// self.queue.into_raw_fd()
// }
// }
pub struct Queue {
tun: Arc<Fd>,
}
impl Queue {
pub fn has_packet_information(&self) -> bool {
// on macos this is always the case
true
}
pub fn set_nonblock(&self) -> io::Result<()> {
self.tun.set_nonblock()
}
pub fn reader(&self) -> posix::Reader {
posix::Reader(self.tun.clone())
}
pub fn writer(&self) -> posix::Writer {
posix::Writer(self.tun.clone())
}
}
// impl AsRawFd for Queue {
// fn as_raw_fd(&self) -> RawFd {
// self.tun.as_raw_fd()
// }
// }
//
// impl IntoRawFd for Queue {
// fn into_raw_fd(self) -> RawFd {
// self.tun.into_raw_fd()
// }
// }
// impl Read for Queue {
// fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
// self.tun.read(buf)
// }
//
// fn read_vectored(&mut self, bufs: &mut [io::IoSliceMut<'_>]) -> io::Result<usize> {
// self.tun.read_vectored(bufs)
// }
// }
//
// impl Write for Queue {
// fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
// self.tun.write(buf)
// }
//
// fn flush(&mut self) -> io::Result<()> {
// self.tun.flush()
// }
//
// fn write_vectored(&mut self, bufs: &[io::IoSlice<'_>]) -> io::Result<usize> {
// self.tun.write_vectored(bufs)
// }
// }
-32
View File
@@ -1,32 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
//! macOS specific functionality.
pub mod sys;
mod device;
pub use self::device::{Device, Queue};
use crate::configuration::Configuration as C;
use crate::error::*;
/// macOS-only interface configuration.
#[derive(Copy, Clone, Default, Debug)]
pub struct Configuration {}
/// Create a TUN device with the given name.
pub fn create(configuration: &C) -> Result<Device> {
Device::new(&configuration)
}
-60
View File
@@ -1,60 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
//! Platform specific modules.
#[cfg(unix)]
pub mod posix;
#[cfg(target_os = "linux")]
pub mod linux;
#[cfg(target_os = "linux")]
pub use self::linux::{create, Configuration, Device, Queue};
#[cfg(target_os = "macos")]
pub mod macos;
#[cfg(target_os = "macos")]
pub use self::macos::{create, Configuration, Device, Queue};
#[cfg(test)]
mod test {
use crate::configuration::Configuration;
use crate::device::Device;
use std::net::Ipv4Addr;
#[test]
fn create() {
let dev = super::create(
Configuration::default()
.name("utun6")
.address("192.168.50.1")
.netmask("255.255.0.0")
.mtu(1400)
.up(),
)
.unwrap();
assert_eq!(
"192.168.50.1".parse::<Ipv4Addr>().unwrap(),
dev.address().unwrap()
);
assert_eq!(
"255.255.0.0".parse::<Ipv4Addr>().unwrap(),
dev.netmask().unwrap()
);
assert_eq!(1400, dev.mtu().unwrap());
}
}
-124
View File
@@ -1,124 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::io::{self, Read, Write};
use std::os::unix::io::{AsRawFd, IntoRawFd, RawFd};
use crate::error::*;
use libc::{self, fcntl, F_GETFL, F_SETFL, O_NONBLOCK};
/// POSIX file descriptor support for `io` traits.
pub struct Fd(pub RawFd);
impl Fd {
pub fn new(value: RawFd) -> Result<Self> {
if value < 0 {
return Err(Error::InvalidDescriptor);
}
Ok(Fd(value))
}
/// Enable non-blocking mode
pub fn set_nonblock(&self) -> io::Result<()> {
match unsafe { fcntl(self.0, F_SETFL, fcntl(self.0, F_GETFL) | O_NONBLOCK) } {
0 => Ok(()),
_ => Err(io::Error::last_os_error()),
}
}
}
impl Read for Fd {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
unsafe {
let amount = libc::read(self.0, buf.as_mut_ptr() as *mut _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(amount as usize)
}
}
fn read_vectored(&mut self, bufs: &mut [io::IoSliceMut<'_>]) -> io::Result<usize> {
unsafe {
let iov = bufs.as_ptr().cast();
let iovcnt = bufs.len().min(libc::c_int::MAX as usize) as _;
let n = libc::readv(self.0, iov, iovcnt);
if n < 0 {
return Err(io::Error::last_os_error());
}
Ok(n as usize)
}
}
}
impl Write for Fd {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
unsafe {
let amount = libc::write(self.0, buf.as_ptr() as *const _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(amount as usize)
}
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
fn write_vectored(&mut self, bufs: &[io::IoSlice<'_>]) -> io::Result<usize> {
unsafe {
let iov = bufs.as_ptr().cast();
let iovcnt = bufs.len().min(libc::c_int::MAX as usize) as _;
let n = libc::writev(self.0, iov, iovcnt);
if n < 0 {
return Err(io::Error::last_os_error());
}
Ok(n as usize)
}
}
}
impl AsRawFd for Fd {
fn as_raw_fd(&self) -> RawFd {
self.0
}
}
impl IntoRawFd for Fd {
fn into_raw_fd(mut self) -> RawFd {
let fd = self.0;
self.0 = -1;
fd
}
}
impl Drop for Fd {
fn drop(&mut self) {
unsafe {
if self.0 >= 0 {
libc::close(self.0);
}
}
}
}
-24
View File
@@ -1,24 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
//! POSIX compliant support.
mod sockaddr;
pub use self::sockaddr::SockAddr;
mod fd;
pub use self::fd::Fd;
mod split;
pub use self::split::{Reader, Writer};
-124
View File
@@ -1,124 +0,0 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::io;
use std::mem;
use std::os::unix::io::{AsRawFd, RawFd};
use std::sync::Arc;
use crate::platform::posix::Fd;
use libc;
/// Read-only end for a file descriptor.
#[derive(Clone)]
pub struct Reader(pub(crate) Arc<Fd>);
/// Write-only end for a file descriptor.
#[derive(Clone)]
pub struct Writer(pub(crate) Arc<Fd>);
impl Reader {
pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
unsafe {
let amount = libc::read(self.0.as_raw_fd(), buf.as_mut_ptr() as *mut _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(amount as usize)
}
}
pub fn read_vectored(&self, bufs: &mut [io::IoSliceMut<'_>]) -> io::Result<usize> {
unsafe {
let mut msg: libc::msghdr = mem::zeroed();
// msg.msg_name: NULL
// msg.msg_namelen: 0
msg.msg_iov = bufs.as_mut_ptr().cast();
msg.msg_iovlen = bufs.len().min(libc::c_int::MAX as usize) as _;
let n = libc::recvmsg(self.0.as_raw_fd(), &mut msg, 0);
if n < 0 {
return Err(io::Error::last_os_error());
}
Ok(n as usize)
}
}
}
impl Writer {
pub fn write(&self, buf: &[u8]) -> io::Result<usize> {
unsafe {
let amount = libc::write(self.0.as_raw_fd(), buf.as_ptr() as *const _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(amount as usize)
}
}
pub fn write_vectored(&self, bufs: &[io::IoSlice<'_>]) -> io::Result<usize> {
unsafe {
let mut msg: libc::msghdr = mem::zeroed();
// msg.msg_name = NULL
// msg.msg_namelen = 0
msg.msg_iov = bufs.as_ptr() as *mut _;
msg.msg_iovlen = bufs.len().min(libc::c_int::MAX as usize) as _;
let n = libc::sendmsg(self.0.as_raw_fd(), &msg, 0);
if n < 0 {
return Err(io::Error::last_os_error());
}
Ok(n as usize)
}
}
pub fn write_all(&self, mut buf: &[u8]) -> io::Result<()> {
while !buf.is_empty() {
match self.write(buf) {
Ok(0) => {
return Err(io::Error::new(
io::ErrorKind::WriteZero,
"failed to write whole buffer",
));
}
Ok(n) => buf = &buf[n..],
Err(ref e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
Ok(())
}
}
impl AsRawFd for Reader {
fn as_raw_fd(&self) -> RawFd {
self.0.as_raw_fd()
}
}
impl AsRawFd for Writer {
fn as_raw_fd(&self) -> RawFd {
self.0.as_raw_fd()
}
}
//
// impl AsRawFd for Writer {
// fn as_raw_fd(&self) -> RawFd {
// self.0.as_raw_fd()
// }
// }
-980
View File
@@ -1,980 +0,0 @@
use std::collections::HashMap;
use std::io;
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
use std::ops::Sub;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::{Duration, Instant};
use byte_pool::{Block, BytePool};
use crossbeam_epoch::{Atomic, Owned};
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use std::net::UdpSocket as StdUdpSocket;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::tcp::OwnedReadHalf;
use tokio::net::{TcpStream, UdpSocket};
use tokio::sync::watch::{channel, Receiver, Sender};
use crate::channel::punch::NatType;
use crate::channel::{Route, RouteKey, Status, TCP_ID, UDP_ID, UDP_V6_ID};
use crate::core::status::VntWorker;
use crate::handle::recv_handler::ChannelDataHandler;
use crate::handle::CurrentDeviceInfo;
use crate::ip_proxy::DashMapNew;
lazy_static::lazy_static! {
static ref POOL:BytePool = BytePool::new();
}
pub struct ContextInner {
//udp用于打洞、服务端通信(可选)
pub(crate) main_channel: Arc<StdUdpSocket>,
pub(crate) main_channel_ipv6: Option<Arc<StdUdpSocket>>,
//在udp的基础上,可以选择使用tcp和服务端通信
pub(crate) main_tcp_channel: Option<tokio::sync::mpsc::Sender<Vec<u8>>>,
pub(crate) route_table: Atomic<HashMap<Ipv4Addr, Vec<Route>>>,
pub(crate) route_table_time: DashMap<(RouteKey, Ipv4Addr), Instant>,
pub(crate) status_receiver: Receiver<Status>,
pub(crate) status_sender: Sender<Status>,
pub(crate) udp_map: Atomic<HashMap<usize, Arc<UdpSocket>>>,
pub(crate) channel_num: usize,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
}
#[derive(Clone)]
pub struct Context {
pub(crate) inner: Arc<ContextInner>,
}
impl Context {
pub fn new(
main_channel: Arc<StdUdpSocket>,
main_channel_ipv6: Option<Arc<StdUdpSocket>>,
main_tcp_channel: Option<tokio::sync::mpsc::Sender<Vec<u8>>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
_channel_num: usize,
) -> Self {
//当前版本只支持一个通道
let channel_num = 1;
let (status_sender, status_receiver) = channel(Status::Cone);
let inner = Arc::new(ContextInner {
main_channel,
main_channel_ipv6,
main_tcp_channel,
route_table: Atomic::new(HashMap::with_capacity(16)),
route_table_time: DashMap::new_cap(16),
status_receiver,
status_sender,
udp_map: Atomic::new(HashMap::with_capacity(16)),
channel_num,
current_device,
});
Self { inner }
}
}
impl Context {
pub fn is_close(&self) -> bool {
*self.inner.status_receiver.borrow() == Status::Close
}
pub fn is_cone(&self) -> bool {
*self.inner.status_receiver.borrow() == Status::Cone
}
pub fn close(&self) -> io::Result<()> {
let _ = self.inner.status_sender.send(Status::Close);
if let Ok(port) = self.main_local_ipv4_port() {
let _ = StdUdpSocket::bind("127.0.0.1:0")?.send_to(
b"stop",
SocketAddr::V4(std::net::SocketAddrV4::new(Ipv4Addr::LOCALHOST, port)),
);
}
if let Ok(port) = self.main_local_ipv6_port() {
let _ = StdUdpSocket::bind("[::]:0")?.send_to(
b"stop",
SocketAddr::V6(std::net::SocketAddrV6::new(Ipv6Addr::LOCALHOST, port, 0, 0)),
);
}
Ok(())
}
pub fn is_main_tcp(&self) -> bool {
self.inner.main_tcp_channel.is_some()
}
pub fn switch(&self, nat_type: NatType) {
match nat_type {
NatType::Symmetric => {
self.switch_to_symmetric();
}
NatType::Cone => {
self.switch_to_cone();
}
}
}
pub fn switch_to_cone(&self) {
let _ = self.inner.status_sender.send(Status::Cone);
}
pub fn switch_to_symmetric(&self) {
let _ = self.inner.status_sender.send(Status::Symmetric);
}
pub fn main_local_ipv4_port(&self) -> io::Result<u16> {
self.inner.main_channel.local_addr().map(|k| k.port())
}
pub fn main_local_ipv6_port(&self) -> io::Result<u16> {
if let Some(ipv6) = &self.inner.main_channel_ipv6 {
ipv6.local_addr().map(|k| k.port())
} else {
Err(io::Error::new(io::ErrorKind::Other, "not ipv6"))
}
}
fn insert_udp(&self, id: usize, udp: Arc<UdpSocket>) {
self.insert_udp_(id, Some(udp))
}
fn remove_udp(&self, id: usize) {
self.insert_udp_(id, None)
}
fn insert_udp_(&self, id: usize, udp: Option<Arc<UdpSocket>>) {
let guard = &crossbeam_epoch::pin();
let udp_map = &self.inner.udp_map;
let mut udp_map_shared = self.inner.udp_map.load(Ordering::Relaxed, guard);
loop {
let mut map = unsafe { udp_map_shared.as_ref().unwrap().clone() };
match udp.clone() {
None => {
map.remove(&id);
}
Some(udp) => {
map.insert(id, udp);
}
}
match udp_map.compare_exchange(
udp_map_shared,
Owned::new(map),
Ordering::Relaxed,
Ordering::Relaxed,
guard,
) {
Ok(p) => unsafe {
guard.defer_destroy(p);
return;
},
Err(e) => {
udp_map_shared = e.current;
}
}
}
}
pub fn send_main_udp(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> {
if addr.is_ipv6() {
if let Some(udp_ipv6) = &self.inner.main_channel_ipv6 {
udp_ipv6.send_to(buf, addr)
} else {
Err(io::Error::new(io::ErrorKind::Other, "not ipv6"))
}
} else {
self.inner.main_channel.send_to(buf, addr)
}
}
pub fn send_main(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> {
if let Some(sender) = &self.inner.main_tcp_channel {
if sender.try_send(buf.to_vec()).is_ok() {
Ok(buf.len())
} else {
Err(io::Error::new(io::ErrorKind::Other, "send_main err"))
}
} else {
self.send_main_udp(buf, addr)
}
}
pub(crate) fn try_send_all(&self, buf: &[u8], addr: SocketAddr) -> io::Result<()> {
let table = unsafe {
let guard = &crossbeam_epoch::pin();
self.inner
.udp_map
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
.clone()
};
if table.is_empty() {
log::error!("udp列表为空,addr={}", addr);
return Ok(());
}
for (_, udp) in table {
//使用ipv6的udp发送ipv4报文会出错
if let Err(e) = udp.try_send_to(buf, addr) {
log::error!("{:?}", e);
}
}
Ok(())
}
pub async fn send_by_id(&self, buf: &[u8], id: &Ipv4Addr) -> io::Result<usize> {
let route = self.get_route_by_id(id)?;
self.send_by_key(buf, &route.route_key()).await
}
pub fn try_send_by_id(&self, buf: &[u8], id: &Ipv4Addr) -> io::Result<usize> {
let route = self.get_route_by_id(id)?;
self.try_send_by_key(buf, &route.route_key())
}
fn get_route_by_id(&self, id: &Ipv4Addr) -> io::Result<Route> {
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
if let Some(v) = table.get(id) {
if v.is_empty() {
return Err(io::Error::new(io::ErrorKind::NotFound, "route not found"));
}
let route = v[0];
if route.rt == 199 {
//这通常是刚加入路由,直接放弃使用,避免抖动
return Err(io::Error::new(io::ErrorKind::NotFound, "route not found"));
}
if !route.is_p2p() {
if let Some(time) = self.inner.route_table_time.get(&(route.route_key(), *id)) {
//借道传输时,长时间不通信的通道不使用
if time.value().elapsed() > Duration::from_secs(6) {
return Err(io::Error::new(io::ErrorKind::NotFound, "route time out"));
}
}
}
return Ok(route);
}
Err(io::Error::new(io::ErrorKind::NotFound, "route not found"))
}
pub async fn send_by_key(&self, buf: &[u8], route_key: &RouteKey) -> io::Result<usize> {
match route_key.index {
TCP_ID => {
if let Some(sender) = &self.inner.main_tcp_channel {
if sender.send(buf.to_vec()).await.is_ok() {
Ok(buf.len())
} else {
Err(io::Error::new(io::ErrorKind::Other, "send_by_key err"))
}
} else {
Err(io::Error::new(io::ErrorKind::Other, "send_by_key err"))
}
}
UDP_ID => self.inner.main_channel.send_to(buf, route_key.addr),
UDP_V6_ID => {
if let Some(udp_ipv6) = &self.inner.main_channel_ipv6 {
udp_ipv6.send_to(buf, route_key.addr)
} else {
Err(io::Error::new(io::ErrorKind::Other, "not ipv6 udp"))
}
}
_ => {
if let Some(udp) = self.get_udp_by_route(route_key) {
return udp.send_to(buf, route_key.addr).await;
}
Err(io::Error::new(io::ErrorKind::NotFound, "route not found"))
}
}
}
pub fn try_send_by_key(&self, buf: &[u8], route_key: &RouteKey) -> io::Result<usize> {
match route_key.index {
TCP_ID => {
if let Some(sender) = &self.inner.main_tcp_channel {
if sender.try_send(buf.to_vec()).is_ok() {
Ok(buf.len())
} else {
Err(io::Error::new(io::ErrorKind::Other, "send_by_key err"))
}
} else {
Err(io::Error::new(io::ErrorKind::Other, "send_by_key err"))
}
}
UDP_ID => self.inner.main_channel.send_to(buf, route_key.addr),
UDP_V6_ID => {
if let Some(udp_ipv6) = &self.inner.main_channel_ipv6 {
udp_ipv6.send_to(buf, route_key.addr)
} else {
Err(io::Error::new(io::ErrorKind::Other, "not ipv6 udp"))
}
}
_ => {
if let Some(udp) = self.get_udp_by_route(route_key) {
return udp.try_send_to(buf, route_key.addr);
}
Err(io::Error::new(io::ErrorKind::NotFound, "route not found"))
}
}
}
fn get_udp_by_route(&self, route_key: &RouteKey) -> Option<Arc<UdpSocket>> {
let guard = &crossbeam_epoch::pin();
let udp_map = unsafe {
self.inner
.udp_map
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
udp_map.get(&route_key.index).cloned()
}
pub fn add_route_if_absent(&self, id: Ipv4Addr, route: Route) {
self.add_route_(id, route, true)
}
pub fn add_route(&self, id: Ipv4Addr, route: Route) {
self.add_route_(id, route, false)
}
fn add_route_(&self, id: Ipv4Addr, route: Route, only_if_absent: bool) {
let key = route.route_key();
let guard = &crossbeam_epoch::pin();
let route_table = &self.inner.route_table;
let mut table_share = route_table.load(Ordering::Relaxed, guard);
loop {
let mut table = unsafe { table_share.as_ref().unwrap().clone() };
let list = table.entry(id).or_insert_with(|| Vec::with_capacity(4));
let mut exist = false;
for x in list.iter_mut() {
if x.metric < route.metric {
//不能比当前的路径更长
return;
}
if x.route_key() == key {
if only_if_absent {
return;
}
x.metric = route.metric;
x.rt = route.rt;
exist = true;
break;
}
}
if exist {
list.sort_by_key(|k| k.sort_key());
} else {
if route.metric == 1 {
//添加了直连的则排除非直连的
list.retain(|k| k.metric == 1);
}
list.push(route);
list.sort_by_key(|k| k.sort_key());
let max_len = self.inner.channel_num + 1;
if list.len() > max_len {
list.truncate(max_len);
}
}
match route_table.compare_exchange(
table_share,
Owned::new(table),
Ordering::Relaxed,
Ordering::Relaxed,
guard,
) {
Ok(p) => unsafe {
guard.defer_destroy(p);
break;
},
Err(e) => {
table_share = e.current;
}
}
}
self.inner
.route_table_time
.insert((key, id), Instant::now().sub(Duration::from_secs(10)));
}
pub fn route(&self, id: &Ipv4Addr) -> Option<Vec<Route>> {
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
if let Some(v) = table.get(id) {
Some(v.clone())
} else {
None
}
}
pub fn route_one(&self, id: &Ipv4Addr) -> Option<Route> {
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
if let Some(v) = table.get(id) {
v.first().map(|v| *v)
} else {
None
}
}
pub fn route_to_id(&self, route_key: &RouteKey) -> Option<Ipv4Addr> {
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
for (k, v) in table.iter() {
for route in v {
if &route.route_key() == route_key && route.is_p2p() {
return Some(*k);
}
}
}
None
}
pub fn need_punch(&self, id: &Ipv4Addr) -> bool {
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
if let Some(v) = table.get(id) {
if v.iter().filter(|k| k.is_p2p()).count() >= self.inner.channel_num {
return false;
}
}
true
}
pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec<Route>)> {
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
table.iter().map(|(k, v)| (k.clone(), v.clone())).collect()
}
pub fn route_table_one(&self) -> Vec<(Ipv4Addr, Route)> {
let mut list = Vec::with_capacity(8);
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
for (k, v) in table {
if let Some(route) = v.first() {
list.push((*k, *route));
}
}
list
}
pub fn direct_route_table_one(&self) -> Vec<(Ipv4Addr, Route)> {
let mut list = Vec::with_capacity(8);
let guard = &crossbeam_epoch::pin();
let table = unsafe {
self.inner
.route_table
.load(Ordering::Relaxed, guard)
.as_ref()
.unwrap()
};
for (k, v) in table {
if let Some(route) = v.first() {
if route.metric == 1 {
list.push((*k, *route));
}
}
}
list
}
pub fn remove_route(&self, id: &Ipv4Addr, route_key: RouteKey) {
let guard = &crossbeam_epoch::pin();
let route_table = &self.inner.route_table;
let mut table_share = route_table.load(Ordering::Relaxed, guard);
loop {
let mut table = unsafe { table_share.as_ref().unwrap().clone() };
if let Some(routes) = table.get_mut(id) {
routes.retain(|x| x.route_key() != route_key);
match route_table.compare_exchange(
table_share,
Owned::new(table),
Ordering::Relaxed,
Ordering::Relaxed,
guard,
) {
Ok(p) => unsafe {
guard.defer_destroy(p);
self.inner.route_table_time.remove(&(route_key, *id));
return;
},
Err(e) => {
table_share = e.current;
}
}
}
}
}
pub fn update_read_time(&self, id: &Ipv4Addr, route_key: &RouteKey) {
if let Some(mut time) = self.inner.route_table_time.get_mut(&(*route_key, *id)) {
*time.value_mut() = Instant::now();
}
}
}
pub struct Channel {
context: Context,
handler: ChannelDataHandler,
}
impl Channel {
pub fn new(context: Context, handler: ChannelDataHandler) -> Self {
Self { context, handler }
}
}
#[derive(Clone)]
struct BufSenderGroup(
usize,
Vec<std::sync::mpsc::SyncSender<(Block<'static>, usize, usize, RouteKey)>>,
);
struct BufReceiverGroup(Vec<std::sync::mpsc::Receiver<(Block<'static>, usize, usize, RouteKey)>>);
impl BufSenderGroup {
pub fn send(&mut self, val: (Block<'static>, usize, usize, RouteKey)) -> bool {
let index = self.0 % self.1.len();
self.0 = self.0.wrapping_add(1);
self.1[index].send(val).is_ok()
}
}
fn buf_channel_group(size: usize) -> (BufSenderGroup, BufReceiverGroup) {
let mut buf_sender_group = Vec::with_capacity(size);
let mut buf_receiver_group = Vec::with_capacity(size);
for _ in 0..size {
let (buf_sender, buf_receiver) =
std::sync::mpsc::sync_channel::<(Block<'static, Vec<u8>>, usize, usize, RouteKey)>(1);
buf_sender_group.push(buf_sender);
buf_receiver_group.push(buf_receiver);
}
(
BufSenderGroup(0, buf_sender_group),
BufReceiverGroup(buf_receiver_group),
)
}
impl Channel {
async fn tcp_handle(
mut tcp_r: OwnedReadHalf,
context: Context,
handler: ChannelDataHandler,
head_reserve: usize,
) -> io::Result<()> {
let mut head = [0; 4];
let addr = tcp_r.peer_addr()?;
let key = RouteKey::new(TCP_ID, addr);
loop {
let mut buf = [0; 4096];
tcp_r.read_exact(&mut head).await?;
let len = (((head[2] as u16) << 8) | head[3] as u16) as usize;
if len < 12 || len > buf.len() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"length overflow",
));
}
tcp_r
.read_exact(&mut buf[head_reserve..head_reserve + len])
.await?;
handler
.handle(&mut buf, head_reserve, head_reserve + len, key, &context)
.await;
}
}
async fn start_tcp(
mut worker: VntWorker,
tcp_stream: TcpStream,
mut receiver: tokio::sync::mpsc::Receiver<Vec<u8>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
context: Context,
handler: ChannelDataHandler,
head_reserve: usize,
) {
let (tcp_r, mut tcp_w) = tcp_stream.into_split();
{
let context = context.clone();
let handler = handler.clone();
tokio::spawn(async move {
if let Err(e) = Self::tcp_handle(tcp_r, context, handler, head_reserve).await {
log::info!("tcp链接断开:{:?}", e);
}
});
}
let mut head = [0; 4];
loop {
tokio::select! {
_=worker.stop_wait()=>{
break;
}
rs=receiver.recv()=>{
if let Some(data) = rs{
let len = data.len();
head[2] = (len >> 8) as u8;
head[3] = (len & 0xFF) as u8;
let mut err = false;
if let Err(e) = tcp_w.write_all(&head).await{
err = true;
log::info!("发送失败,需要重连:{:?}",e);
}else if let Err(e) = tcp_w.write_all(&data).await{
err = true;
log::info!("发送失败,需要重连:{:?}",e);
}
if err {
let _ = tcp_w.shutdown().await;
match TcpStream::connect(current_device.load().connect_server).await {
Ok(tcp_stream) => {
let (r, w) = tcp_stream.into_split();
tcp_w = w;
let context = context.clone();
let handler = handler.clone();
tokio::spawn(async move {
if let Err(e) = Self::tcp_handle(r, context,handler, head_reserve).await {
log::info!("tcp 链接断开:{:?}",e);
}
});
}
Err(e) => {
log::info!("重连失败:{:?}",e);
}
};
}
}else{
break;
}
}
}
}
worker.stop_all();
}
pub async fn start(
self,
mut worker: VntWorker,
tcp: Option<(TcpStream, tokio::sync::mpsc::Receiver<Vec<u8>>)>,
head_reserve: usize, //头部预留字节
symmetric_channel_num: usize, //对称网络,则再加一组监听,提升打洞成功率
relay: bool,
parallel: usize,
) {
let handler = self.handler.clone();
let context = self.context;
let main_channel = context.inner.main_channel.clone();
let buf_sender = if parallel > 1 {
let (buf_sender, buf_receiver) = buf_channel_group(parallel);
for buf_receiver in buf_receiver.0 {
let context = context.clone();
let handler = handler.clone();
std::thread::spawn(move || {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
log::info!("启动异步处理");
runtime.block_on(async move {
while let Ok((mut buf, start, end, route_key)) = buf_receiver.recv() {
handler
.handle(&mut buf, start, end, route_key, &context)
.await;
}
log::warn!("异步处理停止");
});
});
}
Some(buf_sender)
} else {
None
};
if let Some((tcp_stream, receiver)) = tcp {
tokio::spawn(Self::start_tcp(
worker.worker("main_channel_tcp"),
tcp_stream,
receiver,
context.inner.current_device.clone(),
context.clone(),
handler.clone(),
head_reserve,
));
}
if let Some(main_channel_ipv6) = &context.inner.main_channel_ipv6 {
let worker = worker.worker("main_channel_ipv6");
let context = context.clone();
let main_channel_ipv6 = main_channel_ipv6.clone();
let handler = handler.clone();
let buf_sender = buf_sender.clone();
std::thread::spawn(move || {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
log::info!("启动udp v6");
runtime.block_on(Self::main_start_(
worker,
context,
UDP_V6_ID,
main_channel_ipv6,
handler,
buf_sender,
head_reserve,
));
});
}
{
let worker = worker.worker("main_channel_1");
let context = context.clone();
let main_channel = main_channel.clone();
let handler = handler.clone();
let buf_sender = buf_sender.clone();
std::thread::spawn(move || {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
log::info!("启动udp v4");
runtime.block_on(Self::main_start_(
worker,
context,
UDP_ID,
main_channel,
handler,
buf_sender,
head_reserve,
));
});
}
if relay {
worker.stop_wait().await;
return;
}
let mut cur_status = Status::Cone;
let mut status_receiver = context.inner.status_receiver.clone();
loop {
tokio::select! {
_=worker.stop_wait()=>{
break;
}
rs=status_receiver.changed()=>{
match rs {
Ok(_) => {
let s = status_receiver.borrow().clone();
match s {
Status::Cone => {
cur_status = Status::Cone;
}
Status::Symmetric => {
if cur_status == Status::Symmetric {
continue;
}
cur_status = Status::Symmetric;
for _ in 0..symmetric_channel_num {
match UdpSocket::bind("0.0.0.0:0").await {
Ok(udp) => {
let udp = Arc::new(udp);
let context = context.clone();
tokio::spawn(Self::start_(worker.worker("symmetric_channel"),context, udp,handler.clone(),buf_sender.clone(), head_reserve, false));
}
Err(e) => {
log::error!("{}",e);
}
}
}
}
Status::Close => {
break;
}
}
}
Err(_) => {
break;
}
}
}
}
}
worker.stop_all();
}
async fn main_start_(
worker: VntWorker,
context: Context,
id: usize,
udp: Arc<StdUdpSocket>,
handler: ChannelDataHandler,
buf_sender: Option<BufSenderGroup>,
head_reserve: usize,
) {
match buf_sender {
None => {
let mut buf = [0; 4096];
loop {
match udp.recv_from(&mut buf[head_reserve..]) {
Ok((len, addr)) => {
let end = head_reserve + len;
if &buf[head_reserve..end] == b"stop" {
if context.is_close() {
break;
}
}
handler
.handle(
&mut buf,
head_reserve,
end,
RouteKey::new(id, addr),
&context,
)
.await;
}
Err(e) => {
log::error!("udp :{:?}", e);
}
}
}
}
Some(mut buf_sender) => loop {
let mut buf = POOL.alloc(4096);
match udp.recv_from(&mut buf[head_reserve..]) {
Ok((len, addr)) => {
let end = head_reserve + len;
if &buf[head_reserve..end] == b"stop" {
if context.is_close() {
break;
}
}
buf_sender.send((buf, head_reserve, end, RouteKey::new(id, addr)));
}
Err(e) => {
log::error!("udp :{:?}", e);
}
}
},
}
worker.stop_all();
}
async fn start_(
mut worker: VntWorker,
context: Context,
udp: Arc<UdpSocket>,
handler: ChannelDataHandler,
buf_sender: Option<BufSenderGroup>,
head_reserve: usize,
is_core: bool,
) {
let mut status_receiver = context.inner.status_receiver.clone();
#[cfg(target_os = "windows")]
use std::os::windows::io::AsRawSocket;
#[cfg(target_os = "windows")]
let id = 3 + udp.as_raw_socket() as usize;
#[cfg(any(unix))]
use std::os::fd::AsRawFd;
#[cfg(any(unix))]
let id = 3 + udp.as_raw_fd() as usize;
context.insert_udp(id, udp.clone());
match buf_sender {
None => {
let mut buf = [0; 4096];
loop {
tokio::select! {
rs=udp.recv_from(&mut buf[head_reserve..])=>{
match rs {
Ok((len, addr)) => {
handler.handle(&mut buf, head_reserve, head_reserve + len, RouteKey::new(id, addr), &context).await;
}
Err(e) => {
log::error!("{:?}",e)
}
}
}
changed=status_receiver.changed()=>{
match changed {
Ok(_) => {
match *status_receiver.borrow() {
Status::Cone => {
if !is_core{
break;
}
}
Status::Close=>{
break;
}
Status::Symmetric => {}
}
}
Err(_) => {
break;
}
}
}
_=worker.stop_wait()=>{
break;
}
}
}
}
Some(mut buf_sender) => loop {
let mut buf = POOL.alloc(4096);
tokio::select! {
rs=udp.recv_from(&mut buf[head_reserve..])=>{
match rs {
Ok((len, addr)) => {
if !buf_sender.send((buf,head_reserve,head_reserve+len,RouteKey::new(id, addr))){
log::error!("udp buf_sender发送数据失败");
break;
}
}
Err(e) => {
log::error!("{:?}",e)
}
}
}
changed=status_receiver.changed()=>{
match changed {
Ok(_) => {
match *status_receiver.borrow() {
Status::Cone => {
if !is_core{
break;
}
}
Status::Close=>{
break;
}
Status::Symmetric => {}
}
}
Err(_) => {
break;
}
}
}
_=worker.stop_wait()=>{
break;
}
}
},
}
context.remove_udp(id);
if is_core {
worker.stop_all();
}
}
}
+529
View File
@@ -0,0 +1,529 @@
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::Arc;
use std::time::{Duration, Instant};
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::RwLock;
use rand::Rng;
use crate::channel::punch::NatType;
use crate::channel::sender::{AcceptSocketSender, ChannelSender, PacketSender};
use crate::channel::{Route, RouteKey, UseChannelType, DEFAULT_RT};
/// 传输通道上下文,持有udp socket、tcp socket和路由信息
#[derive(Clone)]
pub struct Context {
inner: Arc<ContextInner>,
}
impl Context {
pub fn new(
main_udp_socket: Vec<UdpSocket>,
use_channel_type: UseChannelType,
first_latency: bool,
is_tcp: bool,
packet_loss_rate: Option<f64>,
packet_delay: u32,
use_ipv6: bool,
) -> Self {
let channel_num = main_udp_socket.len();
assert_ne!(channel_num, 0, "not channel");
let packet_loss_rate = packet_loss_rate
.map(|v| {
let v = (v * PACKET_LOSS_RATE_DENOMINATOR as f64) as u32;
if v > PACKET_LOSS_RATE_DENOMINATOR {
PACKET_LOSS_RATE_DENOMINATOR
} else {
v
}
})
.unwrap_or(0);
let inner = ContextInner {
main_udp_socket,
sub_udp_socket: RwLock::new(Vec::with_capacity(64)),
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),
use_ipv6,
};
Self {
inner: Arc::new(inner),
}
}
pub fn sender(&self) -> ChannelSender {
ChannelSender::new(self.clone())
}
}
impl Deref for Context {
type Target = ContextInner;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
/// 对称网络增加的udp socket数目,有助于增加打洞成功率
pub const SYMMETRIC_CHANNEL_NUM: usize = 100;
const PACKET_LOSS_RATE_DENOMINATOR: u32 = 100_0000;
pub struct ContextInner {
// 核心udp socket
pub(crate) main_udp_socket: Vec<UdpSocket>,
// 对称网络增加的udp socket
sub_udp_socket: RwLock<Vec<UdpSocket>>,
// tcp数据发送器
pub(crate) tcp_map: RwLock<HashMap<SocketAddr, PacketSender>>,
// 路由信息
pub route_table: RouteTable,
// 是否使用tcp连接服务器
is_tcp: bool,
//状态
state: AtomicBool,
//控制丢包率,取值v=[0,100_0000] 丢包率r=v/100_0000
packet_loss_rate: u32,
//控制延迟
packet_delay: u32,
main_index: AtomicUsize,
use_ipv6: bool,
}
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()
}
pub fn is_main_tcp(&self) -> bool {
self.is_tcp
}
pub fn is_udp_main(&self, route_key: &RouteKey) -> bool {
!route_key.is_tcp() && route_key.index < self.main_udp_socket.len()
}
pub fn first_latency(&self) -> bool {
self.route_table.first_latency
}
/// 切换NAT类型,不同的nat打洞模式会有不同
pub fn switch(
&self,
nat_type: NatType,
udp_socket_sender: &AcceptSocketSender<Option<Vec<mio::net::UdpSocket>>>,
) -> io::Result<()> {
let mut write_guard = self.sub_udp_socket.write();
match nat_type {
NatType::Symmetric => {
if !write_guard.is_empty() {
return Ok(());
}
let mut vec = Vec::with_capacity(SYMMETRIC_CHANNEL_NUM);
for _ in 0..SYMMETRIC_CHANNEL_NUM {
let udp = UdpSocket::bind("0.0.0.0:0")?;
//副通道使用异步io
udp.set_nonblocking(true)?;
vec.push(udp);
}
let mut mio_vec = Vec::with_capacity(SYMMETRIC_CHANNEL_NUM);
for udp in vec.iter() {
let udp_socket = mio::net::UdpSocket::from_std(udp.try_clone()?);
mio_vec.push(udp_socket);
}
udp_socket_sender.try_add_socket(Some(mio_vec))?;
*write_guard = vec;
}
NatType::Cone => {
if write_guard.is_empty() {
return Ok(());
}
udp_socket_sender.try_add_socket(None)?;
*write_guard = Vec::new();
}
}
Ok(())
}
pub fn channel_num(&self) -> usize {
self.main_udp_socket.len()
}
/// 获取核心udp监听的端口,用于其他客户端连接
pub fn main_local_udp_port(&self) -> io::Result<Vec<u16>> {
let mut ports = Vec::new();
for udp in self.main_udp_socket.iter() {
ports.push(udp.local_addr()?.port())
}
Ok(ports)
}
pub fn send_tcp(&self, buf: &[u8], addr: SocketAddr) -> io::Result<()> {
if let Some(tcp) = self.tcp_map.read().get(&addr) {
tcp.try_send(buf)
} else {
Err(io::Error::from(io::ErrorKind::NotFound))
}
}
pub fn send_main_udp(&self, index: usize, buf: &[u8], mut addr: SocketAddr) -> io::Result<()> {
if self.use_ipv6 {
//如果是v4地址则需要转换成v6
if let SocketAddr::V4(ipv4) = addr {
addr = SocketAddr::V6(SocketAddrV6::new(
ipv4.ip().to_ipv6_mapped(),
ipv4.port(),
0,
0,
));
}
}
self.main_udp_socket[index].send_to(buf, addr)?;
Ok(())
}
/// 将数据发送到默认通道,一般发往服务器才用此方法
pub fn send_default(&self, buf: &[u8], addr: SocketAddr) -> io::Result<()> {
if self.is_tcp {
//服务端地址只在重连时检测变化
self.send_tcp(buf, addr)
} else {
self.send_main_udp(self.main_index.load(Ordering::Relaxed), buf, addr)
}
}
pub fn change_main_index(&self) {
let index = (self.main_index.load(Ordering::Relaxed) + 1) % self.main_udp_socket.len();
self.main_index.store(index, Ordering::Relaxed);
}
/// 此方法仅用于对称网络打洞
pub fn try_send_all(&self, buf: &[u8], addr: SocketAddr) {
self.try_send_all_main(buf, addr);
for udp in self.sub_udp_socket.read().iter() {
if let Err(e) = udp.send_to(buf, addr) {
log::warn!("{:?},add={:?}", e, addr);
}
thread::sleep(Duration::from_millis(1));
}
}
pub fn try_send_all_main(&self, buf: &[u8], addr: SocketAddr) {
for index in 0..self.channel_num() {
if let Err(e) = self.send_main_udp(index, buf, addr) {
log::warn!("{:?},add={:?}", e, addr);
}
}
}
/// 发送网络数据
pub fn send_ipv4_by_id(
&self,
buf: &[u8],
id: &Ipv4Addr,
server_addr: SocketAddr,
send_default: bool,
) -> io::Result<()> {
if self.packet_loss_rate > 0 {
if rand::thread_rng().gen_ratio(self.packet_loss_rate, PACKET_LOSS_RATE_DENOMINATOR) {
return Ok(());
}
}
if self.packet_delay > 0 {
thread::sleep(Duration::from_millis(self.packet_delay as _));
}
//优先发到直连到地址
if let Err(e) = self.send_by_id(buf, id) {
if e.kind() != io::ErrorKind::NotFound {
log::warn!("{}:{:?}", id, e);
}
if !self.route_table.use_channel_type.is_only_p2p() && send_default {
//符合条件再发到服务器转发
self.send_default(buf, server_addr)?;
}
}
Ok(())
}
/// 将数据发到指定id
pub fn send_by_id(&self, buf: &[u8], id: &Ipv4Addr) -> io::Result<()> {
let mut c = 0;
loop {
let route = self.route_table.get_route_by_id(c, id)?;
return if let Err(e) = self.send_by_key(buf, route.route_key()) {
//降低发送速率
if e.kind() == io::ErrorKind::WouldBlock {
c += 1;
if c < 10 {
thread::sleep(Duration::from_micros(200));
continue;
}
}
Err(e)
} else {
Ok(())
};
}
}
/// 将数据发到指定路由
pub fn send_by_key(&self, buf: &[u8], route_key: RouteKey) -> io::Result<()> {
if route_key.is_tcp {
self.send_tcp(buf, route_key.addr)
} else {
if let Some(main_udp) = self.main_udp_socket.get(route_key.index) {
main_udp.send_to(buf, route_key.addr)?;
} else {
if let Some(udp) = self
.sub_udp_socket
.read()
.get(route_key.index - self.main_udp_socket.len())
{
udp.send_to(buf, route_key.addr)?;
} else {
Err(io::Error::from(io::ErrorKind::NotFound))?
}
}
Ok(())
}
}
pub fn remove_route(&self, ip: &Ipv4Addr, route_key: RouteKey) {
if self.route_table.remove_route(ip, route_key) {
if route_key.is_tcp {
if let Some(tcp) = self.tcp_map.write().remove(&route_key.addr) {
if let Err(e) = tcp.shutdown() {
log::warn!("{:?}", e);
}
}
}
}
}
}
pub struct RouteTable {
pub(crate) route_table:
RwLock<HashMap<Ipv4Addr, (AtomicUsize, Vec<(Route, AtomicCell<Instant>)>)>>,
first_latency: bool,
channel_num: usize,
use_channel_type: UseChannelType,
}
impl RouteTable {
fn new(use_channel_type: UseChannelType, first_latency: bool, channel_num: usize) -> Self {
Self {
route_table: RwLock::new(HashMap::with_capacity(64)),
use_channel_type,
first_latency,
channel_num,
}
}
}
impl RouteTable {
fn get_route_by_id(&self, index: usize, id: &Ipv4Addr) -> io::Result<Route> {
if let Some((_count, v)) = self.route_table.read().get(id) {
if self.first_latency {
if let Some((route, _)) = v.first() {
return Ok(*route);
}
} else {
let len = v.len();
if len != 0 {
return Ok(v[index % len].0);
}
}
}
Err(io::Error::new(io::ErrorKind::NotFound, "route not found"))
}
pub fn add_route_if_absent(&self, id: Ipv4Addr, route: Route) {
self.add_route_(id, route, true)
}
pub fn add_route(&self, id: Ipv4Addr, route: Route) {
self.add_route_(id, route, false)
}
fn add_route_(&self, id: Ipv4Addr, route: Route, only_if_absent: bool) {
// 限制通道类型
match self.use_channel_type {
UseChannelType::P2p => {
if !route.is_p2p() {
return;
}
}
_ => {}
}
let key = route.route_key();
let mut route_table = self.route_table.write();
let (_, list) = route_table
.entry(id)
.or_insert_with(|| (AtomicUsize::new(0), Vec::with_capacity(4)));
let mut exist = false;
for (x, time) in list.iter_mut() {
if x.metric < route.metric && !self.first_latency {
//非优先延迟的情况下 不能比当前的路径更长
return;
}
if x.route_key() == key {
if only_if_absent {
return;
}
x.metric = route.metric;
x.rt = route.rt;
exist = true;
time.store(Instant::now());
break;
}
}
if exist {
// 这个排序还有待优化,因为后加入的大概率排最后,被直接淘汰的概率也大,可能导致更好的通道被移除了
list.sort_by_key(|(k, _)| k.rt);
//如果延迟都稳定了,则去除多余通道
for (route, _) in list.iter() {
if route.rt == DEFAULT_RT {
return;
}
}
//延迟优先模式需要更多的通道探测延迟最低的路线
let limit_len = if self.first_latency {
self.channel_num + 2
} else {
self.channel_num
};
self.truncate_(list, limit_len);
} else {
if !self.first_latency {
if route.is_p2p() {
//非优先延迟的情况下 添加了直连的则排除非直连的
list.retain(|(k, _)| k.is_p2p());
}
};
//增加路由表容量,避免波动
let limit_len = self.channel_num * 2;
list.sort_by_key(|(k, _)| k.rt);
self.truncate_(list, limit_len);
list.push((route, AtomicCell::new(Instant::now())));
}
}
fn truncate_(&self, list: &mut Vec<(Route, AtomicCell<Instant>)>, len: usize) {
if list.len() <= len {
return;
}
if self.first_latency {
//找到第一个p2p通道
if let Some(index) =
list.iter()
.enumerate()
.find_map(|(index, (route, _))| if route.is_p2p() { Some(index) } else { None })
{
if index >= len {
//保留第一个p2p通道
let route = list.remove(index);
list.truncate(len - 1);
list.push(route);
return;
}
}
}
list.truncate(len);
}
pub fn route(&self, id: &Ipv4Addr) -> Option<Vec<Route>> {
if let Some((_, v)) = self.route_table.read().get(id) {
Some(v.iter().map(|(i, _)| *i).collect())
} else {
None
}
}
pub fn route_one(&self, id: &Ipv4Addr) -> Option<Route> {
if let Some((_, v)) = self.route_table.read().get(id) {
v.first().map(|(i, _)| *i)
} else {
None
}
}
pub fn route_one_p2p(&self, id: &Ipv4Addr) -> Option<Route> {
if let Some((_, v)) = self.route_table.read().get(id) {
for (i, _) in v {
if i.is_p2p() {
return Some(*i);
}
}
}
None
}
pub fn route_to_id(&self, route_key: &RouteKey) -> Option<Ipv4Addr> {
let table = self.route_table.read();
for (k, (_, v)) in table.iter() {
for (route, _) in v {
if &route.route_key() == route_key && route.is_p2p() {
return Some(*k);
}
}
}
None
}
pub fn need_punch(&self, id: &Ipv4Addr) -> bool {
if let Some((_, v)) = self.route_table.read().get(id) {
//存在p2p的通道则不再打洞
if v.iter().filter(|(k, _)| k.is_p2p()).count() >= 1 {
return false;
}
}
true
}
/// 返回所有路由
pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec<Route>)> {
let table = self.route_table.read();
table
.iter()
.map(|(k, (_, v))| (k.clone(), v.iter().map(|(i, _)| *i).collect()))
.collect()
}
pub fn route_table_p2p(&self) -> Vec<(Ipv4Addr, Route)> {
let table = self.route_table.read();
let mut list = Vec::with_capacity(8);
for (ip, (_, routes)) in table.iter() {
for (route, _) in routes.iter() {
if route.is_p2p() {
list.push((*ip, *route));
break;
}
}
}
list
}
pub fn route_table_one(&self) -> Vec<(Ipv4Addr, Route)> {
let mut list = Vec::with_capacity(8);
let table = self.route_table.read();
for (k, (_, v)) in table.iter() {
if let Some((route, _)) = v.first() {
list.push((*k, *route));
}
}
list
}
pub fn remove_route(&self, id: &Ipv4Addr, route_key: RouteKey) -> bool {
let mut write_guard = self.route_table.write();
if let Some((_, routes)) = write_guard.get_mut(id) {
routes.retain(|(x, _)| x.route_key() != route_key);
if routes.is_empty() {
write_guard.remove(id);
true
} else {
false
}
} else {
return true;
}
}
/// 更新路由入栈包的时刻,长时间没有收到数据的路由将会被剔除
pub fn update_read_time(&self, id: &Ipv4Addr, route_key: &RouteKey) {
if let Some((_, routes)) = self.route_table.read().get(id) {
for (route, time) in routes {
if &route.route_key() == route_key {
time.store(Instant::now());
break;
}
}
}
}
}
+6
View File
@@ -0,0 +1,6 @@
use crate::channel::context::Context;
use crate::channel::RouteKey;
pub trait RecvChannelHandler: Clone + Send + 'static {
fn handle(&mut self, buf: &mut [u8], route_key: RouteKey, context: &Context);
}
+23 -21
View File
@@ -1,10 +1,9 @@
use crate::channel::channel::Context;
use crate::channel::RouteKey;
use std::io;
use std::io::{Error, ErrorKind};
use std::net::Ipv4Addr;
use std::time::Duration;
use crate::channel::context::Context;
use crate::channel::Route;
pub struct Idle {
read_idle: Duration,
context: Context,
@@ -16,28 +15,31 @@ impl Idle {
}
}
pub enum IdleType {
Timeout(Ipv4Addr, Route),
Sleep(Duration),
None,
}
impl Idle {
/// 获取空闲路由
pub async fn next_idle(&self) -> io::Result<(Ipv4Addr, RouteKey)> {
loop {
let mut max = Duration::from_secs(0);
for entry in self.context.inner.route_table_time.iter() {
let last_read = entry.value().elapsed();
pub fn next_idle(&self) -> IdleType {
let mut max = Duration::from_secs(0);
let read_guard = self.context.route_table.route_table.read();
if read_guard.is_empty() {
return IdleType::None;
}
for (ip, (_, routes)) in read_guard.iter() {
for (route, time) in routes {
let last_read = time.load().elapsed();
if last_read >= self.read_idle {
return Ok((entry.key().1.clone(), entry.key().0.clone()));
} else {
if max < last_read {
max = last_read;
}
return IdleType::Timeout(*ip, *route);
} else if max < last_read {
max = last_read;
}
}
if self.read_idle > max {
let sleep_time = self.read_idle - max;
tokio::time::sleep(sleep_time).await;
}
if self.context.is_close() {
return Err(Error::new(ErrorKind::Other, "closed"));
}
}
let sleep_time = self.read_idle - max;
return IdleType::Sleep(sleep_time);
}
}
+204 -10
View File
@@ -1,14 +1,58 @@
use std::net::SocketAddr;
use std::io;
use std::net::{SocketAddr, UdpSocket};
use std::str::FromStr;
pub mod channel;
use crate::channel::context::Context;
use crate::channel::handler::RecvChannelHandler;
use crate::channel::sender::AcceptSocketSender;
use crate::channel::tcp_channel::tcp_listen;
use crate::channel::udp_channel::udp_listen;
use crate::util::{io_convert, StopManager};
pub mod context;
pub mod handler;
pub mod idle;
pub mod notify;
pub mod punch;
pub mod sender;
pub mod tcp_channel;
pub mod udp_channel;
const TCP_ID: usize = 0;
const UDP_ID: usize = 1;
const UDP_V6_ID: usize = 2;
const BUFFER_SIZE: usize = 1024 * 16;
#[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub enum UseChannelType {
Relay,
P2p,
All,
}
impl UseChannelType {
pub fn is_only_relay(&self) -> bool {
self == &UseChannelType::Relay
}
pub fn is_only_p2p(&self) -> bool {
self == &UseChannelType::P2p
}
pub fn is_all(&self) -> bool {
self == &UseChannelType::All
}
}
impl FromStr for UseChannelType {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().trim() {
"relay" => Ok(UseChannelType::Relay),
"p2p" => Ok(UseChannelType::P2p),
"all" => Ok(UseChannelType::All),
_ => Err(format!("not match '{}', enum: relay/p2p/all", s)),
}
}
}
impl Default for UseChannelType {
fn default() -> Self {
UseChannelType::All
}
}
#[derive(Copy, Clone, Eq, PartialEq)]
pub enum Status {
Cone,
@@ -18,6 +62,7 @@ pub enum Status {
#[derive(Copy, Clone, Debug)]
pub struct Route {
pub is_tcp: bool,
index: usize,
pub addr: SocketAddr,
pub metric: u8,
@@ -29,10 +74,11 @@ pub struct RouteSortKey {
pub metric: u8,
pub rt: i64,
}
const DEFAULT_RT: i64 = 999;
impl Route {
pub fn new(index: usize, addr: SocketAddr, metric: u8, rt: i64) -> Self {
pub fn new(is_tcp: bool, index: usize, addr: SocketAddr, metric: u8, rt: i64) -> Self {
Self {
is_tcp,
index,
addr,
metric,
@@ -41,14 +87,25 @@ impl Route {
}
pub fn from(route_key: RouteKey, metric: u8, rt: i64) -> Self {
Self {
is_tcp: route_key.is_tcp,
index: route_key.index,
addr: route_key.addr,
metric,
rt,
}
}
pub fn from_default_rt(route_key: RouteKey, metric: u8) -> Self {
Self {
is_tcp: route_key.is_tcp,
index: route_key.index,
addr: route_key.addr,
metric,
rt: DEFAULT_RT,
}
}
pub fn route_key(&self) -> RouteKey {
RouteKey {
is_tcp: self.is_tcp,
index: self.index,
addr: self.addr,
}
@@ -66,15 +123,152 @@ impl Route {
#[derive(Copy, Clone, Ord, PartialOrd, Eq, PartialEq, Hash, Debug)]
pub struct RouteKey {
is_tcp: bool,
index: usize,
pub addr: SocketAddr,
}
impl RouteKey {
pub(crate) fn new(index: usize, addr: SocketAddr) -> Self {
Self { index, addr }
pub(crate) fn new(is_tcp: bool, index: usize, addr: SocketAddr) -> Self {
Self {
is_tcp,
index,
addr,
}
}
pub fn is_tcp(&self) -> bool {
self.index == TCP_ID
self.is_tcp
}
pub fn index(&self) -> usize {
self.index
}
}
pub fn init_context(
ports: Vec<u16>,
use_channel_type: UseChannelType,
first_latency: bool,
is_tcp: bool,
packet_loss_rate: Option<f64>,
packet_delay: u32,
) -> io::Result<(Context, mio::net::TcpListener)> {
assert!(!ports.is_empty(), "not channel");
let mut udps = Vec::with_capacity(ports.len());
//检查系统是否支持ipv6
let use_ipv6 = match socket2::Socket::new(socket2::Domain::IPV6, socket2::Type::DGRAM, None) {
Ok(_) => true,
Err(e) => {
log::warn!("{:?}", e);
false
}
};
for port in &ports {
//监听v6+v4双栈
let (socket, address) = if use_ipv6 {
let address: SocketAddr = format!("[::]:{}", port).parse().unwrap();
let socket = socket2::Socket::new(socket2::Domain::IPV6, socket2::Type::DGRAM, None)?;
io_convert(socket.set_only_v6(false), |_| {
format!("set_only_v6 failed: {}", &address)
})?;
(socket, address)
} else {
let address: SocketAddr = format!("0.0.0.0:{}", port).parse().unwrap();
(
socket2::Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None)?,
address,
)
};
io_convert(socket.set_reuse_address(true), |_| {
format!("set_reuse_address failed: {}", &address)
})?;
io_convert(socket.set_send_buffer_size(2 * 1024 * 1024), |_| {
format!("set_send_buffer_size failed: {}", &address)
})?;
io_convert(socket.set_recv_buffer_size(2 * 1024 * 1024), |_| {
format!("set_recv_buffer_size failed: {}", &address)
})?;
io_convert(socket.bind(&address.into()), |_| {
format!("bind failed: {}", &address)
})?;
let main_channel: UdpSocket = socket.into();
main_channel.set_nonblocking(true)?;
udps.push(main_channel);
}
let context = Context::new(
udps,
use_channel_type,
first_latency,
is_tcp,
packet_loss_rate,
packet_delay,
use_ipv6,
);
let port = context.main_local_udp_port()?[0];
//监听v6+v4双栈,tcp通道使用异步io
let (socket, address) = if use_ipv6 {
let address: SocketAddr = format!("[::]:{}", port).parse().unwrap();
let socket = socket2::Socket::new(socket2::Domain::IPV6, socket2::Type::STREAM, None)?;
io_convert(socket.set_only_v6(false), |_| {
format!("set_only_v6 failed: {}", &address)
})?;
(socket, address)
} else {
let address: SocketAddr = format!("0.0.0.0:{}", port).parse().unwrap();
let socket = socket2::Socket::new(socket2::Domain::IPV4, socket2::Type::STREAM, None)?;
(socket, address)
};
io_convert(socket.set_reuse_address(true), |_| {
format!("set_reuse_address failed: {}", &address)
})?;
if let Err(e) = socket.bind(&address.into()) {
if ports[0] == 0 {
//端口可能冲突,则使用任意端口
log::warn!("监听tcp端口失败 {:?},重试一次", address);
let address: SocketAddr = if use_ipv6 {
format!("[::]:{}", 0).parse().unwrap()
} else {
format!("0.0.0.0:{}", port).parse().unwrap()
};
io_convert(socket.bind(&address.into()), |_| {
format!("bind failed: {}", &address)
})?;
} else {
//手动指定的ip,直接报错
io_convert(Err(e), |_| format!("bind failed: {}", &address))?;
}
}
socket.listen(128)?;
socket.set_nonblocking(true)?;
socket.set_nodelay(false)?;
let tcp_listener = mio::net::TcpListener::from_std(socket.into());
Ok((context, tcp_listener))
}
pub fn init_channel<H>(
tcp_listener: mio::net::TcpListener,
context: Context,
stop_manager: StopManager,
recv_handler: H,
) -> io::Result<(
AcceptSocketSender<Option<Vec<mio::net::UdpSocket>>>,
AcceptSocketSender<(mio::net::TcpStream, SocketAddr, Option<Vec<u8>>)>,
)>
where
H: RecvChannelHandler,
{
// udp监听,udp_socket_sender 用于NAT类型切换
let udp_socket_sender =
udp_listen(stop_manager.clone(), recv_handler.clone(), context.clone())?;
// 建立tcp监听,tcp_socket_sender 用于tcp 直连
let tcp_socket_sender = tcp_listen(
tcp_listener,
stop_manager.clone(),
recv_handler.clone(),
context.clone(),
)?;
Ok((udp_socket_sender, tcp_socket_sender))
}
+126
View File
@@ -0,0 +1,126 @@
use mio::{Token, Waker};
use parking_lot::Mutex;
use std::io;
use std::ops::Deref;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
#[derive(Clone)]
pub struct WritableNotify {
inner: Arc<WritableNotifyInner>,
}
impl WritableNotify {
pub fn new(waker: Waker) -> Self {
Self {
inner: Arc::new(WritableNotifyInner {
waker,
state: AtomicUsize::new(0),
tokens: Mutex::new(Vec::with_capacity(8)),
}),
}
}
}
impl Deref for WritableNotify {
type Target = WritableNotifyInner;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
pub struct WritableNotifyInner {
waker: Waker,
state: AtomicUsize,
tokens: Mutex<Vec<(Token, bool)>>,
}
impl WritableNotifyInner {
pub fn notify(&self, token: Token, state: bool) -> io::Result<()> {
{
let mut guard = self.tokens.lock();
if guard.is_empty() || !guard.contains(&(token, state)) {
guard.push((token, state));
}
drop(guard);
}
self.need_write()
}
pub fn stop(&self) -> io::Result<()> {
self.state.store(0b001, Ordering::Release);
self.waker.wake()
}
pub fn need_write(&self) -> io::Result<()> {
self.state.fetch_or(0b010, Ordering::AcqRel);
self.waker.wake()
}
pub fn add_socket(&self) -> io::Result<()> {
self.state.fetch_or(0b100, Ordering::AcqRel);
self.waker.wake()
}
pub fn take_all(&self) -> Option<Vec<(Token, bool)>> {
let mut guard = self.tokens.lock();
if guard.is_empty() {
None
} else {
Some(guard.drain(..).collect())
}
}
pub fn is_stop(&self) -> bool {
self.state.load(Ordering::Acquire) & 0b001 == 0b001
}
pub fn is_need_write(&self) -> bool {
self.state.fetch_and(!0b010, Ordering::AcqRel) & 0b010 == 0b010
}
pub fn is_add_socket(&self) -> bool {
self.state.fetch_and(!0b100, Ordering::AcqRel) & 0b100 == 0b100
}
}
#[derive(Clone)]
pub struct AcceptNotify {
inner: Arc<AcceptNotifyInner>,
}
impl AcceptNotify {
pub fn new(waker: Waker) -> Self {
Self {
inner: Arc::new(AcceptNotifyInner {
waker,
state: AtomicUsize::new(0),
}),
}
}
}
impl Deref for AcceptNotify {
type Target = AcceptNotifyInner;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
pub struct AcceptNotifyInner {
waker: Waker,
state: AtomicUsize,
}
impl AcceptNotifyInner {
pub fn is_stop(&self) -> bool {
self.state.load(Ordering::Acquire) & 0b001 == 0b001
}
pub fn is_add_socket(&self) -> bool {
self.state.fetch_and(!0b100, Ordering::AcqRel) & 0b100 == 0b100
}
pub fn stop(&self) -> io::Result<()> {
self.state.store(0b001, Ordering::Release);
self.waker.wake()
}
pub fn add_socket(&self) -> io::Result<()> {
self.state.fetch_or(0b100, Ordering::AcqRel);
self.waker.wake()
}
}
+232 -66
View File
@@ -1,12 +1,15 @@
use std::collections::HashMap;
use std::io;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4, SocketAddrV6};
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6};
use std::str::FromStr;
use std::time::Duration;
use std::{io, thread};
use mio::net::TcpStream;
use rand::prelude::SliceRandom;
use rand::Rng;
use crate::channel::channel::Context;
use crate::channel::context::Context;
use crate::channel::sender::AcceptSocketSender;
#[derive(Copy, Clone, Eq, PartialEq, Debug)]
pub enum PunchModel {
@@ -22,19 +25,28 @@ impl FromStr for PunchModel {
match s.to_lowercase().trim() {
"ipv4" => Ok(PunchModel::IPv4),
"ipv6" => Ok(PunchModel::IPv6),
_ => Ok(PunchModel::All),
"all" => Ok(PunchModel::All),
_ => Err(format!("not match '{}', enum: ipv4/ipv6/all", s)),
}
}
}
impl Default for PunchModel {
fn default() -> Self {
PunchModel::All
}
}
#[derive(Clone, Debug)]
pub struct NatInfo {
pub public_ips: Vec<Ipv4Addr>,
pub public_port: u16,
pub public_ports: Vec<u16>,
pub public_port_range: u16,
pub local_ipv4_addr: SocketAddrV4,
pub ipv6_addr: SocketAddrV6,
pub nat_type: NatType,
pub(crate) local_ipv4: Option<Ipv4Addr>,
pub(crate) ipv6: Option<Ipv6Addr>,
pub(crate) udp_ports: Vec<u16>,
pub tcp_port: u16,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug, Hash)]
@@ -46,22 +58,119 @@ pub enum NatType {
impl NatInfo {
pub fn new(
mut public_ips: Vec<Ipv4Addr>,
public_port: u16,
public_ports: Vec<u16>,
public_port_range: u16,
local_ipv4_addr: SocketAddrV4,
ipv6_addr: SocketAddrV6,
nat_type: NatType,
mut local_ipv4: Option<Ipv4Addr>,
mut ipv6: Option<Ipv6Addr>,
udp_ports: Vec<u16>,
tcp_port: u16,
mut nat_type: NatType,
) -> Self {
public_ips.retain(|ip| !ip.is_loopback() && !ip.is_private());
public_ips.retain(|ip| {
!ip.is_multicast()
&& !ip.is_broadcast()
&& !ip.is_unspecified()
&& !ip.is_loopback()
&& !ip.is_private()
});
if public_ips.len() > 1 {
nat_type = NatType::Symmetric;
}
if let Some(ip) = local_ipv4 {
if ip.is_multicast() || ip.is_broadcast() || ip.is_unspecified() || ip.is_loopback() {
local_ipv4 = None
}
}
if let Some(ip) = ipv6 {
if ip.is_multicast() || ip.is_unspecified() || ip.is_loopback() {
ipv6 = None
}
}
Self {
public_ips,
public_port,
public_ports,
public_port_range,
local_ipv4_addr,
ipv6_addr,
local_ipv4,
ipv6,
udp_ports,
tcp_port,
nat_type,
}
}
pub fn update_addr(&mut self, index: usize, ip: Ipv4Addr, port: u16) {
if port != 0 {
if let Some(public_port) = self.public_ports.get_mut(index) {
*public_port = port;
}
}
if !ip.is_multicast()
&& !ip.is_broadcast()
&& !ip.is_unspecified()
&& !ip.is_loopback()
&& !ip.is_private()
{
if !self.public_ips.contains(&ip) {
self.public_ips.push(ip);
}
}
}
pub fn local_ipv4(&self) -> Option<Ipv4Addr> {
self.local_ipv4
}
pub fn ipv6(&self) -> Option<Ipv6Addr> {
self.ipv6
}
pub fn local_udp_ipv4addr(&self, index: usize) -> Option<SocketAddr> {
let len = self.udp_ports.len();
if len == 0 {
return None;
}
if let Some(local_ipv4) = self.local_ipv4 {
Some(SocketAddr::V4(SocketAddrV4::new(
local_ipv4,
self.udp_ports[index % len],
)))
} else {
None
}
}
pub fn local_udp_ipv6addr(&self, index: usize) -> Option<SocketAddr> {
let len = self.udp_ports.len();
if len == 0 {
return None;
}
if let Some(ipv6) = self.ipv6 {
Some(SocketAddr::V6(SocketAddrV6::new(
ipv6,
self.udp_ports[index % len],
0,
0,
)))
} else {
None
}
}
pub fn local_tcp_ipv6addr(&self) -> Option<SocketAddr> {
if self.tcp_port == 0 {
return None;
}
if let Some(ipv6) = self.ipv6 {
Some(SocketAddr::V6(SocketAddrV6::new(ipv6, self.tcp_port, 0, 0)))
} else {
None
}
}
pub fn local_tcp_ipv4addr(&self) -> Option<SocketAddr> {
if self.tcp_port == 0 {
return None;
}
if let Some(ipv4) = self.local_ipv4 {
Some(SocketAddr::V4(SocketAddrV4::new(ipv4, self.tcp_port)))
} else {
None
}
}
}
#[derive(Clone)]
@@ -70,10 +179,17 @@ pub struct Punch {
port_vec: Vec<u16>,
port_index: HashMap<Ipv4Addr, usize>,
punch_model: PunchModel,
is_tcp: bool,
tcp_socket_sender: AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
}
impl Punch {
pub fn new(context: Context, punch_model: PunchModel) -> Self {
pub fn new(
context: Context,
punch_model: PunchModel,
is_tcp: bool,
tcp_socket_sender: AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
) -> Self {
let mut port_vec: Vec<u16> = (1..65535).collect();
port_vec.push(65535);
let mut rng = rand::thread_rng();
@@ -83,30 +199,73 @@ impl Punch {
port_vec,
port_index: HashMap::new(),
punch_model,
is_tcp,
tcp_socket_sender,
}
}
}
impl Punch {
pub async fn punch(&mut self, buf: &[u8], id: Ipv4Addr, nat_info: NatInfo) -> io::Result<()> {
if !self.context.need_punch(&id) {
fn connect_tcp(&self, buf: &[u8], addr: SocketAddr) -> bool {
// mio是非阻塞的,不能立马判断是否能连接成功,所以用标准库的tcp
match std::net::TcpStream::connect_timeout(&addr, Duration::from_secs(3)) {
Ok(tcp_stream) => {
if tcp_stream.set_nonblocking(true).is_err() {
return false;
}
return self
.tcp_socket_sender
.try_add_socket((TcpStream::from_std(tcp_stream), addr, Some(buf.to_vec())))
.is_ok();
}
Err(e) => {
log::warn!("连接到tcp失败,addr={},err={}", addr, e);
}
}
false
}
pub fn punch(&mut self, buf: &[u8], id: Ipv4Addr, nat_info: NatInfo) -> io::Result<()> {
if !self.context.route_table.need_punch(&id) {
log::info!("已打洞成功,无需打洞:{:?}", id);
return Ok(());
}
if !nat_info.local_ipv4_addr.ip().is_unspecified() && nat_info.local_ipv4_addr.port() != 0 {
let _ = self
.context
.send_main_udp(buf, SocketAddr::V4(nat_info.local_ipv4_addr));
if self.is_tcp && nat_info.tcp_port != 0 {
//向tcp发起连接
if let Some(ipv6_addr) = nat_info.local_tcp_ipv6addr() {
if self.connect_tcp(buf, ipv6_addr) {
return Ok(());
}
}
//向tcp发起连接
if let Some(ipv4_addr) = nat_info.local_tcp_ipv4addr() {
if self.connect_tcp(buf, ipv4_addr) {
return Ok(());
}
}
if nat_info.nat_type == NatType::Cone && nat_info.public_ips.len() == 1 {
let addr =
SocketAddr::V4(SocketAddrV4::new(nat_info.public_ips[0], nat_info.tcp_port));
if self.connect_tcp(buf, addr) {
return Ok(());
}
}
}
if self.punch_model != PunchModel::IPv4
&& !nat_info.ipv6_addr.ip().is_unspecified()
&& nat_info.ipv6_addr.port() != 0
{
let rs = self
.context
.send_main_udp(buf, SocketAddr::V6(nat_info.ipv6_addr));
log::info!("发送到ipv6地址:{:?},rs={:?}", nat_info.ipv6_addr, rs);
if rs.is_ok() && self.punch_model == PunchModel::IPv6 {
return Ok(());
let channel_num = self.context.channel_num();
for index in 0..channel_num {
if let Some(ipv4_addr) = nat_info.local_udp_ipv4addr(index) {
let _ = self.context.send_main_udp(index, buf, ipv4_addr);
}
}
if self.punch_model != PunchModel::IPv4 {
for index in 0..channel_num {
if let Some(ipv6_addr) = nat_info.local_udp_ipv6addr(index) {
let rs = self.context.send_main_udp(index, buf, ipv6_addr);
log::info!("发送到ipv6地址:{:?},rs={:?}", ipv6_addr, rs);
if rs.is_ok() && self.punch_model == PunchModel::IPv6 {
return Ok(());
}
}
}
}
match nat_info.nat_type {
@@ -119,17 +278,16 @@ impl Punch {
//预测范围内最多发送max_k1个包
let max_k1 = 60;
//全局最多发送max_k2个包
let max_k2 = 800;
let max_k2 = rand::thread_rng().gen_range(600..800);
let port = nat_info.public_ports.get(0).map(|e| *e).unwrap_or(0);
if nat_info.public_port_range < max_k1 * 3 {
//端口变化不大时,在预测的范围内随机发送
let min_port = if nat_info.public_port > nat_info.public_port_range {
nat_info.public_port - nat_info.public_port_range
let min_port = if port > nat_info.public_port_range {
port - nat_info.public_port_range
} else {
1
};
let (max_port, overflow) = nat_info
.public_port
.overflowing_add(nat_info.public_port_range);
let (max_port, overflow) = port.overflowing_add(nat_info.public_port_range);
let max_port = if overflow { 65535 } else { max_port };
let k = if max_port - min_port + 1 > max_k1 {
max_k1 as usize
@@ -138,64 +296,72 @@ impl Punch {
};
let mut nums: Vec<u16> = (min_port..max_port).collect();
nums.push(max_port);
{
let mut rng = rand::thread_rng();
nums.shuffle(&mut rng);
}
self.punch_symmetric(&nums[..k], buf, &nat_info.public_ips, max_k1 as usize)
.await?;
nums.shuffle(&mut rand::thread_rng());
self.punch_symmetric(&nums[..k], buf, &nat_info.public_ips, max_k1 as usize)?;
}
let start = *self.port_index.entry(id.clone()).or_insert(0);
let mut end = start + max_k2;
let mut index = end;
if end >= self.port_vec.len() {
if end > self.port_vec.len() {
end = self.port_vec.len();
}
let mut index = start
+ self.punch_symmetric(
&self.port_vec[start..end],
buf,
&nat_info.public_ips,
max_k2,
)?;
if index >= self.port_vec.len() {
index = 0
}
self.punch_symmetric(
&self.port_vec[start..end],
buf,
&nat_info.public_ips,
max_k2,
)
.await?;
self.port_index.insert(id, index);
}
NatType::Cone => {
let is_cone = self.context.is_cone();
for ip in nat_info.public_ips {
let addr = SocketAddr::V4(SocketAddrV4::new(ip, nat_info.public_port));
self.context.send_main_udp(buf, addr)?;
'a: for index in 0..nat_info.public_ports.len().min(channel_num) {
for ip in &nat_info.public_ips {
let port = nat_info.public_ports[index];
if port == 0 || ip.is_unspecified() {
continue;
}
let addr = SocketAddr::V4(SocketAddrV4::new(*ip, port));
if is_cone {
self.context.send_main_udp(index, buf, addr)?;
} else {
//只有一方是对称,则对称方要使用全部端口发送数据,符合上述计算的概率
self.context.try_send_all(buf, addr);
}
thread::sleep(Duration::from_millis(2));
}
if !is_cone {
//只有一方是对称,则对称方要使用全部端口发送数据,符合上述计算的概率
self.context.try_send_all(buf, addr)?;
//对称网络数据只发一遍
break 'a;
}
tokio::time::sleep(Duration::from_millis(2)).await;
}
}
}
Ok(())
}
async fn punch_symmetric(
fn punch_symmetric(
&self,
ports: &[u16],
buf: &[u8],
ips: &Vec<Ipv4Addr>,
max: usize,
) -> io::Result<()> {
) -> io::Result<usize> {
let mut count = 0;
for port in ports {
for (index, port) in ports.iter().enumerate() {
for pub_ip in ips {
count += 1;
if count == max {
return Ok(());
return Ok(index);
}
let addr = SocketAddr::V4(SocketAddrV4::new(*pub_ip, *port));
self.context.send_main_udp(buf, addr)?;
tokio::time::sleep(Duration::from_millis(2)).await;
self.context.send_main_udp(0, buf, addr)?;
thread::sleep(Duration::from_millis(2));
}
}
Ok(())
Ok(ports.len())
}
}
+85 -1
View File
@@ -1,5 +1,12 @@
use crate::channel::channel::Context;
use std::io;
use std::ops::Deref;
use std::sync::mpsc::{SyncSender, TrySendError};
use std::sync::Arc;
use mio::Token;
use crate::channel::context::Context;
use crate::channel::notify::{AcceptNotify, WritableNotify};
#[derive(Clone)]
pub struct ChannelSender {
@@ -19,3 +26,80 @@ impl Deref for ChannelSender {
&self.context
}
}
pub struct AcceptSocketSender<T> {
sender: SyncSender<T>,
notify: AcceptNotify,
}
impl<T> Clone for AcceptSocketSender<T> {
fn clone(&self) -> Self {
Self {
sender: self.sender.clone(),
notify: self.notify.clone(),
}
}
}
impl<T> AcceptSocketSender<T> {
pub fn new(notify: AcceptNotify, sender: SyncSender<T>) -> Self {
Self { sender, notify }
}
pub fn try_add_socket(&self, t: T) -> io::Result<()> {
match self.sender.try_send(t) {
Ok(_) => self.notify.add_socket(),
Err(e) => match e {
TrySendError::Full(_) => Err(io::Error::from(io::ErrorKind::WouldBlock)),
TrySendError::Disconnected(_) => Err(io::Error::from(io::ErrorKind::WriteZero)),
},
}
}
}
#[derive(Clone)]
pub struct PacketSender {
inner: Arc<PacketSenderInner>,
}
impl PacketSender {
pub fn new(notify: WritableNotify, buffer: SyncSender<Vec<u8>>, token: Token) -> Self {
Self {
inner: Arc::new(PacketSenderInner {
token,
notify,
buffer,
}),
}
}
#[inline]
pub fn try_send(&self, buf: &[u8]) -> io::Result<()> {
self.inner.try_send(buf)
}
pub fn shutdown(&self) -> io::Result<()> {
self.inner.shutdown()
}
}
pub struct PacketSenderInner {
token: Token,
notify: WritableNotify,
buffer: SyncSender<Vec<u8>>,
}
impl PacketSenderInner {
#[inline]
fn try_send(&self, buf: &[u8]) -> io::Result<()> {
let len = buf.len();
let mut buf_vec = Vec::with_capacity(buf.len() + 4);
buf_vec.extend_from_slice(&[0, 0, (len >> 8) as u8, (len & 0xFF) as u8]);
buf_vec.extend_from_slice(buf);
match self.buffer.try_send(buf_vec) {
Ok(_) => self.notify.notify(self.token, true),
Err(e) => match e {
TrySendError::Disconnected(_) => Err(io::Error::from(io::ErrorKind::WriteZero)),
TrySendError::Full(_) => Err(io::Error::from(io::ErrorKind::WouldBlock)),
},
}
}
fn shutdown(&self) -> io::Result<()> {
self.notify.notify(self.token, false)
}
}
+456
View File
@@ -0,0 +1,456 @@
use std::collections::HashMap;
use std::io::{Read, Write};
use std::net::{Shutdown, SocketAddr};
#[cfg(any(unix))]
use std::os::fd::FromRawFd;
#[cfg(any(unix))]
use std::os::fd::IntoRawFd;
#[cfg(windows)]
use std::os::windows::io::FromRawSocket;
#[cfg(windows)]
use std::os::windows::io::IntoRawSocket;
use std::sync::mpsc::{sync_channel, Receiver, SyncSender, TryRecvError, TrySendError};
use std::{io, thread};
use mio::net::{TcpListener, TcpStream};
use mio::{Events, Interest, Poll, Registry, Token, Waker};
use crate::channel::context::Context;
use crate::channel::handler::RecvChannelHandler;
use crate::channel::notify::{AcceptNotify, WritableNotify};
use crate::channel::sender::{AcceptSocketSender, PacketSender};
use crate::channel::{RouteKey, BUFFER_SIZE};
use crate::util::StopManager;
const SERVER: Token = Token(0);
const NOTIFY: Token = Token(1);
/// 监听tcp端口,等待客户端连接
pub fn tcp_listen<H>(
tcp_server: TcpListener,
stop_manager: StopManager,
recv_handler: H,
context: Context,
) -> io::Result<AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>>
where
H: RecvChannelHandler,
{
let (tcp_sender, tcp_receiver) = sync_channel(64);
let poll = Poll::new()?;
let waker = AcceptNotify::new(Waker::new(poll.registry(), NOTIFY)?);
let accept = AcceptSocketSender::new(waker.clone(), tcp_sender);
let worker = {
let waker = waker.clone();
stop_manager.add_listener("tcp_listen".into(), move || {
if let Err(e) = waker.stop() {
log::error!("{:?}", e);
}
})?
};
thread::Builder::new()
.name("tcpRead".into())
.spawn(move || {
if let Err(e) = tcp_listen0(
poll,
tcp_server,
&stop_manager,
waker,
tcp_receiver,
recv_handler,
context,
) {
log::error!("{:?}", e);
}
worker.stop_all();
})?;
Ok(accept)
}
fn tcp_listen0<H>(
mut poll: Poll,
mut tcp_server: TcpListener,
stop_manager: &StopManager,
accept_notify: AcceptNotify,
accept_tcp_receiver: Receiver<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
mut recv_handler: H,
context: Context,
) -> io::Result<()>
where
H: RecvChannelHandler,
{
let (tcp_sender, tcp_receiver) = sync_channel(64);
let write_waker = init_writable_handler(tcp_receiver, stop_manager.clone(), context.clone())?;
poll.registry()
.register(&mut tcp_server, SERVER, Interest::READABLE)?;
let mut events = Events::with_capacity(1024);
let mut read_map: HashMap<Token, (RouteKey, TcpStream, Box<[u8; BUFFER_SIZE]>, usize)> =
HashMap::with_capacity(32);
loop {
poll.poll(&mut events, None)?;
for event in events.iter() {
match event.token() {
SERVER => loop {
match tcp_server.accept() {
Ok((stream, addr)) => {
accept_handle(
stream,
addr,
None,
&write_waker,
&mut read_map,
&tcp_sender,
poll.registry(),
)?;
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
}
},
NOTIFY => {
if accept_notify.is_stop() {
return Ok(());
}
if accept_notify.is_add_socket() {
while let Ok((stream, addr, init_buf)) = accept_tcp_receiver.try_recv() {
accept_handle(
stream,
addr,
init_buf,
&write_waker,
&mut read_map,
&tcp_sender,
poll.registry(),
)?;
}
}
}
token => {
if event.is_readable() {
if let Err(e) =
readable_handle(&token, &mut read_map, &mut recv_handler, &context)
{
closed_handle_r(&token, &mut read_map);
log::warn!("{:?}", e);
if let Err(e) = write_waker.notify(token, false) {
log::warn!("{:?}", e);
}
}
} else {
closed_handle_r(&token, &mut read_map);
if let Err(e) = write_waker.notify(token, false) {
log::warn!("{:?}", e);
}
}
}
}
}
}
}
/// 处理写事件
fn init_writable_handler(
receiver: Receiver<(TcpStream, Token, SocketAddr, Option<Vec<u8>>)>,
stop_manager: StopManager,
context: Context,
) -> io::Result<WritableNotify> {
let poll = Poll::new()?;
let writable_notify = WritableNotify::new(Waker::new(poll.registry(), NOTIFY)?);
let worker = {
let writable_notify = writable_notify.clone();
stop_manager.add_listener("tcp_writable_handler".into(), move || {
if let Err(e) = writable_notify.stop() {
log::error!("{:?}", e);
}
})?
};
{
let writable_notify = writable_notify.clone();
thread::Builder::new()
.name("tcpWriteableListen".into())
.spawn(move || {
if let Err(e) = tcp_writable_listen(receiver, poll, writable_notify, &context) {
log::error!("{:?}", e);
}
worker.stop_all();
})?;
}
Ok(writable_notify)
}
/// 处理写事件
fn tcp_writable_listen(
receiver: Receiver<(TcpStream, Token, SocketAddr, Option<Vec<u8>>)>,
mut poll: Poll,
writable_notify: WritableNotify,
context: &Context,
) -> io::Result<()> {
let mut events = Events::with_capacity(1024);
let mut write_map: HashMap<
Token,
(
TcpStream,
SocketAddr,
Receiver<Vec<u8>>,
Option<(Vec<u8>, usize)>,
),
> = HashMap::with_capacity(32);
loop {
poll.poll(&mut events, None)?;
for event in events.iter() {
match event.token() {
NOTIFY => {
if writable_notify.is_stop() {
//服务停止
return Ok(());
}
if writable_notify.is_need_write() {
// 需要写入数据
if let Some(tokens) = writable_notify.take_all() {
for (token, state) in tokens {
if !state {
closed_handle_w(&token, &mut write_map, &context);
continue;
}
if let Err(e) = writable_handle(&token, &mut write_map) {
closed_handle_w(&token, &mut write_map, &context);
log::warn!("{:?}", e);
}
}
}
}
if writable_notify.is_add_socket() {
//添加tcp连接,并监听写事件
while let Ok((mut stream, token, addr, init_buf)) = receiver.try_recv() {
if let Err(e) = stream.set_nodelay(true) {
log::warn!("set_nodelay err={:?}", e);
}
if let Err(e) =
poll.registry()
.register(&mut stream, token, Interest::WRITABLE)
{
log::warn!("registry err={:?}", e);
continue;
}
let (sender, receiver) = sync_channel(128);
let packet_sender =
PacketSender::new(writable_notify.clone(), sender, token);
if let Some(init_buf) = init_buf {
packet_sender.try_send(&init_buf)?;
}
context.tcp_map.write().insert(addr, packet_sender);
write_map.insert(token, (stream, addr, receiver, None));
}
}
}
token => {
if event.is_writable() {
if let Err(e) = writable_handle(&token, &mut write_map) {
closed_handle_w(&token, &mut write_map, &context);
log::warn!("{:?}", e);
}
} else {
closed_handle_w(&token, &mut write_map, &context);
}
}
}
}
}
}
fn accept_handle(
stream: TcpStream,
addr: SocketAddr,
init_buf: Option<Vec<u8>>,
write_waker: &WritableNotify,
read_map: &mut HashMap<Token, (RouteKey, TcpStream, Box<[u8; BUFFER_SIZE]>, usize)>,
tcp_sender: &SyncSender<(TcpStream, Token, SocketAddr, Option<Vec<u8>>)>,
registry: &Registry,
) -> io::Result<()> {
#[cfg(windows)]
let (tcp_stream, index) = unsafe {
let fd = stream.into_raw_socket();
(std::net::TcpStream::from_raw_socket(fd), fd as usize)
};
#[cfg(any(unix))]
let (tcp_stream, index) = unsafe {
let fd = stream.into_raw_fd();
(std::net::TcpStream::from_raw_fd(fd), fd as usize)
};
if index == 0 || index == 1 {
log::error!("index err={:?}", addr);
return Ok(());
}
let token = Token(index);
match tcp_stream.try_clone() {
Ok(tcp_writer) => {
match tcp_sender.try_send((TcpStream::from_std(tcp_writer), token, addr, init_buf)) {
Ok(_) => {
if let Err(e) = write_waker.add_socket() {
log::error!("write_waker,err={:?},addr={:?}", e, addr);
return Ok(());
}
}
Err(e) => {
return match e {
TrySendError::Full(_) => {
log::error!("Full,addr={:?}", addr);
Ok(())
}
TrySendError::Disconnected(_) => {
Err(io::Error::new(io::ErrorKind::Other, "write thread exit"))
}
};
}
}
}
Err(e) => {
log::error!("try_clone err={:?},addr={:?}", e, addr);
return Ok(());
}
}
let mut stream = TcpStream::from_std(tcp_stream);
if let Err(e) = registry.register(&mut stream, token, Interest::READABLE) {
log::error!("registry err={:?},addr={:?}", e, addr);
return Ok(());
}
read_map.insert(
token,
(
RouteKey::new(true, index, addr),
stream,
Box::new([0; BUFFER_SIZE]),
0,
),
);
Ok(())
}
fn readable_handle<H>(
token: &Token,
map: &mut HashMap<Token, (RouteKey, TcpStream, Box<[u8; BUFFER_SIZE]>, usize)>,
recv_handler: &mut H,
context: &Context,
) -> io::Result<()>
where
H: RecvChannelHandler,
{
if let Some((route_key, stream, buf, begin)) = map.get_mut(token) {
loop {
let end = if *begin >= 4 {
4 + (((buf[2] as u16) << 8) | buf[3] as u16) as usize
} else {
4
};
if end > BUFFER_SIZE {
return Err(io::Error::from(io::ErrorKind::InvalidData));
}
match stream.read(&mut buf[*begin..end]) {
Ok(len) => {
if len == 0 {
return Err(io::Error::from(io::ErrorKind::UnexpectedEof));
}
*begin += len;
if end > 4 && *begin == end {
recv_handler.handle(&mut buf[4..end], *route_key, context);
*begin = 0;
}
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
}
}
}
Ok(())
}
fn writable_handle(
token: &Token,
map: &mut HashMap<
Token,
(
TcpStream,
SocketAddr,
Receiver<Vec<u8>>,
Option<(Vec<u8>, usize)>,
),
>,
) -> io::Result<()> {
if let Some((stream, _, receiver, last)) = map.get_mut(token) {
loop {
if let Some((buf, begin)) = last {
match stream.write(&buf[*begin..]) {
Ok(len) => {
if len == 0 {
return Err(io::Error::from(io::ErrorKind::WriteZero));
}
if len + *begin == buf.len() {
*last = None;
} else {
*begin += len;
continue;
}
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
}
}
match receiver.try_recv() {
Ok(buf) => *last = Some((buf, 0)),
Err(e) => match e {
TryRecvError::Empty => {
break;
}
TryRecvError::Disconnected => {
return Err(io::Error::from(io::ErrorKind::Other));
}
},
}
}
}
Ok(())
}
fn closed_handle_r(
token: &Token,
map: &mut HashMap<Token, (RouteKey, TcpStream, Box<[u8; BUFFER_SIZE]>, usize)>,
) {
if let Some((_, tcp, _, _)) = map.remove(token) {
let _ = tcp.shutdown(Shutdown::Both);
}
}
fn closed_handle_w(
token: &Token,
map: &mut HashMap<
Token,
(
TcpStream,
SocketAddr,
Receiver<Vec<u8>>,
Option<(Vec<u8>, usize)>,
),
>,
context: &Context,
) {
if let Some((tcp, addr, _, _)) = map.remove(token) {
context.tcp_map.write().remove(&addr);
let _ = tcp.shutdown(Shutdown::Both);
}
}
+284
View File
@@ -0,0 +1,284 @@
use std::collections::HashMap;
use std::sync::mpsc::{sync_channel, Receiver};
use std::sync::Arc;
use std::{io, thread};
use mio::event::Source;
use mio::net::UdpSocket;
use mio::{Events, Interest, Poll, Token, Waker};
use crate::channel::context::Context;
use crate::channel::handler::RecvChannelHandler;
use crate::channel::notify::AcceptNotify;
use crate::channel::sender::AcceptSocketSender;
use crate::channel::{RouteKey, BUFFER_SIZE};
use crate::util::StopManager;
pub fn udp_listen<H>(
stop_manager: StopManager,
recv_handler: H,
context: Context,
) -> io::Result<AcceptSocketSender<Option<Vec<UdpSocket>>>>
where
H: RecvChannelHandler,
{
main_udp_listen(stop_manager.clone(), recv_handler.clone(), context.clone())?;
sub_udp_listen(stop_manager, recv_handler, context)
}
const NOTIFY: Token = Token(0);
fn sub_udp_listen<H>(
stop_manager: StopManager,
recv_handler: H,
context: Context,
) -> io::Result<AcceptSocketSender<Option<Vec<UdpSocket>>>>
where
H: RecvChannelHandler,
{
let (udp_sender, udp_receiver) = sync_channel(64);
let poll = Poll::new()?;
let waker = AcceptNotify::new(Waker::new(poll.registry(), NOTIFY)?);
let worker = {
let waker = waker.clone();
stop_manager.add_listener("sub_udp_listen".into(), move || {
if let Err(e) = waker.stop() {
log::error!("{:?}", e);
}
})?
};
let accept = AcceptSocketSender::new(waker.clone(), udp_sender);
thread::Builder::new()
.name("subUdp".into())
.spawn(move || {
if let Err(e) = sub_udp_listen0(poll, recv_handler, context, waker, udp_receiver) {
log::error!("{:?}", e);
}
worker.stop_all();
})?;
Ok(accept)
}
fn sub_udp_listen0<H>(
mut poll: Poll,
mut recv_handler: H,
context: Context,
accept_notify: AcceptNotify,
accept_receiver: Receiver<Option<Vec<UdpSocket>>>,
) -> io::Result<()>
where
H: RecvChannelHandler,
{
let mut events = Events::with_capacity(1024);
let mut buf = [0; BUFFER_SIZE];
let mut read_map: HashMap<Token, UdpSocket> = HashMap::with_capacity(32);
loop {
poll.poll(&mut events, None)?;
for event in events.iter() {
match event.token() {
NOTIFY => {
if accept_notify.is_stop() {
return Ok(());
}
if accept_notify.is_add_socket() {
while let Ok(option) = accept_receiver.try_recv() {
match option {
None => {
log::info!("切换成锥形模式");
for (_, mut udp_socket) in read_map.drain() {
if let Err(e) = udp_socket.deregister(poll.registry()) {
log::error!("{:?}", e);
}
}
}
Some(socket_list) => {
log::info!("切换成对称模式 监听端口数:{}", socket_list.len());
for (index, mut udp_socket) in
socket_list.into_iter().enumerate()
{
let token = Token(index + context.channel_num());
poll.registry().register(
&mut udp_socket,
token,
Interest::READABLE,
)?;
read_map.insert(token, udp_socket);
}
}
}
}
}
}
token => {
if let Some(udp_socket) = read_map.get(&token) {
loop {
match udp_socket.recv_from(&mut buf) {
Ok((len, addr)) => {
recv_handler.handle(
&mut buf[..len],
RouteKey::new(false, token.0, addr),
&context,
);
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::error!("{:?}", e);
}
}
}
}
}
}
}
}
}
/// 阻塞监听
fn main_udp_listen<H>(
stop_manager: StopManager,
recv_handler: H,
context: Context,
) -> io::Result<()>
where
H: RecvChannelHandler,
{
let poll = Poll::new()?;
let waker = Arc::new(Waker::new(poll.registry(), NOTIFY)?);
let _waker = waker.clone();
let worker = stop_manager.add_listener("main_udp".into(), move || {
if let Err(e) = waker.wake() {
log::error!("{:?}", e);
}
})?;
thread::Builder::new()
.name("mainUdp".into())
.spawn(move || {
if let Err(e) = main_udp_listen0(poll, recv_handler, context) {
log::error!("{:?}", e);
}
drop(_waker);
worker.stop_all();
})?;
Ok(())
}
pub fn main_udp_listen0<H>(mut poll: Poll, mut recv_handler: H, context: Context) -> io::Result<()>
where
H: RecvChannelHandler,
{
let mut buf = [0; BUFFER_SIZE];
let mut udps = Vec::with_capacity(context.main_udp_socket.len());
for (index, udp) in context.main_udp_socket.iter().enumerate() {
let udp_socket = udp.try_clone()?;
udp_socket.set_nonblocking(true)?;
let mut mio_udp = UdpSocket::from_std(udp_socket);
poll.registry()
.register(&mut mio_udp, Token(index + 1), Interest::READABLE)?;
udps.push(mio_udp);
}
let mut events = Events::with_capacity(udps.len());
loop {
poll.poll(&mut events, None)?;
for x in events.iter() {
let index = match x.token() {
NOTIFY => return Ok(()),
Token(index) => index - 1,
};
let udp = if let Some(udp) = udps.get(index) {
udp
} else {
log::error!("{:?}", x);
continue;
};
loop {
match udp.recv_from(&mut buf) {
Ok((len, addr)) => {
recv_handler.handle(
&mut buf[..len],
RouteKey::new(false, index, addr),
&context,
);
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::error!("main_udp_listen_{}={:?}", index, e);
}
}
}
}
}
}
// /// 用recvmmsg没什么帮助,这里记录下,以下是完整代码
// #[cfg(unix)]
// pub fn main_udp_listen0<H>(index: usize, mut recv_handler: H, context: Context) -> io::Result<()>
// where
// H: RecvChannelHandler,
// {
// use libc::{c_uint, mmsghdr, sockaddr_storage, socklen_t, timespec};
// use std::os::fd::AsRawFd;
//
// let udp_socket = context.main_udp_socket[index].try_clone()?;
// let fd = udp_socket.as_raw_fd();
// const MAX_MESSAGES: usize = 16;
// let mut iov: [libc::iovec; MAX_MESSAGES] = unsafe { std::mem::zeroed() };
// let mut buf: [[u8; BUFFER_SIZE]; MAX_MESSAGES] = [[0; BUFFER_SIZE]; MAX_MESSAGES];
// let mut msgs: [mmsghdr; MAX_MESSAGES] = unsafe { std::mem::zeroed() };
// let mut addrs: [sockaddr_storage; MAX_MESSAGES] = unsafe { std::mem::zeroed() };
// for i in 0..MAX_MESSAGES {
// iov[i].iov_base = buf[i].as_mut_ptr() as *mut libc::c_void;
// iov[i].iov_len = BUFFER_SIZE;
// msgs[i].msg_hdr.msg_iov = &mut iov[i];
// msgs[i].msg_hdr.msg_iovlen = 1;
// msgs[i].msg_hdr.msg_name = &mut addrs[i] as *const _ as *mut libc::c_void;
// msgs[i].msg_hdr.msg_namelen = std::mem::size_of::<sockaddr_storage>() as socklen_t;
// }
// let mut time: timespec = unsafe { std::mem::zeroed() };
// loop {
// if context.is_stop() {
// return Ok(());
// }
// let res =
// unsafe { libc::recvmmsg(fd, msgs.as_mut_ptr(), MAX_MESSAGES as c_uint, 0, &mut time) };
// if res == -1 {
// log::error!("main_udp_listen_{}={:?}", index, io::Error::last_os_error());
// continue;
// }
//
// let nmsgs = res as usize;
// for i in 0..nmsgs {
// let msg = &mut buf[i][0..msgs[i].msg_len as usize];
// let addr = sockaddr_to_socket_addr(&addrs[i], msgs[i].msg_hdr.msg_namelen);
// if msg == b"stop" {
// if context.is_stop() {
// return Ok(());
// }
// }
// recv_handler.handle(msg, RouteKey::new(false, index, addr), &context);
// }
// }
// }
//
// #[cfg(unix)]
// fn sockaddr_to_socket_addr(addr: &libc::sockaddr_storage, _len: libc::socklen_t) -> SocketAddr {
// match addr.ss_family as libc::c_int {
// libc::AF_INET => {
// let addr_in = unsafe { *(addr as *const _ as *const libc::sockaddr_in) };
// let ip = u32::from_be(addr_in.sin_addr.s_addr);
// let port = u16::from_be(addr_in.sin_port);
// SocketAddr::V4(std::net::SocketAddrV4::new(Ipv4Addr::from(ip), port))
// }
// libc::AF_INET6 => {
// let addr_in6 = unsafe { *(addr as *const _ as *const libc::sockaddr_in6) };
// let ip = std::net::Ipv6Addr::from(addr_in6.sin6_addr.s6_addr);
// let port = u16::from_be(addr_in6.sin6_port);
// SocketAddr::V6(std::net::SocketAddrV6::new(ip, port, 0, 0))
// }
// _ => panic!("Unsupported address family"),
// }
// }
+8 -6
View File
@@ -51,10 +51,6 @@ impl AesEcbCipher {
//未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
}
if net_packet.payload().len() < 16 {
log::error!("数据异常,长度{}小于{}", net_packet.payload().len(), 16);
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if let Some(finger) = &self.finger {
let mut nonce_raw = [0; 12];
@@ -75,6 +71,10 @@ impl AesEcbCipher {
}
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"));
}
let mut out = [0u8; 1024 * 5];
let rs = match self.key {
AesEcbEnum::AES128ECB(key) => Aes128EcbDec::new(&key.into())
@@ -111,7 +111,7 @@ impl AesEcbCipher {
}
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("解密失败:{}", e),
format!("aes_ecb解密失败:{}", e),
)),
}
}
@@ -154,7 +154,7 @@ impl AesEcbCipher {
}
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("加密失败:{}", e),
format!("aes_ecb加密失败:{}", e),
)),
};
}
@@ -164,6 +164,8 @@ impl AesEcbCipher {
fn test_aes_ecb() {
let d = AesEcbCipher::new_128([0; 16], Some(Finger::new("123")));
let mut p = NetPacket::new_encrypt([0; 100]).unwrap();
let src = p.buffer().to_vec();
d.encrypt_ipv4(&mut p).unwrap();
d.decrypt_ipv4(&mut p).unwrap();
assert_eq!(p.buffer(), &src)
}
+202 -14
View File
@@ -1,47 +1,136 @@
#[cfg(feature = "aes_ecb")]
#[cfg(not(any(feature = "openssl-vendored", feature = "openssl")))]
use crate::cipher::aes_ecb::AesEcbCipher;
#[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;
#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))]
#[cfg(feature = "ring-cipher")]
use crate::cipher::ring_aes_gcm_cipher::AesGcmCipher;
use crate::cipher::{aes_cbc, Finger};
#[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::Finger;
use crate::protocol::NetPacket;
use aes_cbc::AesCbcCipher;
#[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 = "aes_cbc")]
AesCbc,
#[cfg(feature = "aes_ecb")]
AesEcb,
#[cfg(feature = "sm4_cbc")]
Sm4Cbc,
None,
}
impl FromStr for CipherModel {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
#[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 = "aes_cbc")]
"aes_cbc" => Ok(CipherModel::AesCbc),
#[cfg(feature = "aes_ecb")]
"aes_ecb" => Ok(CipherModel::AesEcb),
_ => Err(format!("not match '{}'", s)),
#[cfg(feature = "sm4_cbc")]
"sm4_cbc" => Ok(CipherModel::Sm4Cbc),
_ => {
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");
let str = if enums.is_empty() {
"no encrypt"
} else {
&enums[1..]
};
Err(format!("not match '{}', enum:{}", s, str))
}
}
}
}
#[derive(Clone)]
pub enum Cipher {
#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))]
AesGcm((AesGcmCipher, Vec<u8>)),
#[cfg(feature = "aes_cbc")]
AesCbc(AesCbcCipher),
#[cfg(feature = "aes_ecb")]
AesEcb(AesEcbCipher),
#[cfg(feature = "sm4_cbc")]
Sm4Cbc(Sm4CbcCipher),
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<String>,
_token: Option<String>,
) -> 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<String>,
@@ -53,6 +142,7 @@ impl Cipher {
hasher.update(password.as_bytes());
let key: [u8; 32] = hasher.finalize().into();
match model {
#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))]
CipherModel::AesGcm => {
if password.len() < 8 {
let aes = AesGcmCipher::new_128(key[..16].try_into().unwrap(), finger);
@@ -62,6 +152,7 @@ impl Cipher {
Cipher::AesGcm((aes, key.to_vec()))
}
}
#[cfg(feature = "aes_cbc")]
CipherModel::AesCbc => {
if password.len() < 8 {
let aes = AesCbcCipher::new_128(key[..16].try_into().unwrap(), finger);
@@ -71,6 +162,7 @@ impl Cipher {
Cipher::AesCbc(aes)
}
}
#[cfg(feature = "aes_ecb")]
CipherModel::AesEcb => {
if password.len() < 8 {
let aes = AesEcbCipher::new_128(key[..16].try_into().unwrap(), finger);
@@ -80,18 +172,43 @@ impl Cipher {
Cipher::AesEcb(aes)
}
}
#[cfg(feature = "sm4_cbc")]
CipherModel::Sm4Cbc => {
let aes = Sm4CbcCipher::new_128(key[..16].try_into().unwrap(), finger);
Cipher::Sm4Cbc(aes)
}
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<Self> {
Err(io::Error::new(io::ErrorKind::Other, "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<Self> {
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())))
@@ -104,9 +221,14 @@ impl Cipher {
net_packet: &mut NetPacket<B>,
) -> io::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 = "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::None => {
if net_packet.is_encrypt() {
return Err(io::Error::new(io::ErrorKind::Other, "not key"));
@@ -115,35 +237,101 @@ impl Cipher {
}
}
}
#[cfg(not(any(
feature = "aes_gcm",
feature = "server_encrypt",
feature = "aes_cbc",
feature = "aes_ecb",
feature = "sm4_cbc"
)))]
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
_net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
Ok(())
}
#[cfg(any(
feature = "aes_gcm",
feature = "server_encrypt",
feature = "aes_cbc",
feature = "aes_ecb",
feature = "sm4_cbc"
))]
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
match self {
#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))]
Cipher::AesGcm((aes_gcm, _)) => aes_gcm.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::None => Ok(()),
}
}
#[cfg(not(any(
feature = "aes_gcm",
feature = "server_encrypt",
feature = "aes_cbc",
feature = "aes_ecb",
feature = "sm4_cbc"
)))]
pub fn check_finger<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
_net_packet: &NetPacket<B>,
) -> io::Result<()> {
Ok(())
}
#[cfg(any(
feature = "aes_gcm",
feature = "server_encrypt",
feature = "aes_cbc",
feature = "aes_ecb",
feature = "sm4_cbc"
))]
pub fn check_finger<B: AsRef<[u8]>>(&self, net_packet: &NetPacket<B>) -> io::Result<()> {
let finger = match self {
Cipher::AesGcm((aes_gcm, _)) => aes_gcm.finger.as_ref(),
Cipher::AesCbc(aes_cbc) => aes_cbc.finger.as_ref(),
Cipher::AesEcb(aes_ecb) => aes_ecb.finger.as_ref(),
Cipher::None => None,
};
if let Some(finger) = finger {
finger.check_finger(net_packet)
} else {
Ok(())
match self {
#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))]
Cipher::AesGcm((aes_gcm, _)) => aes_gcm
.finger
.as_ref()
.map(|f| f.check_finger(net_packet))
.unwrap_or(Ok(())),
#[cfg(feature = "aes_cbc")]
Cipher::AesCbc(aes_cbc) => aes_cbc
.finger
.as_ref()
.map(|f| f.check_finger(net_packet))
.unwrap_or(Ok(())),
#[cfg(feature = "aes_ecb")]
Cipher::AesEcb(aes_ecb) => aes_ecb
.finger
.as_ref()
.map(|f| f.check_finger(net_packet))
.unwrap_or(Ok(())),
#[cfg(feature = "sm4_cbc")]
Cipher::Sm4Cbc(sm4_cbc) => sm4_cbc
.finger
.as_ref()
.map(|f| f.check_finger(net_packet))
.unwrap_or(Ok(())),
Cipher::None => Ok(()),
}
}
pub fn key(&self) -> Option<&[u8]> {
match self {
#[cfg(any(feature = "aes_gcm", feature = "server_encrypt"))]
Cipher::AesGcm((_, key)) => Some(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::None => None,
}
}
+26 -4
View File
@@ -1,18 +1,40 @@
#[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"
))]
mod finger;
#[cfg(feature = "ring-cipher")]
mod ring_aes_gcm_cipher;
mod rsa_cipher;
#[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"
))]
pub use finger::Finger;
#[cfg(feature = "server_encrypt")]
mod rsa_cipher;
#[cfg(feature = "server_encrypt")]
pub use rsa_cipher::RsaCipher;
+4 -4
View File
@@ -105,10 +105,6 @@ impl AesEcbCipher {
//未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
}
if net_packet.payload().len() < 16 {
log::error!("数据异常,长度{}小于{}", net_packet.payload().len(), 16);
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if let Some(finger) = &self.finger {
let mut nonce_raw = [0; 12];
@@ -129,6 +125,10 @@ impl AesEcbCipher {
}
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"));
}
let input = net_packet.payload();
let mut out = [0u8; 1024 * 5];
let mut out_len = 0;
+14 -9
View File
@@ -1,17 +1,18 @@
use crate::protocol::body::{RsaSecretBody, RSA_ENCRYPTION_RESERVED};
use crate::protocol::NetPacket;
use rand::Rng;
use rsa::pkcs8::der::Decode;
use rsa::{PublicKey, RsaPublicKey};
use sha2::Digest;
use spki::{DecodePublicKey, EncodePublicKey};
use std::io;
use {
crate::protocol::body::{RsaSecretBody, RSA_ENCRYPTION_RESERVED},
rand::Rng,
rsa::pkcs8::der::Decode,
rsa::RsaPublicKey,
sha2::Digest,
spki::{DecodePublicKey, EncodePublicKey},
};
#[derive(Clone)]
pub struct RsaCipher {
inner: Inner,
}
#[derive(Clone)]
struct Inner {
public_key: RsaPublicKey,
@@ -30,9 +31,10 @@ impl RsaCipher {
)),
}
}
pub fn finger(&self) -> io::Result<String> {
match self.inner.public_key.to_public_key_der() {
Ok(der) => match rsa::pkcs8::SubjectPublicKeyInfo::from_der(der.as_bytes()) {
Ok(der) => match rsa::pkcs8::SubjectPublicKeyInfoRef::from_der(der.as_bytes()) {
Ok(spki) => match spki.fingerprint_base64() {
Ok(finger) => Ok(finger),
Err(e) => Err(io::Error::new(
@@ -51,6 +53,9 @@ impl RsaCipher {
)),
}
}
pub fn public_key(&self) -> io::Result<&RsaPublicKey> {
return Ok(&self.inner.public_key);
}
}
impl RsaCipher {
@@ -83,7 +88,7 @@ impl RsaCipher {
secret_body.set_finger(&key[16..])?;
match self.inner.public_key.encrypt(
&mut rng,
rsa::PaddingScheme::PKCS1v15Encrypt,
rsa::pkcs1v15::Pkcs1v15Encrypt,
secret_body.buffer(),
) {
Ok(enc_data) => {
+171
View File
@@ -0,0 +1,171 @@
use crate::cipher::Finger;
use crate::protocol::{NetPacket, HEAD_LEN};
use libsm::sm4::cipher_mode::CipherMode;
use libsm::sm4::Sm4CipherMode;
use rand::RngCore;
use std::io;
pub struct Sm4CbcCipher {
key: [u8; 16],
pub(crate) cipher: Sm4CipherMode,
pub(crate) finger: Option<Finger>,
}
impl Clone for Sm4CbcCipher {
fn clone(&self) -> Self {
let cipher = Sm4CipherMode::new(&self.key, CipherMode::Cbc).unwrap();
Self {
key: self.key,
cipher,
finger: self.finger.clone(),
}
}
}
impl Sm4CbcCipher {
pub fn key(&self) -> &[u8] {
&self.key
}
}
impl Sm4CbcCipher {
pub fn new_128(key: [u8; 16], finger: Option<Finger>) -> Self {
let cipher = Sm4CipherMode::new(&key, CipherMode::Cbc).unwrap();
Self {
key,
cipher,
finger,
}
}
pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
if !net_packet.is_encrypt() {
//未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
}
if let Some(finger) = &self.finger {
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 len = net_packet.payload().len();
if len < 12 {
return Err(io::Error::new(io::ErrorKind::Other, "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"));
}
net_packet.set_data_len(net_packet.data_len() - finger.len())?;
}
let payload = net_packet.payload();
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"));
}
let mut out = [0u8; 1024 * 4];
let data = &payload[..len - 16];
let iv = &payload[len - 16..];
match self.cipher.decrypt(data, iv, &mut out) {
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"));
}
if src_net_packet.destination() != net_packet.destination() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.protocol() != net_packet.protocol() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.transport_protocol() != net_packet.transport_protocol() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.is_gateway() != net_packet.is_gateway() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.source_ttl() != net_packet.source_ttl() {
return Err(io::Error::new(io::ErrorKind::Other, "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),
)),
}
}
/// net_packet 必须预留足够长度
/// data_len是有效载荷的长度
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
let mut out = [0u8; 1024 * 4];
let mut iv = [0u8; 16];
rand::thread_rng().fill_bytes(&mut iv);
if net_packet.data_len() > 1024 * 4 - 32 {
log::error!(
"数据异常,长度{}大于1024 * 4 - 32",
net_packet.buffer().len()
);
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
match self.cipher.encrypt(net_packet.buffer(), &iv, &mut out) {
Ok(len) => {
net_packet.set_data_len(HEAD_LEN + len + 16)?;
net_packet.payload_mut()[..len].copy_from_slice(&out[..len]);
net_packet.payload_mut()[len..].copy_from_slice(&iv);
if let Some(finger) = &self.finger {
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 finger = finger.calculate_finger(&nonce_raw, net_packet.payload());
let src_data_len = net_packet.data_len();
//设置实际长度
net_packet.set_data_len(src_data_len + finger.len())?;
net_packet.buffer_mut()[src_data_len..].copy_from_slice(&finger);
}
net_packet.set_encrypt_flag(true);
Ok(())
}
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("sm4_cbc加密失败:{}", e),
)),
}
}
}
#[test]
fn test_sm4_ecb() {
let d = Sm4CbcCipher::new_128([0; 16], Some(Finger::new("123")));
let mut p = NetPacket::new_encrypt([1; 1024]).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 = Sm4CbcCipher::new_128([0; 16], None);
let mut p = NetPacket::new_encrypt([1; 102]).unwrap();
let src = p.buffer().to_vec();
d.encrypt_ipv4(&mut p).unwrap();
d.decrypt_ipv4(&mut p).unwrap();
assert_eq!(p.buffer(), &src)
}
+380
View File
@@ -0,0 +1,380 @@
use std::collections::HashMap;
use std::io;
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::time::Duration;
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::{Mutex, RwLock};
use rand::Rng;
use tun::device::IFace;
use crate::channel::context::Context;
use crate::channel::idle::Idle;
use crate::channel::punch::{NatInfo, Punch};
use crate::channel::{init_channel, init_context, Route, RouteKey};
use crate::cipher::Cipher;
#[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher;
use crate::core::Config;
use crate::external_route::{AllowExternalRoute, ExternalRoute};
use crate::handle::handshaker::Handshake;
use crate::handle::maintain::PunchReceiver;
use crate::handle::recv_data::RecvDataHandler;
use crate::handle::{
maintain, tun_tap, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo,
};
use crate::nat::NatTest;
use crate::util::{
Scheduler, SingleU64Adder, StopManager, U64Adder, WatchSingleU64Adder, WatchU64Adder,
};
use crate::{nat, tun_tap_device, DeviceInfo, VntCallback};
#[derive(Clone)]
pub struct Vnt {
stop_manager: StopManager,
config: Config,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
nat_test: NatTest,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
context: Context,
peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>>,
down_count_watcher: WatchU64Adder,
up_count_watcher: WatchSingleU64Adder,
}
impl Vnt {
pub fn new<Call: VntCallback>(config: Config, callback: Call) -> io::Result<Self> {
log::info!("config:{:?}", config);
//服务端非对称加密
#[cfg(feature = "server_encrypt")]
let rsa_cipher: Arc<Mutex<Option<RsaCipher>>> = Arc::new(Mutex::new(None));
//服务端对称加密
let server_cipher: Cipher = if config.server_encrypt {
let mut key = [0u8; 32];
rand::thread_rng().fill(&mut key);
Cipher::new_key(key, config.token.clone())?
} else {
Cipher::None
};
let finger = if config.finger {
Some(config.token.clone())
} else {
None
};
//客户端对称加密
let client_cipher =
Cipher::new_password(config.cipher_model, config.password.clone(), finger);
//当前设备信息
let current_device = Arc::new(AtomicCell::new(CurrentDeviceInfo::new0(
config.server_address,
)));
//设备列表
let device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>> =
Arc::new(Mutex::new((0, Vec::with_capacity(16))));
//基础信息
let config_info = BaseConfigInfo::new(
config.name.clone(),
config.token.clone(),
config.ip,
config.password.is_some(),
config.device_id.clone(),
config.server_address_str.clone(),
);
let ports = config.ports.as_ref().map_or(vec![0, 0], |v| {
if v.is_empty() {
vec![0, 0]
} else {
v.clone()
}
});
//通道上下文
let (context, tcp_listener) = init_context(
ports,
config.use_channel_type,
config.first_latency,
config.tcp,
config.packet_loss_rate,
config.packet_delay,
)?;
let local_ipv4 = nat::local_ipv4();
let local_ipv6 = nat::local_ipv6();
let udp_ports = context.main_local_udp_port()?;
let tcp_port = tcp_listener.local_addr()?.port();
//nat检测工具
let nat_test = NatTest::new(
context.channel_num(),
config.stun_server.clone(),
local_ipv4,
local_ipv6,
udp_ports,
tcp_port,
);
// 虚拟网卡
let device = tun_tap_device::create_device(&config)?;
let tun_info = DeviceInfo::new(device.name()?, device.version()?);
callback.create_tun(tun_info);
// 服务停止管理器
let stop_manager = {
let callback = callback.clone();
StopManager::new(move || callback.stop())
};
// 定时器
let scheduler = Scheduler::new(stop_manager.clone())?;
let external_route = ExternalRoute::new(config.in_ips.clone());
let out_external_route = AllowExternalRoute::new(config.out_ips.clone());
#[cfg(feature = "ip_proxy")]
let proxy_map = if !config.out_ips.is_empty() && !config.no_proxy {
Some(crate::ip_proxy::init_proxy(
context.clone(),
scheduler.clone(),
stop_manager.clone(),
current_device.clone(),
client_cipher.clone(),
)?)
} else {
None
};
let (punch_sender, punch_receiver) = maintain::punch_channel();
let peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>> =
Arc::new(RwLock::new(HashMap::with_capacity(16)));
let down_counter =
U64Adder::with_capacity(config.ports.as_ref().map(|v| v.len()).unwrap_or_default() + 8);
let down_count_watcher = down_counter.watch();
let handshake = Handshake::new();
let handler = RecvDataHandler::new(
#[cfg(feature = "server_encrypt")]
rsa_cipher,
server_cipher.clone(),
client_cipher.clone(),
current_device.clone(),
device.clone(),
device_list.clone(),
config_info.clone(),
nat_test.clone(),
callback.clone(),
punch_sender,
peer_nat_info_map.clone(),
external_route.clone(),
out_external_route,
#[cfg(feature = "ip_proxy")]
proxy_map.clone(),
down_counter,
handshake.clone(),
);
//初始化网络数据通道
let (udp_socket_sender, tcp_socket_sender) =
init_channel(tcp_listener, context.clone(), stop_manager.clone(), handler)?;
// 打洞逻辑
let punch = Punch::new(
context.clone(),
config.punch_model,
config.tcp,
tcp_socket_sender.clone(),
);
let up_counter = SingleU64Adder::new();
let up_count_watcher = up_counter.watch();
tun_tap::tun_handler::start(
stop_manager.clone(),
context.clone(),
device.clone(),
current_device.clone(),
external_route,
#[cfg(feature = "ip_proxy")]
proxy_map,
client_cipher.clone(),
server_cipher.clone(),
config.parallel,
up_counter,
)?;
maintain::idle_gateway(
&scheduler,
context.clone(),
current_device.clone(),
config_info.clone(),
tcp_socket_sender.clone(),
callback.clone(),
0,
handshake,
);
{
let context = context.clone();
let nat_test = nat_test.clone();
let device_list = device_list.clone();
let down_count_watcher = down_count_watcher.clone();
let up_count_watcher = up_count_watcher.clone();
let current_device = current_device.clone();
if !config.use_channel_type.is_only_relay() {
// 定时nat探测
maintain::retrieve_nat_type(
&scheduler,
context.clone(),
nat_test.clone(),
udp_socket_sender,
);
}
//延迟启动
scheduler.timeout(Duration::from_secs(3), move |scheduler| {
start(
scheduler,
context,
nat_test,
device_list,
current_device,
client_cipher,
server_cipher,
punch_receiver,
config_info,
punch,
callback,
down_count_watcher,
up_count_watcher,
);
});
}
Ok(Self {
stop_manager,
config,
current_device,
nat_test,
device_list,
context,
peer_nat_info_map,
down_count_watcher,
up_count_watcher,
})
}
}
pub fn start<Call: VntCallback>(
scheduler: &Scheduler,
context: Context,
nat_test: NatTest,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
server_cipher: Cipher,
punch_receiver: PunchReceiver,
config_info: BaseConfigInfo,
punch: Punch,
callback: Call,
down_count_watcher: WatchU64Adder,
up_count_watcher: WatchSingleU64Adder,
) {
// 定时心跳
maintain::heartbeat(
&scheduler,
context.clone(),
current_device.clone(),
device_list.clone(),
client_cipher.clone(),
server_cipher.clone(),
);
// 路由空闲检测逻辑
let idle = Idle::new(Duration::from_secs(10), context.clone());
// 定时空闲检查
maintain::idle_route(
&scheduler,
idle,
context.clone(),
current_device.clone(),
callback,
);
// 定时客户端中继检测
if !context.use_channel_type().is_only_p2p() {
maintain::client_relay(
&scheduler,
context.clone(),
current_device.clone(),
device_list.clone(),
client_cipher.clone(),
);
}
// 定时地址探测
maintain::addr_request(
&scheduler,
context.clone(),
current_device.clone(),
server_cipher.clone(),
config_info.clone(),
);
if !context.use_channel_type().is_only_relay() {
// 定时打洞
maintain::punch(
&scheduler,
context.clone(),
nat_test.clone(),
device_list.clone(),
current_device.clone(),
client_cipher.clone(),
punch_receiver,
punch,
);
}
maintain::up_status(
scheduler,
context.clone(),
current_device.clone(),
down_count_watcher,
up_count_watcher,
)
}
impl Vnt {
pub fn name(&self) -> &str {
&self.config.name
}
pub fn server_encrypt(&self) -> bool {
self.config.server_encrypt
}
pub fn client_encrypt(&self) -> bool {
self.config.password.is_some()
}
pub fn current_device(&self) -> CurrentDeviceInfo {
self.current_device.load()
}
pub fn peer_nat_info(&self, ip: &Ipv4Addr) -> Option<NatInfo> {
self.peer_nat_info_map.read().get(ip).cloned()
}
pub fn connection_status(&self) -> ConnectStatus {
self.current_device.load().status
}
pub fn nat_info(&self) -> NatInfo {
self.nat_test.nat_info()
}
pub fn device_list(&self) -> Vec<PeerDeviceInfo> {
let device_list_lock = self.device_list.lock();
let (_epoch, device_list) = device_list_lock.clone();
drop(device_list_lock);
device_list
}
pub fn route(&self, ip: &Ipv4Addr) -> Option<Route> {
self.context.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<Ipv4Addr> {
self.context.route_table.route_to_id(route_key)
}
pub fn route_table(&self) -> Vec<(Ipv4Addr, Vec<Route>)> {
self.context.route_table.route_table()
}
pub fn up_stream(&self) -> u64 {
self.up_count_watcher.get()
}
pub fn down_stream(&self) -> u64 {
self.down_count_watcher.get()
}
pub fn stop(&self) {
self.stop_manager.stop()
}
pub fn wait(&self) {
self.stop_manager.wait()
}
}
+54 -556
View File
@@ -1,554 +1,17 @@
use std::io;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::Arc;
use std::time::Duration;
use std::net::{Ipv4Addr, SocketAddr};
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use parking_lot::Mutex;
use rand::Rng;
use std::net::UdpSocket;
use tokio::net::TcpStream;
use tokio::sync::mpsc::channel;
pub use conn::Vnt;
use crate::channel::channel::{Channel, Context};
use crate::channel::idle::Idle;
use crate::channel::punch::{NatInfo, Punch, PunchModel};
use crate::channel::sender::ChannelSender;
use crate::channel::{Route, RouteKey};
use crate::cipher::{Cipher, CipherModel, RsaCipher};
use crate::core::status::VntStatusManger;
use crate::error::Error;
use crate::external_route::{AllowExternalRoute, ExternalRoute};
use crate::handle::handshake_handler::HandshakeEnum;
use crate::handle::recv_handler::ChannelDataHandler;
use crate::handle::registration_handler::{RegResponse, ReqEnum};
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
use crate::handle::tun_tap::tap_handler;
use crate::handle::tun_tap::tun_handler;
use crate::handle::{
handshake_handler, heartbeat_handler, punch_handler, registration_handler, ConnectStatus,
CurrentDeviceInfo, PeerDeviceInfo,
};
use crate::igmp_server::IgmpServer;
use crate::ip_proxy::DashMapNew;
use crate::nat::NatTest;
use crate::tun_tap_device;
use crate::tun_tap_device::{DeviceReader, DeviceWriter};
use crate::channel::punch::PunchModel;
use crate::channel::UseChannelType;
use crate::cipher::CipherModel;
pub mod status;
pub mod sync;
#[derive(Clone)]
pub struct Vnt {
config: Config,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
context: Context,
vnt_status_manager: VntStatusManger,
device_writer: DeviceWriter,
/// 0. 机器纪元,每一次上线或者下线都会增1,用于感知网络中机器变化
/// 服务端和客户端的不一致,则服务端会推送新的设备列表
/// 1. 网络中的虚拟ip列表
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
nat_test: NatTest,
connect_status: Arc<AtomicCell<ConnectStatus>>,
peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>>,
}
pub struct VntUtil {
config: Config,
main_channel: UdpSocket,
main_channel_ipv6: Option<UdpSocket>,
main_tcp_channel: Option<TcpStream>,
response: Option<RegResponse>,
iface: Option<(DeviceWriter, DeviceReader)>,
server_cipher: Cipher,
rsa_cipher: Option<RsaCipher>,
}
impl VntUtil {
pub async fn new(config: Config) -> io::Result<VntUtil> {
//单个udp用同步的性能更好,但是代理和多端口监听用异步更方便,这里将两者结合起来
let main_channel = UdpSocket::bind("0.0.0.0:0")?;
main_channel.set_write_timeout(Some(Duration::from_secs(5)))?;
main_channel.set_read_timeout(Some(Duration::from_secs(2)))?;
let main_channel_ipv6 = if config.punch_model != PunchModel::IPv4 {
match UdpSocket::bind("[::]:0") {
Ok(main_channel_ipv6) => {
main_channel_ipv6.set_write_timeout(Some(Duration::from_secs(5)))?;
Some(main_channel_ipv6)
}
Err(e) => {
log::warn!("绑定ipv6地址失败:{}", e);
None
}
}
} else {
None
};
let server_cipher = if config.server_encrypt {
let mut key = [0 as u8; 32];
rand::thread_rng().fill(&mut key);
Cipher::new_key(key, config.token.clone())?
} else {
Cipher::None
};
Ok(VntUtil {
config,
main_channel,
main_channel_ipv6,
main_tcp_channel: None,
response: None,
iface: None,
server_cipher,
rsa_cipher: None,
})
}
///链接
pub async fn connect(&mut self) -> io::Result<()> {
if self.config.tcp {
let tcp = TcpStream::connect(self.config.server_address).await?;
let _ = self.main_tcp_channel.insert(tcp);
}
Ok(())
}
///握手 用于获取公钥
pub async fn handshake(&mut self) -> Result<Option<RsaCipher>, HandshakeEnum> {
let rsa_cipher = handshake_handler::handshake(
&self.main_channel,
self.main_tcp_channel.as_mut(),
self.config.server_address,
self.config.server_encrypt,
)
.await?;
self.rsa_cipher = rsa_cipher.clone();
Ok(rsa_cipher)
}
/// 加密握手 用于同步密钥
pub async fn secret_handshake(&mut self) -> Result<(), HandshakeEnum> {
handshake_handler::secret_handshake(
&self.main_channel,
self.main_tcp_channel.as_mut(),
self.config.server_address,
self.rsa_cipher.as_ref().unwrap(),
&self.server_cipher,
self.config.token.clone(),
)
.await
}
/// 注册
pub async fn register(&mut self) -> Result<RegResponse, ReqEnum> {
match registration_handler::registration(
&self.main_channel,
self.main_tcp_channel.as_mut(),
&self.server_cipher,
self.config.server_address,
self.config.token.clone(),
self.config.device_id.clone(),
self.config.name.clone(),
self.config.ip.unwrap_or(Ipv4Addr::UNSPECIFIED),
self.config.password.is_some(),
)
.await
{
Ok(res) => {
let _ = self.response.insert(res.clone());
Ok(res)
}
Err(e) => Err(e),
}
}
#[cfg(any(target_os = "android"))]
pub fn create_iface(&mut self, vpn_fd: i32) {
let (device_writer, device_reader) = tun_tap_device::create(vpn_fd);
let _ = self.iface.insert((device_writer, device_reader));
}
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
pub fn create_iface(&mut self) -> io::Result<tun_tap_device::DriverInfo> {
if self.iface.is_some() {
return Err(io::Error::from(io::ErrorKind::AlreadyExists));
}
let response = match &self.response {
None => {
return Err(io::Error::from(io::ErrorKind::AlreadyExists));
}
Some(res) => res,
};
let device_type = if self.config.tap {
{
//删除tun网卡避免ip冲突,因为非正常退出会保留网卡
tun_tap_device::delete_device(tun_tap_device::DeviceType::Tun);
}
tun_tap_device::DeviceType::Tap
} else {
{
//删除tap网卡避免ip冲突,非正常退出会保留网卡
tun_tap_device::delete_device(tun_tap_device::DeviceType::Tap);
}
tun_tap_device::DeviceType::Tun
};
let mtu = match self.config.mtu {
None => {
if self.config.password.is_none() {
1450
} else {
1410
}
}
Some(mtu) => mtu,
};
let in_ips = self
.config
.in_ips
.iter()
.map(|(dest, mask, _)| (Ipv4Addr::from(*dest & *mask), Ipv4Addr::from(*mask)))
.collect::<Vec<(Ipv4Addr, Ipv4Addr)>>();
let (device_writer, device_reader, driver_info) = tun_tap_device::create_device(
device_type,
response.virtual_ip,
response.virtual_netmask,
response.virtual_gateway,
in_ips,
mtu,
)?;
let _ = self.iface.insert((device_writer, device_reader));
Ok(driver_info)
}
pub async fn build(self) -> crate::Result<Vnt> {
//将读的超时时间清空
self.main_channel.set_read_timeout(None)?;
let response = match self.response {
None => {
return Err(Error::Stop("response None".to_string()));
}
Some(res) => res,
};
let (device_writer, device_reader) = match self.iface {
None => {
return Err(Error::Stop("iface None".to_string()));
}
Some(res) => res,
};
let config = self.config.clone();
let vnt_status_manager = VntStatusManger::new();
let finger = if config.finger {
Some(config.token.clone())
} else {
None
};
let client_cipher =
Cipher::new_password(config.cipher_model, config.password.clone(), finger);
let virtual_ip = response.virtual_ip;
let virtual_gateway = response.virtual_gateway;
let virtual_netmask = response.virtual_netmask;
let current_device = Arc::new(AtomicCell::new(CurrentDeviceInfo::new(
virtual_ip,
virtual_gateway,
virtual_netmask,
config.server_address,
)));
let (cone_sender, cone_receiver) = channel(3);
let (symmetric_sender, symmetric_receiver) = channel(2);
let (tcp_sender, tcp) = if let Some(main_tcp_channel) = self.main_tcp_channel {
let (tcp_sender, tcp_receiver) = channel::<Vec<u8>>(100);
(Some(tcp_sender), Some((main_tcp_channel, tcp_receiver)))
} else {
(None, None)
};
let context = Context::new(
Arc::new(self.main_channel),
self.main_channel_ipv6.map(|v| Arc::new(v)),
tcp_sender,
current_device.clone(),
1,
);
let punch = Punch::new(context.clone(), config.punch_model);
let idle = Idle::new(Duration::from_secs(16), context.clone());
let channel_sender = ChannelSender::new(context.clone());
let register = Arc::new(registration_handler::Register::new(
self.server_cipher.clone(),
channel_sender.clone(),
config.server_address,
config.token.clone(),
config.device_id.clone(),
config.name.clone(),
config.password.is_some(),
));
let device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>> =
Arc::new(Mutex::new((response.epoch, response.device_info_list)));
let peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>> = Arc::new(DashMap::new0());
let connect_status = Arc::new(AtomicCell::new(ConnectStatus::Connected));
let public_ip = response.public_ip;
let public_port = response.public_port;
let local_port = context.main_local_ipv4_port().unwrap_or(0);
let local_ipv4_addr = crate::nat::local_ipv4_addr(local_port);
let ipv6_port = context.main_local_ipv6_port().unwrap_or(0);
let ipv6_addr = crate::nat::local_ipv6_addr(ipv6_port);
// NAT检测
let nat_test = NatTest::new(
config.stun_server.clone(),
public_ip,
public_port,
local_ipv4_addr,
ipv6_addr,
);
let in_external_route = if config.in_ips.is_empty() {
None
} else {
Some(ExternalRoute::new(config.in_ips))
};
let (tcp_proxy, udp_proxy, ip_proxy_map) = if config.out_ips.is_empty() {
(None, None, None)
} else {
let (tcp_proxy, udp_proxy, ip_proxy_map) = crate::ip_proxy::init_proxy(
#[cfg(not(target_os = "android"))]
channel_sender.clone(),
#[cfg(not(target_os = "android"))]
current_device.clone(),
#[cfg(not(target_os = "android"))]
client_cipher.clone(),
)
.await?;
(Some(tcp_proxy), Some(udp_proxy), Some(ip_proxy_map))
};
let out_external_route = AllowExternalRoute::new(config.out_ips);
let igmp_server = if config.simulate_multicast {
Some(IgmpServer::new(device_writer.clone()))
} else {
None
};
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
if config.tap {
tap_handler::start(
vnt_status_manager.worker("tap_handler"),
channel_sender.clone(),
device_reader,
device_writer.clone(),
igmp_server.clone(),
current_device.clone(),
in_external_route,
ip_proxy_map.clone(),
client_cipher.clone(),
self.server_cipher.clone(),
config.parallel,
);
} else {
tun_handler::start(
vnt_status_manager.worker("tun_handler"),
channel_sender.clone(),
device_reader,
device_writer.clone(),
igmp_server.clone(),
current_device.clone(),
in_external_route,
ip_proxy_map.clone(),
client_cipher.clone(),
self.server_cipher.clone(),
config.parallel,
);
}
#[cfg(any(target_os = "android"))]
tun_handler::start(
vnt_status_manager.worker("android tun_handler"),
channel_sender.clone(),
device_reader,
device_writer.clone(),
igmp_server.clone(),
current_device.clone(),
in_external_route,
ip_proxy_map.clone(),
client_cipher.clone(),
self.server_cipher.clone(),
config.parallel,
);
//外部数据接收处理
let channel_recv_handler = ChannelDataHandler::new(
current_device.clone(),
device_list.clone(),
register.clone(),
nat_test.clone(),
igmp_server,
device_writer.clone(),
connect_status.clone(),
peer_nat_info_map.clone(),
ip_proxy_map,
out_external_route,
cone_sender,
symmetric_sender,
client_cipher.clone(),
self.server_cipher.clone(),
self.rsa_cipher.clone(),
config.relay,
config.token.clone(),
);
{
let channel = Channel::new(context.clone(), channel_recv_handler);
let channel_worker = vnt_status_manager.worker("channel_worker");
let relay = config.relay;
tokio::spawn(async move {
channel
.start(channel_worker, tcp, 14, 65, relay, config.parallel)
.await
});
}
{
let nat_test = nat_test.clone();
let device_list = device_list.clone();
let current_device = current_device.clone();
// 定时心跳
heartbeat_handler::start_heartbeat(
vnt_status_manager.worker("heartbeat"),
channel_sender.clone(),
device_list.clone(),
current_device.clone(),
config.server_address_str,
client_cipher.clone(),
self.server_cipher.clone(),
);
// 空闲检查
heartbeat_handler::start_idle(
vnt_status_manager.worker("idle"),
idle,
channel_sender.clone(),
);
if !config.relay {
// 打洞处理
punch_handler::start(
vnt_status_manager.worker("cone_receiver"),
cone_receiver,
punch.clone(),
current_device.clone(),
client_cipher.clone(),
);
punch_handler::start(
vnt_status_manager.worker("symmetric_receiver"),
symmetric_receiver,
punch,
current_device.clone(),
client_cipher.clone(),
);
tokio::spawn(punch_handler::start_punch(
vnt_status_manager.worker("punch_handler"),
nat_test,
device_list,
channel_sender,
current_device,
client_cipher.clone(),
));
}
}
{
//代理
if let Some(tcp_proxy) = tcp_proxy {
tokio::spawn(tcp_proxy.start());
}
if let Some(udp_proxy) = udp_proxy {
tokio::spawn(udp_proxy.start());
}
let context = context.clone();
let nat_test = nat_test.clone();
tokio::spawn(async move {
let info = nat_test
.re_test(public_ip, public_port, local_ipv4_addr, ipv6_addr)
.await;
context.switch(info.nat_type);
});
}
Ok(Vnt {
config: self.config,
current_device,
context,
vnt_status_manager,
device_writer,
nat_test,
device_list,
connect_status,
peer_nat_info_map,
})
}
}
impl Vnt {
pub fn name(&self) -> &str {
&self.config.name
}
pub fn server_encrypt(&self) -> bool {
self.config.server_encrypt
}
pub fn client_encrypt(&self) -> bool {
self.config.password.is_some()
}
pub fn current_device(&self) -> CurrentDeviceInfo {
self.current_device.load()
}
pub fn peer_nat_info(&self, ip: &Ipv4Addr) -> Option<NatInfo> {
self.peer_nat_info_map.get(ip).map(|e| e.value().clone())
}
pub fn connection_status(&self) -> ConnectStatus {
self.connect_status.load()
}
pub fn nat_info(&self) -> NatInfo {
self.nat_test.nat_info()
}
pub fn device_list(&self) -> Vec<PeerDeviceInfo> {
let device_list_lock = self.device_list.lock();
let (_epoch, device_list) = device_list_lock.clone();
drop(device_list_lock);
device_list
}
pub fn route(&self, ip: &Ipv4Addr) -> Option<Route> {
self.context.route_one(ip)
}
pub fn route_key(&self, route_key: &RouteKey) -> Option<Ipv4Addr> {
self.context.route_to_id(route_key)
}
pub fn route_table(&self) -> Vec<(Ipv4Addr, Route)> {
self.context.route_table_one()
}
pub fn stop(&self) -> io::Result<()> {
let _ = self.context.close();
self.vnt_status_manager.stop_all();
let _ = self.device_writer.close();
let virtual_gateway = self.current_device.load().virtual_gateway;
let _ = UdpSocket::bind("0.0.0.0:0")?.send_to(
b"stop",
SocketAddr::V4(SocketAddrV4::new(virtual_gateway, 10000)),
);
Ok(())
}
pub async fn wait_stop(&mut self) {
self.vnt_status_manager.wait().await;
let _ = self.stop();
}
pub async fn wait_stop_ms(&mut self, ms: Duration) -> bool {
tokio::select! {
_=self.vnt_status_manager.wait()=>{
let _ = self.stop();
return true;
}
_=tokio::time::sleep(ms)=>{
return false;
}
}
}
}
impl Drop for Vnt {
fn drop(&mut self) {
let _ = self.stop();
}
}
mod conn;
#[derive(Clone, Debug)]
pub struct Config {
#[cfg(any(target_os = "windows", target_os = "linux"))]
pub tap: bool,
pub token: String,
pub device_id: String,
@@ -559,21 +22,31 @@ pub struct Config {
pub in_ips: Vec<(u32, u32, Ipv4Addr)>,
pub out_ips: Vec<(u32, u32)>,
pub password: Option<String>,
pub simulate_multicast: bool,
pub mtu: Option<u16>,
pub mtu: Option<u32>,
pub tcp: bool,
pub ip: Option<Ipv4Addr>,
pub relay: bool,
#[cfg(feature = "ip_proxy")]
pub no_proxy: bool,
pub server_encrypt: bool,
pub parallel: usize,
pub cipher_model: CipherModel,
pub finger: bool,
pub punch_model: PunchModel,
pub ports: Option<Vec<u16>>,
pub first_latency: bool,
#[cfg(not(target_os = "android"))]
pub device_name: Option<String>,
#[cfg(target_os = "android")]
pub device_fd: i32,
pub use_channel_type: UseChannelType,
//控制丢包率
pub packet_loss_rate: Option<f64>,
pub packet_delay: u32,
}
impl Config {
pub fn new(
tap: bool,
#[cfg(any(target_os = "windows", target_os = "linux"))] tap: bool,
token: String,
device_id: String,
name: String,
@@ -583,23 +56,39 @@ impl Config {
in_ips: Vec<(u32, u32, Ipv4Addr)>,
out_ips: Vec<(u32, u32)>,
password: Option<String>,
simulate_multicast: bool,
mtu: Option<u16>,
mtu: Option<u32>,
tcp: bool,
ip: Option<Ipv4Addr>,
relay: bool,
#[cfg(feature = "ip_proxy")] no_proxy: bool,
server_encrypt: bool,
parallel: usize,
cipher_model: CipherModel,
finger: bool,
punch_model: PunchModel,
) -> Self {
ports: Option<Vec<u16>>,
first_latency: bool,
#[cfg(not(target_os = "android"))] device_name: Option<String>,
#[cfg(target_os = "android")] device_fd: i32,
use_channel_type: UseChannelType,
packet_loss_rate: Option<f64>,
packet_delay: u32,
) -> io::Result<Self> {
for x in stun_server.iter_mut() {
if !x.contains(":") {
x.push_str(":3478");
}
}
Self {
if token.is_empty() || token.len() > 128 {
return Err(io::Error::new(io::ErrorKind::Other, "token too long"));
}
if device_id.is_empty() || device_id.len() > 128 {
return Err(io::Error::new(io::ErrorKind::Other, "device_id too long"));
}
if name.is_empty() || name.len() > 128 {
return Err(io::Error::new(io::ErrorKind::Other, "name too long"));
}
Ok(Self {
#[cfg(any(target_os = "windows", target_os = "linux"))]
tap,
token,
device_id,
@@ -610,16 +99,25 @@ impl Config {
in_ips,
out_ips,
password,
simulate_multicast,
mtu,
tcp,
ip,
relay,
#[cfg(feature = "ip_proxy")]
no_proxy,
server_encrypt,
parallel,
cipher_model,
finger,
punch_model,
}
ports,
first_latency,
#[cfg(not(target_os = "android"))]
device_name,
#[cfg(target_os = "android")]
device_fd,
use_channel_type,
packet_loss_rate,
packet_delay,
})
}
}
-92
View File
@@ -1,92 +0,0 @@
use crate::util::wait::WaitGroup;
use std::sync::Arc;
use tokio::sync::watch;
use tokio::sync::watch::{Receiver, Sender};
#[derive(Copy, Clone, Eq, PartialEq)]
pub enum VntStatus {
Starting,
Stopping,
}
pub struct VntWorker {
name: String,
wg: WaitGroup,
status_s: Arc<Sender<VntStatus>>,
status_r: Receiver<VntStatus>,
}
impl VntWorker {
pub fn worker(&self, name: &str) -> Self {
self.wg.add();
VntWorker {
name: name.to_string(),
wg: self.wg.clone(),
status_s: self.status_s.clone(),
status_r: self.status_r.clone(),
}
}
}
impl Drop for VntWorker {
fn drop(&mut self) {
log::info!("任务停止:{}", self.name);
self.wg.done();
}
}
impl VntWorker {
pub fn stop_all(&self) {
let _ = self.status_s.send(VntStatus::Stopping);
}
pub async fn stop_wait(&mut self) {
loop {
if *self.status_r.borrow() == VntStatus::Stopping {
return;
}
match self.status_r.changed().await {
Ok(_) => {
if *self.status_r.borrow() == VntStatus::Stopping {
return;
}
}
Err(_) => {
return;
}
}
}
}
}
#[derive(Clone)]
pub struct VntStatusManger {
wg: WaitGroup,
status_s: Arc<Sender<VntStatus>>,
status_r: Receiver<VntStatus>,
}
impl VntStatusManger {
pub fn new() -> Self {
let (status_s, status_r) = watch::channel(VntStatus::Starting);
Self {
wg: WaitGroup::new(),
status_s: Arc::new(status_s),
status_r,
}
}
pub fn stop_all(&self) {
let _ = self.status_s.send(VntStatus::Stopping);
}
pub async fn wait(&mut self) {
self.wg.wait().await
}
pub fn worker(&self, name: &str) -> VntWorker {
self.wg.add();
VntWorker {
name: name.to_string(),
wg: self.wg.clone(),
status_s: self.status_s.clone(),
status_r: self.status_r.clone(),
}
}
}
-81
View File
@@ -1,81 +0,0 @@
use crate::cipher::RsaCipher;
use crate::core::{Config, Vnt, VntUtil};
use crate::handle::handshake_handler::HandshakeEnum;
use crate::handle::registration_handler::{RegResponse, ReqEnum};
use std::io;
use std::ops::Deref;
use std::time::Duration;
use tokio::runtime::Runtime;
pub struct VntUtilSync {
vnt_util: VntUtil,
runtime: Runtime,
}
pub struct VntSync {
vnt: Vnt,
runtime: Runtime,
}
impl VntUtilSync {
pub fn new(config: Config) -> io::Result<VntUtilSync> {
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()?;
let vnt_util = runtime.block_on(VntUtil::new(config))?;
Ok(VntUtilSync { vnt_util, runtime })
}
pub fn connect(&mut self) -> io::Result<()> {
self.runtime.block_on(self.vnt_util.connect())
}
pub fn handshake(&mut self) -> Result<Option<RsaCipher>, HandshakeEnum> {
self.runtime.block_on(self.vnt_util.handshake())
}
pub fn secret_handshake(&mut self) -> Result<(), HandshakeEnum> {
self.runtime.block_on(self.vnt_util.secret_handshake())
}
pub fn register(&mut self) -> Result<RegResponse, ReqEnum> {
self.runtime.block_on(self.vnt_util.register())
}
#[cfg(any(target_os = "android"))]
pub fn create_iface(&mut self, vpn_fd: i32) {
self.vnt_util.create_iface(vpn_fd)
}
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
pub fn create_iface(&mut self) -> io::Result<crate::tun_tap_device::DriverInfo> {
self.vnt_util.create_iface()
}
pub fn build(self) -> crate::Result<VntSync> {
let runtime = self.runtime;
let vnt = runtime.block_on(self.vnt_util.build())?;
{
let mut vnt = vnt.clone();
std::thread::spawn(move || runtime.block_on(vnt.wait_stop()));
}
Ok(VntSync {
vnt,
runtime: tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap(),
})
}
}
impl VntSync {
pub fn wait_stop(&mut self) {
self.runtime.block_on(self.vnt.wait_stop())
}
pub fn wait_stop_ms(&mut self, ms: u64) -> bool {
self.runtime
.block_on(self.vnt.wait_stop_ms(Duration::from_millis(ms)))
}
}
impl Deref for VntSync {
type Target = Vnt;
fn deref(&self) -> &Self::Target {
&self.vnt
}
}
-21
View File
@@ -1,21 +0,0 @@
use std::io;
use thiserror::Error;
#[derive(Error, Debug)]
pub enum Error {
#[error("Io error")]
Io(#[from] io::Error),
#[error("Protobuf error")]
Protobuf(#[from] protobuf::Error),
#[error("Invalid packet")]
InvalidPacket,
#[error("Not support")]
NotSupport,
#[error("Stop")]
Stop(String),
#[error("Warn")]
Warn(String),
}
pub type Result<T> = std::result::Result<T, Error>;
+14 -4
View File
@@ -5,16 +5,17 @@ use std::sync::Arc;
#[derive(Clone)]
pub struct ExternalRoute {
route_table: Arc<Vec<(u32, u32, Ipv4Addr)>>,
route_table: Vec<(u32, u32, Ipv4Addr)>,
}
impl ExternalRoute {
pub fn new(route_table: Vec<(u32, u32, Ipv4Addr)>) -> Self {
Self {
route_table: Arc::new(route_table),
}
Self { route_table }
}
pub fn route(&self, ip: &Ipv4Addr) -> Option<Ipv4Addr> {
if self.route_table.is_empty() {
return None;
}
let ip = u32::from_be_bytes(ip.octets());
for (dest, mask, gateway) in self.route_table.iter() {
if *mask & ip == *mask & *dest {
@@ -23,6 +24,12 @@ impl ExternalRoute {
}
None
}
pub fn to_route(&self) -> Vec<(Ipv4Addr, Ipv4Addr)> {
self.route_table
.iter()
.map(|(dest, mask, _)| (Ipv4Addr::from(*dest & *mask), Ipv4Addr::from(*mask)))
.collect::<Vec<(Ipv4Addr, Ipv4Addr)>>()
}
}
#[derive(Clone)]
@@ -37,6 +44,9 @@ impl AllowExternalRoute {
}
}
pub fn allow(&self, ip: &Ipv4Addr) -> bool {
if self.route_table.is_empty() {
return false;
}
let ip = u32::from_be_bytes(ip.octets());
for (dest, mask) in self.route_table.iter() {
if *mask & ip == *mask & *dest {
+206
View File
@@ -0,0 +1,206 @@
#[cfg(feature = "server_encrypt")]
use rsa::RsaPublicKey;
use std::fmt::{Display, Formatter};
use std::io;
use std::net::{Ipv4Addr, SocketAddr};
#[derive(Debug)]
pub struct DeviceInfo {
pub name: String,
pub version: String,
}
impl Display for DeviceInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!("name={} ,version={}", self.name, self.version))
}
}
impl DeviceInfo {
pub fn new(name: String, version: String) -> Self {
return Self { name, version };
}
}
#[derive(Debug)]
pub struct ConnectInfo {
// 第几次连接,从1开始
pub count: usize,
// 服务端地址
pub address: SocketAddr,
}
impl Display for ConnectInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!("count={} ,address={}", self.count, self.address))
}
}
impl ConnectInfo {
pub fn new(count: usize, address: SocketAddr) -> Self {
Self { count, address }
}
}
#[derive(Debug)]
pub struct HandshakeInfo {
//服务端公钥
#[cfg(feature = "server_encrypt")]
pub public_key: Option<RsaPublicKey>,
//服务端指纹
#[cfg(feature = "server_encrypt")]
pub finger: Option<String>,
//服务端版本
pub version: String,
}
impl Display for HandshakeInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
#[cfg(feature = "server_encrypt")]
return match &self.finger {
None => f.write_str(&format!("no_secret server version={}", self.version)),
Some(finger) => f.write_str(&format!(
"finger={} ,server version={}",
finger, self.version
)),
};
#[cfg(not(feature = "server_encrypt"))]
f.write_str(&format!("server version={}", self.version))
}
}
#[cfg(feature = "server_encrypt")]
impl HandshakeInfo {
pub fn new(public_key: RsaPublicKey, finger: String, version: String) -> Self {
Self {
public_key: Some(public_key),
finger: Some(finger),
version,
}
}
pub fn new_no_secret(version: String) -> Self {
Self {
public_key: None,
finger: None,
version,
}
}
}
#[cfg(not(feature = "server_encrypt"))]
impl HandshakeInfo {
pub fn new_no_secret(version: String) -> Self {
Self { version }
}
}
#[derive(Debug)]
pub struct RegisterInfo {
//本机虚拟IP
pub virtual_ip: Ipv4Addr,
//子网掩码
pub virtual_netmask: Ipv4Addr,
//虚拟网关
pub virtual_gateway: Ipv4Addr,
}
impl Display for RegisterInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!(
"ip={} ,netmask={} ,gateway={}",
self.virtual_ip, self.virtual_netmask, self.virtual_gateway,
))
}
}
impl RegisterInfo {
pub fn new(virtual_ip: Ipv4Addr, virtual_netmask: Ipv4Addr, virtual_gateway: Ipv4Addr) -> Self {
Self {
virtual_ip,
virtual_netmask,
virtual_gateway,
}
}
}
#[derive(Debug)]
pub struct ErrorInfo {
pub code: ErrorType,
pub msg: Option<String>,
pub source: Option<io::Error>,
}
impl Display for ErrorInfo {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.write_str(&format!("ErrorType={:?} ", self.code))?;
if let Some(msg) = &self.msg {
f.write_str(&format!(",msg={:?} ", msg))?;
}
if let Some(source) = &self.source {
f.write_str(&format!(",source={:?} ", source))?;
}
Ok(())
}
}
impl ErrorInfo {
pub fn new(code: ErrorType) -> Self {
Self {
code,
msg: None,
source: None,
}
}
pub fn new_msg(code: ErrorType, msg: String) -> Self {
Self {
code,
msg: Some(msg),
source: None,
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum ErrorType {
TokenError,
Disconnect,
AddressExhausted,
IpAlreadyExists,
InvalidIp,
LocalIpExists,
Unknown,
}
impl Into<u8> for ErrorType {
fn into(self) -> u8 {
match self {
ErrorType::TokenError => 1,
ErrorType::Disconnect => 2,
ErrorType::AddressExhausted => 3,
ErrorType::IpAlreadyExists => 4,
ErrorType::InvalidIp => 5,
ErrorType::LocalIpExists => 6,
ErrorType::Unknown => 255,
}
}
}
pub trait VntCallback: Clone + Send + Sync + 'static {
/// 启动成功
fn success(&self) {}
/// 创建网卡的信息
fn create_tun(&self, _info: DeviceInfo) {}
/// 连接
fn connect(&self, _info: ConnectInfo) {}
/// 握手,返回false则拒绝握手,可在此处检查服务端信息
fn handshake(&self, _info: HandshakeInfo) -> bool {
true
}
/// 注册,返回false则拒绝注册
fn register(&self, _info: RegisterInfo) -> bool {
true
}
/// 异常信息
fn error(&self, _info: ErrorInfo) {}
/// 服务停止
fn stop(&self) {}
}
-256
View File
@@ -1,256 +0,0 @@
use std::net::SocketAddr;
use protobuf::Message;
use std::net::UdpSocket;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use crate::channel::channel::Context;
use crate::channel::RouteKey;
use crate::cipher::{Cipher, RsaCipher};
use crate::proto::message::{HandshakeRequest, HandshakeResponse, SecretHandshakeRequest};
use crate::protocol::body::RSA_ENCRYPTION_RESERVED;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
pub enum HandshakeEnum {
NotSecret,
KeyError,
Timeout,
ServerError(String),
Other(String),
}
fn handshake_request_packet(secret: bool) -> crate::Result<NetPacket<Vec<u8>>> {
let mut request = HandshakeRequest::new();
request.secret = secret;
request.version = crate::VNT_VERSION.to_string();
let bytes = request.write_to_bytes()?;
let buf = vec![0u8; 12 + bytes.len()];
let mut net_packet = NetPacket::new(buf)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::HandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(net_packet)
}
fn secret_handshake_request_packet(
rsa_cipher: &RsaCipher,
token: String,
key: &[u8],
) -> crate::Result<NetPacket<Vec<u8>>> {
let mut request = SecretHandshakeRequest::new();
request.token = token;
request.key = key.to_vec();
let bytes = request.write_to_bytes()?;
let mut net_packet = NetPacket::new0(
12 + bytes.len(),
vec![0u8; 12 + bytes.len() + RSA_ENCRYPTION_RESERVED],
)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::SecretHandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(rsa_cipher.encrypt(&mut net_packet)?)
}
/// 第一次握手,拿到公钥
pub async fn handshake(
main_channel: &UdpSocket,
main_tcp_channel: Option<&mut TcpStream>,
server_address: SocketAddr,
secret: bool,
) -> Result<Option<RsaCipher>, HandshakeEnum> {
let request_packet = handshake_request_packet(secret).unwrap();
let send_buf = request_packet.buffer();
let mut recv_buf = [0u8; 10240];
let len = send_recv(
main_channel,
main_tcp_channel,
server_address,
send_buf,
&mut recv_buf,
)
.await?;
let net_packet = match NetPacket::new(&recv_buf[..len]) {
Ok(net_packet) => net_packet,
Err(e) => {
return Err(HandshakeEnum::Other(format!("net_packet {}", e)));
}
};
match net_packet.protocol() {
Protocol::Service => {
match service_packet::Protocol::from(net_packet.transport_protocol()) {
service_packet::Protocol::HandshakeResponse => {
match HandshakeResponse::parse_from_bytes(net_packet.payload()) {
Ok(response) => {
if !response.secret && secret {
//客户端要加密,服务端不支持加密
return Err(HandshakeEnum::NotSecret);
}
if secret {
//转换公钥
match RsaCipher::new(&response.public_key) {
Ok(rsa) => {
match rsa.finger() {
Ok(finger) => {
if finger != response.key_finger {
return Err(HandshakeEnum::Other(
"finger error".to_string(),
));
}
}
Err(e) => {
return Err(HandshakeEnum::Other(format!(
"finger {}",
e
)));
}
}
Ok(Some(rsa))
}
Err(e) => {
return Err(HandshakeEnum::Other(format!(
"RsaCipher {}",
e
)));
}
}
} else {
Ok(None)
}
}
Err(e) => {
return Err(HandshakeEnum::Other(format!("parse_from_bytes {}", e)));
}
}
}
_ => {
return Err(HandshakeEnum::Other("not match".to_string()));
}
}
}
_ => {
return Err(HandshakeEnum::Other("not match".to_string()));
}
}
}
async fn send_recv(
main_channel: &UdpSocket,
main_tcp_channel: Option<&mut TcpStream>,
server_address: SocketAddr,
send_buf: &[u8],
recv_buf: &mut [u8],
) -> Result<usize, HandshakeEnum> {
if let Some(main_tcp_channel) = main_tcp_channel {
let mut head = [0; 4];
let len = send_buf.len();
head[2] = (len >> 8) as u8;
head[3] = (len & 0xFF) as u8;
if let Err(e) = main_tcp_channel.write_all(&head).await {
return Err(HandshakeEnum::Other(format!("send error:{}", e)));
}
if let Err(e) = main_tcp_channel.write_all(send_buf).await {
return Err(HandshakeEnum::Other(format!("send error:{}", e)));
}
if let Err(e) = main_tcp_channel.read_exact(&mut head).await {
return Err(HandshakeEnum::Other(format!("read error:{}", e)));
}
let len = (((head[2] as u16) << 8) | head[3] as u16) as usize;
if len > recv_buf.len() {
return Err(HandshakeEnum::Other("too long".to_string()));
}
if let Err(e) = main_tcp_channel.read_exact(&mut recv_buf[..len]).await {
return Err(HandshakeEnum::Other(format!("read error:{}", e)));
}
Ok(len)
} else {
if let Err(e) = main_channel.send_to(send_buf, server_address) {
return Err(HandshakeEnum::Other(format!("send error:{}", e)));
}
match main_channel.recv_from(recv_buf) {
Ok((len, addr)) => {
if server_address != addr {
Err(HandshakeEnum::Other(format!("invalid data,from {}", addr)))
} else {
Ok(len)
}
}
Err(e) => Err(HandshakeEnum::Other(format!("receiver error:{}", e))),
}
}
}
/// 第二次握手,同步对称密钥,后续将使用对称加密
pub async fn secret_handshake(
main_channel: &UdpSocket,
main_tcp_channel: Option<&mut TcpStream>,
server_address: SocketAddr,
rsa_cipher: &RsaCipher,
server_cipher: &Cipher,
token: String,
) -> Result<(), HandshakeEnum> {
let secret_packet =
match secret_handshake_request_packet(rsa_cipher, token, server_cipher.key().unwrap()) {
Ok(secret_packet) => secret_packet,
Err(e) => {
return Err(HandshakeEnum::Other(format!(
"secret_handshake_request_packet {}",
e
)));
}
};
let send_buf = secret_packet.buffer();
let mut recv_buf = [0u8; 10240];
let len = send_recv(
main_channel,
main_tcp_channel,
server_address,
send_buf,
&mut recv_buf,
)
.await?;
let mut net_packet = match NetPacket::new(&mut recv_buf[..len]) {
Ok(net_packet) => net_packet,
Err(e) => {
return Err(HandshakeEnum::Other(format!("secret_net_packet {}", e)));
}
};
match server_cipher.decrypt_ipv4(&mut net_packet) {
Ok(_) => {
if net_packet.is_gateway()
&& net_packet.protocol() == Protocol::Service
&& service_packet::Protocol::from(net_packet.transport_protocol())
== service_packet::Protocol::SecretHandshakeResponse
{
Ok(())
} else {
Err(HandshakeEnum::Other("not match".to_string()))
}
}
Err(e) => Err(HandshakeEnum::Other(format!("decrypt_ipv4 {}", e))),
}
}
pub fn secret_handshake_req(
context: &Context,
server_address: SocketAddr,
rsa_cipher: &RsaCipher,
server_cipher: &Cipher,
token: String,
route_key: &RouteKey,
) -> crate::Result<()> {
let secret_packet =
secret_handshake_request_packet(rsa_cipher, token, server_cipher.key().unwrap())?;
if route_key.is_tcp() {
context.send_main(secret_packet.buffer(), server_address)?;
} else {
context.send_main_udp(secret_packet.buffer(), server_address)?;
}
Ok(())
}
+104
View File
@@ -0,0 +1,104 @@
use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use crossbeam_utils::atomic::AtomicCell;
use protobuf::Message;
use crate::channel::context::Context;
#[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher;
use crate::handle::{GATEWAY_IP, SELF_IP};
use crate::proto::message::HandshakeRequest;
#[cfg(feature = "server_encrypt")]
use crate::proto::message::SecretHandshakeRequest;
#[cfg(feature = "server_encrypt")]
use crate::protocol::body::RSA_ENCRYPTION_RESERVED;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
pub enum HandshakeEnum {
NotSecret,
KeyError,
Timeout,
ServerError(String),
Other(String),
}
#[derive(Clone)]
pub struct Handshake {
time: Arc<AtomicCell<Instant>>,
}
impl Handshake {
pub fn new() -> Self {
Handshake {
time: Arc::new(AtomicCell::new(Instant::now() - Duration::from_secs(60))),
}
}
pub fn send(&self, context: &Context, secret: bool, addr: SocketAddr) -> io::Result<()> {
let last = self.time.load();
//短时间不重复发送
if last.elapsed() < Duration::from_secs(3) {
return Ok(());
}
let request_packet = handshake_request_packet(secret)?;
log::info!("发送握手请求,secret={},{:?}", secret, addr);
context.send_default(request_packet.buffer(), addr)?;
self.time.store(Instant::now());
Ok(())
}
}
/// 第一次握手数据
pub fn handshake_request_packet(secret: bool) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = HandshakeRequest::new();
request.secret = secret;
request.version = crate::VNT_VERSION.to_string();
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("handshake_request_packet {:?}", e),
)
})?;
let buf = vec![0u8; 12 + bytes.len()];
let mut net_packet = NetPacket::new(buf)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::HandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(net_packet)
}
/// 第二次加密握手
#[cfg(feature = "server_encrypt")]
pub fn secret_handshake_request_packet(
rsa_cipher: &RsaCipher,
token: String,
key: &[u8],
) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = SecretHandshakeRequest::new();
request.token = token;
request.key = key.to_vec();
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("secret_handshake_request_packet {:?}", e),
)
})?;
let mut net_packet = NetPacket::new0(
12 + bytes.len(),
vec![0u8; 12 + bytes.len() + RSA_ENCRYPTION_RESERVED],
)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::SecretHandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(rsa_cipher.encrypt(&mut net_packet)?)
}
-256
View File
@@ -1,256 +0,0 @@
use std::io;
use std::net::{Ipv4Addr, ToSocketAddrs};
use std::sync::Arc;
use std::time::Duration;
use crate::channel::idle::Idle;
use crate::channel::sender::ChannelSender;
use crate::channel::Route;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use rand::prelude::SliceRandom;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::PingPacket;
use crate::protocol::{control_packet, NetPacket, Protocol, Version, MAX_TTL};
pub fn start_idle(mut worker: VntWorker, idle: Idle, sender: ChannelSender) {
tokio::spawn(async move {
tokio::select! {
_=worker.stop_wait()=>{
return;
}
rs=start_idle_(idle, sender)=>{
if let Err(e) = rs {
log::warn!("空闲检测任务停止:{:?}", e);
}
}
}
worker.stop_all();
});
}
async fn start_idle_(idle: Idle, sender: ChannelSender) -> io::Result<()> {
log::info!("启动空闲检查任务");
loop {
let (peer_ip, route) = idle.next_idle().await?;
log::info!("路由空闲 peer_ip:{:?},route:{:?}", peer_ip, route);
sender.remove_route(&peer_ip, route);
}
}
pub fn start_heartbeat(
mut worker: VntWorker,
sender: ChannelSender,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
server_address_str: String,
client_cipher: Cipher,
server_cipher: Cipher,
) {
tokio::spawn(async move {
tokio::select! {
_=worker.stop_wait()=>{
return;
}
rs=start_heartbeat_(sender, device_list, current_device,server_address_str,client_cipher,server_cipher)=>{
if let Err(e) = rs {
log::warn!("心跳任务停止:{:?}", e);
}
}
}
worker.stop_all();
});
}
fn heartbeat_packet(
ttl: u8,
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
client_cipher: &Cipher,
server_cipher: &Cipher,
gateway: bool,
src: Ipv4Addr,
dest: Ipv4Addr,
) -> NetPacket<[u8; 12 + 4 + ENCRYPTION_RESERVED]> {
let mut net_packet = NetPacket::new_encrypt([0u8; 12 + 4 + ENCRYPTION_RESERVED]).unwrap();
net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::Control);
net_packet.set_transport_protocol(control_packet::Protocol::Ping.into());
net_packet.first_set_ttl(ttl);
net_packet.set_source(src);
net_packet.set_destination(dest);
{
let mut ping = PingPacket::new(net_packet.payload_mut()).unwrap();
let epoch = { device_list.lock().0 };
ping.set_epoch(epoch);
ping.set_time(crate::handle::now_time() as u16);
}
if gateway {
net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet).unwrap();
} else {
client_cipher.encrypt_ipv4(&mut net_packet).unwrap();
}
net_packet
}
async fn start_heartbeat_(
sender: ChannelSender,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
server_address_str: String,
client_cipher: Cipher,
server_cipher: Cipher,
) -> io::Result<()> {
let mut count = 0;
log::info!("启动心跳任务");
loop {
if sender.is_close() {
return Ok(());
}
let mut current_dev = current_device.load();
//如果和服务端使用tcp连接,则维持udp洞的频率要更高些
if (sender.is_main_tcp() && count % 2 == 0) || (!sender.is_main_tcp() && count % 20 == 1) {
let mut packet = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED])?;
packet.set_version(Version::V1);
packet.set_gateway_flag(true);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::AddrRequest.into());
packet.first_set_ttl(MAX_TTL);
packet.set_source(current_dev.virtual_ip());
packet.set_destination(current_dev.virtual_gateway);
server_cipher.encrypt_ipv4(&mut packet)?;
let _ = sender.send_main_udp(packet.buffer(), current_dev.connect_server);
}
if count % 20 == 19 {
if let Ok(mut addr) = server_address_str.to_socket_addrs() {
if let Some(addr) = addr.next() {
if addr != current_dev.connect_server {
let mut tmp = current_dev.clone();
tmp.connect_server = addr;
log::info!(
"服务端地址变化,旧地址:{},新地址:{}",
current_dev.connect_server,
addr
);
if current_device.compare_exchange(current_dev, tmp).is_ok() {
current_dev.connect_server = addr;
}
}
}
}
}
let src = current_dev.virtual_ip();
let server_packet = heartbeat_packet(
MAX_TTL,
&device_list,
&client_cipher,
&server_cipher,
true,
src,
current_dev.virtual_gateway,
);
if let Err(e) = sender.send_main(server_packet.buffer(), current_dev.connect_server) {
log::warn!("connect_server:{:?},e:{:?}", current_dev.connect_server, e);
}
if count < 7 || count % 7 == 0 {
let mut route_list: Option<Vec<(Ipv4Addr, Vec<Route>)>> = None;
let peer_list = { device_list.lock().1.clone() };
for peer in peer_list {
if peer.virtual_ip == current_dev.virtual_ip {
continue;
}
let client_packet = heartbeat_packet(
MAX_TTL,
&device_list,
&client_cipher,
&server_cipher,
false,
src,
peer.virtual_ip,
);
if let Some(route) = sender.route_one(&peer.virtual_ip) {
if let Err(e) =
sender.try_send_by_key(client_packet.buffer(), &route.route_key())
{
log::warn!("virtual_ip:{},route:{:?},e:{:?}", peer.virtual_ip, route, e);
}
if route.is_p2p() {
continue;
}
} else {
//没有直连路由则发送到网关
if let Err(e) =
sender.send_main(client_packet.buffer(), current_dev.connect_server)
{
log::warn!(
"virtual_ip:{},connect_server:{:?},e:{:?}",
peer.virtual_ip,
current_dev.connect_server,
e
);
}
}
//再随机发送到其他地址,看有没有客户端符合转发条件
let route_list = route_list.get_or_insert_with(|| {
let mut l = sender.route_table();
l.shuffle(&mut rand::thread_rng());
l
});
let mut num = 0;
'a: for (peer_ip, route_list) in route_list.iter() {
for route in route_list {
if peer_ip != &peer.virtual_ip && route.is_p2p() {
if let Err(e) =
sender.try_send_by_key(client_packet.buffer(), &route.route_key())
{
log::warn!(
"virtual_ip:{},route:{:?},e:{:?}",
peer.virtual_ip,
route,
e
);
}
num += 1;
break;
}
if num >= 2 {
break 'a;
}
}
}
tokio::time::sleep(Duration::from_millis(1)).await;
}
} else {
for (peer_ip, route_list) in sender.route_table().iter() {
if peer_ip == &current_dev.virtual_gateway {
continue;
}
let client_packet = heartbeat_packet(
MAX_TTL,
&device_list,
&client_cipher,
&server_cipher,
false,
src,
*peer_ip,
);
for route in route_list {
if let Err(e) =
sender.try_send_by_key(client_packet.buffer(), &route.route_key())
{
log::warn!("peer_ip:{:?},route:{:?},e:{:?}", peer_ip, route, e);
}
tokio::time::sleep(Duration::from_millis(2)).await;
}
}
}
count += 1;
tokio::time::sleep(Duration::from_millis(5000)).await;
}
}
+67
View File
@@ -0,0 +1,67 @@
use std::sync::Arc;
use std::time::Duration;
use crossbeam_utils::atomic::AtomicCell;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::{BaseConfigInfo, CurrentDeviceInfo};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{control_packet, NetPacket, Protocol, Version, MAX_TTL};
use crate::util::Scheduler;
pub fn addr_request(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
server_cipher: Cipher,
_config: BaseConfigInfo,
) {
pub_address_request(
scheduler,
context,
current_device_info.clone(),
server_cipher,
);
}
pub fn pub_address_request(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
server_cipher: Cipher,
) {
addr_request0(&context, &current_device_info, &server_cipher);
// 17秒发送一次
let rs = scheduler.timeout(Duration::from_secs(17), |s| {
pub_address_request(s, context, current_device_info, server_cipher)
});
if !rs {
log::info!("定时任务停止");
}
}
pub fn addr_request0(
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
server_cipher: &Cipher,
) {
let current_dev = current_device.load();
if current_dev.connect_server.is_ipv4() && current_dev.status.online() {
// 如果连接的是ipv4服务,则探测公网端口
let gateway_ip = current_dev.virtual_gateway;
let src_ip = current_dev.virtual_ip;
let mut packet = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED]).unwrap();
packet.set_version(Version::V1);
packet.set_gateway_flag(true);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::AddrRequest.into());
packet.first_set_ttl(MAX_TTL);
packet.set_source(src_ip);
packet.set_destination(gateway_ip);
if let Err(e) = server_cipher.encrypt_ipv4(&mut packet) {
log::warn!("AddrRequest err={:?}", e)
} else {
context.try_send_all_main(packet.buffer(), current_dev.connect_server);
}
}
}
+249
View File
@@ -0,0 +1,249 @@
use std::io;
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::time::Duration;
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use rand::prelude::SliceRandom;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::PingPacket;
use crate::protocol::{control_packet, NetPacket, Protocol, Version};
use crate::util::Scheduler;
/// 定时发送心跳包
pub fn heartbeat(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
client_cipher: Cipher,
server_cipher: Cipher,
) {
heartbeat0(
&context,
&current_device_info.load(),
&device_list,
&client_cipher,
&server_cipher,
);
// 心跳包 3秒发送一次
let rs = scheduler.timeout(Duration::from_secs(3), |s| {
heartbeat(
s,
context,
current_device_info,
device_list,
client_cipher,
server_cipher,
)
});
if !rs {
log::info!("定时任务停止");
}
}
fn heartbeat0(
context: &Context,
current_device: &CurrentDeviceInfo,
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) {
let gateway_ip = current_device.virtual_gateway;
let src_ip = current_device.virtual_ip;
// 可能服务器ip发生变化,导致发送失败
let mut is_send_gateway = false;
match heartbeat_packet_server(device_list, server_cipher, src_ip, gateway_ip) {
Ok(net_packet) => {
if let Err(e) = context.send_default(net_packet.buffer(), current_device.connect_server)
{
log::warn!("heartbeat err={:?}", e)
} else {
is_send_gateway = true
}
}
Err(e) => {
log::error!("heartbeat_packet err={:?}", e);
}
}
for (dest_ip, routes) in context.route_table.route_table() {
let net_packet = if current_device.is_gateway(&dest_ip) {
if is_send_gateway {
continue;
}
heartbeat_packet_server(device_list, server_cipher, src_ip, gateway_ip)
} else {
heartbeat_packet_client(client_cipher, src_ip, dest_ip)
};
let net_packet = match net_packet {
Ok(net_packet) => net_packet,
Err(e) => {
log::error!("heartbeat_packet err={:?}", e);
continue;
}
};
for route in routes {
if let Err(e) = context.send_by_key(net_packet.buffer(), route.route_key()) {
log::warn!("heartbeat err={:?}", e)
}
}
}
let peer_list = { device_list.lock().1.clone() };
for peer in &peer_list {
if !peer.status.is_online() {
continue;
}
if current_device.is_gateway(&peer.virtual_ip) {
continue;
}
if current_device.status.offline() {
continue;
}
if context.route_table.route_one(&peer.virtual_ip).is_none() {
//路由为空,则向服务端地址发送
let net_packet = match heartbeat_packet_client(client_cipher, src_ip, peer.virtual_ip) {
Ok(net_packet) => net_packet,
Err(e) => {
log::error!("heartbeat_packet err={:?}", e);
continue;
}
};
if let Err(e) = context.send_default(net_packet.buffer(), current_device.connect_server)
{
log::error!("heartbeat_packet send_default err={:?}", e);
}
}
}
}
/// 客户端中继路径探测,延迟启动
pub fn client_relay(
scheduler: &Scheduler,
context: Context,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
client_cipher: Cipher,
) {
let rs = scheduler.timeout(Duration::from_secs(30), move |s| {
client_relay_(s, context, current_device, device_list, client_cipher)
});
if !rs {
log::info!("定时任务停止");
}
}
/// 客户端中继路径探测,每30秒探测一次
fn client_relay_(
scheduler: &Scheduler,
context: Context,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
client_cipher: Cipher,
) {
if let Err(e) = client_relay0(
&context,
&current_device.load(),
&device_list,
&client_cipher,
) {
log::error!("{:?}", e);
}
let rs = scheduler.timeout(Duration::from_secs(30), move |s| {
client_relay_(s, context, current_device, device_list, client_cipher)
});
if !rs {
log::info!("定时任务停止");
}
}
fn client_relay0(
context: &Context,
current_device: &CurrentDeviceInfo,
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
client_cipher: &Cipher,
) -> io::Result<()> {
// 离线了不再探测
if current_device.status.offline() {
return Ok(());
}
let peer_list = { device_list.lock().1.clone() };
let mut routes = context.route_table.route_table_p2p();
for peer in &peer_list {
if !peer.status.is_online() || peer.virtual_ip == current_device.virtual_ip {
continue;
}
if context
.route_table
.route_one_p2p(&peer.virtual_ip)
.is_some()
&& !context.first_latency()
{
continue;
}
let client_packet =
heartbeat_packet_client(client_cipher, current_device.virtual_ip, peer.virtual_ip)?;
//随机发送到其他地址,看有没有客户端符合转发条件
routes.shuffle(&mut rand::thread_rng());
for (index, (ip, route)) in routes.iter().enumerate() {
if current_device.is_gateway(ip) {
continue;
}
if let Err(e) = context.send_by_key(client_packet.buffer(), route.route_key()) {
log::error!("{:?}", e);
}
if index >= 2 {
break;
}
}
}
Ok(())
}
/// 构建心跳包
fn heartbeat_packet(
src: Ipv4Addr,
dest: Ipv4Addr,
) -> io::Result<NetPacket<[u8; 12 + 4 + ENCRYPTION_RESERVED]>> {
let mut net_packet = NetPacket::new_encrypt([0u8; 12 + 4 + ENCRYPTION_RESERVED])?;
net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::Control);
net_packet.set_transport_protocol(control_packet::Protocol::Ping.into());
net_packet.first_set_ttl(5);
net_packet.set_source(src);
net_packet.set_destination(dest);
let mut ping = PingPacket::new(net_packet.payload_mut())?;
ping.set_time(crate::handle::now_time() as u16);
Ok(net_packet)
}
fn heartbeat_packet_client(
client_cipher: &Cipher,
src: Ipv4Addr,
dest: Ipv4Addr,
) -> io::Result<NetPacket<[u8; 12 + 4 + ENCRYPTION_RESERVED]>> {
let mut net_packet = heartbeat_packet(src, dest)?;
client_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
fn heartbeat_packet_server(
device_list: &Mutex<(u16, Vec<PeerDeviceInfo>)>,
server_cipher: &Cipher,
src: Ipv4Addr,
dest: Ipv4Addr,
) -> io::Result<NetPacket<[u8; 12 + 4 + ENCRYPTION_RESERVED]>> {
let mut net_packet = heartbeat_packet(src, dest)?;
let mut ping = PingPacket::new(net_packet.payload_mut())?;
ping.set_epoch(device_list.lock().0);
net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
+211
View File
@@ -0,0 +1,211 @@
use std::io;
use std::net::{SocketAddr, ToSocketAddrs};
use std::sync::Arc;
use std::time::{Duration, Instant};
use crossbeam_utils::atomic::AtomicCell;
use mio::net::TcpStream;
use crate::channel::context::Context;
use crate::channel::idle::{Idle, IdleType};
use crate::channel::sender::AcceptSocketSender;
use crate::handle::callback::{ConnectInfo, ErrorType};
use crate::handle::handshaker::Handshake;
use crate::handle::{handshaker, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo};
use crate::util::Scheduler;
use crate::{ErrorInfo, VntCallback};
pub fn idle_route<Call: VntCallback>(
scheduler: &Scheduler,
idle: Idle,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
call: Call,
) {
let delay = idle_route0(&idle, &context, &current_device_info, &call);
let rs = scheduler.timeout(delay, move |s| {
idle_route(s, idle, context, current_device_info, call)
});
if !rs {
log::info!("定时任务停止");
}
}
pub fn idle_gateway<Call: VntCallback>(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
config: BaseConfigInfo,
tcp_socket_sender: AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
call: Call,
connect_count: usize,
handshake: Handshake,
) {
let time = Instant::now();
idle_gateway_(
scheduler,
context,
current_device_info,
config,
tcp_socket_sender,
call,
connect_count,
handshake,
time,
);
}
pub fn idle_gateway_<Call: VntCallback>(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
config: BaseConfigInfo,
tcp_socket_sender: AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
call: Call,
mut connect_count: usize,
handshake: Handshake,
mut time: Instant,
) {
idle_gateway0(
&context,
&current_device_info,
&config,
&tcp_socket_sender,
&call,
&mut connect_count,
&handshake,
&mut time,
);
let rs = scheduler.timeout(Duration::from_secs(5), move |s| {
idle_gateway_(
s,
context,
current_device_info,
config,
tcp_socket_sender,
call,
connect_count,
handshake,
time,
)
});
if !rs {
log::info!("定时任务停止");
}
}
fn idle_gateway0<Call: VntCallback>(
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
config: &BaseConfigInfo,
tcp_socket_sender: &AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
call: &Call,
connect_count: &mut usize,
handshake: &Handshake,
time: &mut Instant,
) {
if let Err(e) = check_gateway_channel(
context,
current_device,
config,
tcp_socket_sender,
call,
connect_count,
handshake,
time,
) {
let cur = current_device.load();
call.error(ErrorInfo::new_msg(
ErrorType::Disconnect,
format!("connect:{},error:{:?}", cur.connect_server, e),
));
}
}
fn idle_route0<Call: VntCallback>(
idle: &Idle,
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
call: &Call,
) -> Duration {
let cur = current_device.load();
match idle.next_idle() {
IdleType::Timeout(ip, route) => {
log::info!("route Timeout {:?},{:?}", ip, route);
context.remove_route(&ip, route.route_key());
if cur.is_gateway(&ip) {
//网关路由过期,则需要改变状态
crate::handle::change_status(current_device, ConnectStatus::Connecting);
call.error(ErrorInfo::new(ErrorType::Disconnect));
}
Duration::from_millis(100)
}
IdleType::Sleep(duration) => duration,
IdleType::None => Duration::from_millis(3000),
}
}
fn check_gateway_channel<Call: VntCallback>(
context: &Context,
current_device_info: &AtomicCell<CurrentDeviceInfo>,
config: &BaseConfigInfo,
tcp_socket_sender: &AcceptSocketSender<(TcpStream, SocketAddr, Option<Vec<u8>>)>,
call: &Call,
count: &mut usize,
handshake: &Handshake,
time: &mut Instant,
) -> io::Result<()> {
let mut current_device = current_device_info.load();
if current_device.status.offline() {
*count += 1;
if time.elapsed() < Duration::from_secs(6 * 60) {
// 探测服务器地址
current_device = domain_request0(current_device_info, config);
*time = Instant::now()
}
//需要重连
call.connect(ConnectInfo::new(*count, current_device.connect_server));
log::info!("发送握手请求,{:?}", config);
if let Err(e) = handshake.send(context, config.client_secret, current_device.connect_server)
{
log::warn!("{:?}", e);
if context.is_main_tcp() {
let request_packet = handshaker::handshake_request_packet(config.client_secret)?;
//tcp需要重连
let tcp_stream = std::net::TcpStream::connect_timeout(
&current_device.connect_server,
Duration::from_secs(5),
)?;
tcp_stream.set_nonblocking(true)?;
if let Err(e) = tcp_socket_sender.try_add_socket((
TcpStream::from_std(tcp_stream),
current_device.connect_server,
Some(request_packet.into_buffer()),
)) {
log::warn!("{:?}", e)
}
}
}
}
Ok(())
}
pub fn domain_request0(
current_device: &AtomicCell<CurrentDeviceInfo>,
config: &BaseConfigInfo,
) -> CurrentDeviceInfo {
let mut current_dev = current_device.load();
// 探测服务端地址变化
if let Ok(mut addr) = config.server_addr.to_socket_addrs() {
if let Some(addr) = addr.next() {
if addr != current_dev.connect_server {
let mut tmp = current_dev.clone();
tmp.connect_server = addr;
let rs = current_device.compare_exchange(current_dev, tmp);
current_dev.connect_server = addr;
log::info!(
"服务端地址变化,旧地址:{},新地址:{},替换结果:{}",
current_dev.connect_server,
addr,
rs.is_ok()
);
}
}
}
current_dev
}
+19
View File
@@ -0,0 +1,19 @@
mod heartbeat;
pub use heartbeat::client_relay;
pub use heartbeat::heartbeat;
mod re_nat_type;
pub use re_nat_type::retrieve_nat_type;
mod addr_request;
pub use addr_request::addr_request;
mod punch;
pub use punch::*;
mod idle;
pub use idle::idle_gateway;
pub use idle::idle_route;
mod up_status;
pub use up_status::*;
+279
View File
@@ -0,0 +1,279 @@
use std::net::Ipv4Addr;
use std::sync::mpsc::{sync_channel, Receiver, SyncSender};
use std::sync::Arc;
use std::time::Duration;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use protobuf::Message;
use rand::prelude::SliceRandom;
use crate::channel::context::Context;
use crate::channel::punch::{NatInfo, NatType, Punch};
use crate::cipher::Cipher;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::nat::NatTest;
use crate::proto::message::{PunchInfo, PunchNatType};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{control_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL};
use crate::util::Scheduler;
#[derive(Clone)]
pub struct PunchSender {
sender_self: SyncSender<(Ipv4Addr, NatInfo)>,
sender_peer: SyncSender<(Ipv4Addr, NatInfo)>,
sender_cone_self: SyncSender<(Ipv4Addr, NatInfo)>,
sender_cone_peer: SyncSender<(Ipv4Addr, NatInfo)>,
}
impl PunchSender {
pub fn send(&self, src_peer: bool, ip: Ipv4Addr, info: NatInfo) -> bool {
log::info!(
"发送打洞协商消息,是否对端发起:{},ip:{},info:{:?}",
src_peer,
ip,
info
);
let sender = match info.nat_type {
NatType::Symmetric => {
if src_peer {
&self.sender_peer
} else {
&self.sender_self
}
}
NatType::Cone => {
if src_peer {
&self.sender_cone_peer
} else {
&self.sender_cone_self
}
}
};
sender.try_send((ip, info)).is_ok()
}
}
pub struct PunchReceiver {
receiver_peer: Receiver<(Ipv4Addr, NatInfo)>,
receiver_self: Receiver<(Ipv4Addr, NatInfo)>,
receiver_cone_peer: Receiver<(Ipv4Addr, NatInfo)>,
receiver_cone_self: Receiver<(Ipv4Addr, NatInfo)>,
}
pub fn punch_channel() -> (PunchSender, PunchReceiver) {
let (sender_self, receiver_self) = sync_channel(1);
let (sender_peer, receiver_peer) = sync_channel(1);
let (sender_cone_peer, receiver_cone_peer) = sync_channel(1);
let (sender_cone_self, receiver_cone_self) = sync_channel(1);
(
PunchSender {
sender_self,
sender_peer,
sender_cone_peer,
sender_cone_self,
},
PunchReceiver {
receiver_peer,
receiver_self,
receiver_cone_peer,
receiver_cone_self,
},
)
}
pub fn punch(
scheduler: &Scheduler,
context: Context,
nat_test: NatTest,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
receiver: PunchReceiver,
punch: Punch,
) {
punch_request(
scheduler,
context,
nat_test,
device_list,
current_device.clone(),
client_cipher.clone(),
0,
);
let f = |receiver: Receiver<(Ipv4Addr, NatInfo)>| {
let punch = punch.clone();
let current_device = current_device.clone();
let client_cipher = client_cipher.clone();
thread::Builder::new()
.name("punch".into())
.spawn(move || {
punch_start(receiver, punch, current_device, client_cipher);
})
.expect("punch");
};
f(receiver.receiver_peer);
f(receiver.receiver_self);
f(receiver.receiver_cone_peer);
f(receiver.receiver_cone_self);
}
/// 接收打洞消息,配合对端打洞
fn punch_start(
receiver: Receiver<(Ipv4Addr, NatInfo)>,
mut punch: Punch,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) {
while let Ok((peer_ip, nat_info)) = receiver.recv() {
let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED]).unwrap();
packet.set_version(Version::V1);
packet.first_set_ttl(1);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::PunchRequest.into());
packet.set_source(current_device.load().virtual_ip());
packet.set_destination(peer_ip);
log::info!("发起打洞,目标:{:?},{:?}", peer_ip, nat_info);
if let Err(e) = client_cipher.encrypt_ipv4(&mut packet) {
log::error!("{:?}", e);
continue;
}
if let Err(e) = punch.punch(packet.buffer(), peer_ip, nat_info) {
log::warn!("{:?}", e)
}
}
}
/// 定时发起打洞请求
fn punch_request(
scheduler: &Scheduler,
context: Context,
nat_test: NatTest,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
count: usize,
) {
let curr = current_device.load();
let secs = if curr.status.online() {
if let Err(e) = punch0(&context, &nat_test, &device_list, curr, &client_cipher) {
log::warn!("{:?}", e)
}
let sleep_time = [3, 5, 7, 11, 13, 17, 19, 23, 29];
Duration::from_secs(sleep_time[count % sleep_time.len()])
} else {
Duration::from_secs(3)
};
let rs = scheduler.timeout(secs, move |s| {
punch_request(
s,
context,
nat_test,
device_list,
current_device,
client_cipher,
count + 1,
);
});
if !rs {
log::info!("定时任务停止");
}
}
/// 随机对需要打洞的客户端发起打洞请求
fn punch0(
context: &Context,
nat_test: &NatTest,
device_list: &Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
current_device: CurrentDeviceInfo,
client_cipher: &Cipher,
) -> io::Result<()> {
let nat_info = nat_test.nat_info();
let current_ip = current_device.virtual_ip;
let mut list: Vec<PeerDeviceInfo> = device_list
.lock()
.1
.iter()
.filter(|info| info.status.is_online() && info.virtual_ip > current_ip)
.cloned()
.collect();
list.shuffle(&mut rand::thread_rng());
let mut count = 0;
// // 优先没打洞的 need_punch会过滤掉已经打洞成功的
// list.sort_by(|v1, v2| {
// if context.route_table.route_one_p2p(&v1.virtual_ip).is_none() {
// Ordering::Less
// } else if context.route_table.route_one_p2p(&v2.virtual_ip).is_none() {
// Ordering::Greater
// } else {
// Ordering::Equal
// }
// });
for info in list {
if !info.status.is_online() {
continue;
}
if info.virtual_ip <= current_device.virtual_ip {
continue;
}
if !context.route_table.need_punch(&info.virtual_ip) {
continue;
}
count += 1;
if count > 2 {
break;
}
let packet = punch_packet(
client_cipher,
current_device.virtual_ip(),
&nat_info,
info.virtual_ip,
)?;
log::info!(
"发起打洞协商请求,目标:{:?},{:?}",
info.virtual_ip,
nat_info
);
context.send_default(packet.buffer(), current_device.connect_server)?;
}
Ok(())
}
fn punch_packet(
client_cipher: &Cipher,
virtual_ip: Ipv4Addr,
nat_info: &NatInfo,
dest: Ipv4Addr,
) -> io::Result<NetPacket<Vec<u8>>> {
let mut punch_reply = PunchInfo::new();
punch_reply.reply = false;
punch_reply.public_ip_list = nat_info
.public_ips
.iter()
.map(|ip| u32::from_be_bytes(ip.octets()))
.collect();
punch_reply.public_port = nat_info.public_ports.get(0).map_or(0, |v| *v as u32);
punch_reply.public_ports = nat_info.public_ports.iter().map(|e| *e as u32).collect();
punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.local_ip = u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED));
punch_reply.local_port = nat_info.udp_ports[0] as u32;
punch_reply.tcp_port = nat_info.tcp_port as u32;
punch_reply.udp_ports = nat_info.udp_ports.iter().map(|e| *e as u32).collect();
if let Some(ipv6) = nat_info.ipv6 {
punch_reply.ipv6_port = nat_info.udp_ports[0] as u32;
punch_reply.ipv6 = ipv6.octets().to_vec();
}
punch_reply.nat_type = protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
log::info!("请求打洞={:?}", punch_reply);
let bytes = punch_reply
.write_to_bytes()
.map_err(|e| io::Error::new(io::ErrorKind::Other, format!("punch_packet {:?}", e)))?;
let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::OtherTurn);
net_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip);
net_packet.set_destination(dest);
net_packet.set_payload(&bytes)?;
client_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
+48
View File
@@ -0,0 +1,48 @@
use std::thread;
use std::time::Duration;
use crate::channel::context::Context;
use crate::channel::sender::AcceptSocketSender;
use crate::nat;
use crate::nat::NatTest;
use crate::util::Scheduler;
/// 10分钟探测一次nat
pub fn retrieve_nat_type(
scheduler: &Scheduler,
context: Context,
nat_test: NatTest,
udp_socket_sender: AcceptSocketSender<Option<Vec<mio::net::UdpSocket>>>,
) {
retrieve_nat_type0(context.clone(), nat_test.clone(), udp_socket_sender.clone());
scheduler.timeout(Duration::from_secs(60 * 10), move |s| {
retrieve_nat_type(s, context, nat_test, udp_socket_sender)
});
}
fn retrieve_nat_type0(
context: Context,
nat_test: NatTest,
udp_socket_sender: AcceptSocketSender<Option<Vec<mio::net::UdpSocket>>>,
) {
thread::Builder::new()
.name("natTest".into())
.spawn(move || {
if nat_test.can_update() {
let local_ipv4 = nat::local_ipv4();
let local_ipv6 = nat::local_ipv6();
match nat_test.re_test(local_ipv4, local_ipv6) {
Ok(nat_info) => {
log::info!("当前nat信息:{:?}", nat_info);
if let Err(e) = context.switch(nat_info.nat_type, &udp_socket_sender) {
log::warn!("{:?}", e);
}
}
Err(e) => {
log::warn!("nat re_test {:?}", e);
}
};
}
})
.expect("natTest");
}
+104
View File
@@ -0,0 +1,104 @@
use crate::channel::context::Context;
use crate::handle::CurrentDeviceInfo;
use crate::proto::message::{ClientStatusInfo, PunchNatType, RouteItem};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, HEAD_LEN, MAX_TTL};
use crate::util::{Scheduler, WatchSingleU64Adder, WatchU64Adder};
use crossbeam_utils::atomic::AtomicCell;
use protobuf::Message;
use std::io;
use std::sync::Arc;
use std::time::Duration;
/// 上报状态给服务器
pub fn up_status(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
down_count_watcher: WatchU64Adder,
up_count_watcher: WatchSingleU64Adder,
) {
let _ = scheduler.timeout(Duration::from_secs(60), move |x| {
up_status0(
x,
context,
current_device_info,
down_count_watcher,
up_count_watcher,
)
});
}
fn up_status0(
scheduler: &Scheduler,
context: Context,
current_device_info: Arc<AtomicCell<CurrentDeviceInfo>>,
down_count_watcher: WatchU64Adder,
up_count_watcher: WatchSingleU64Adder,
) {
if let Err(e) = send_up_status_packet(
&context,
&current_device_info,
&down_count_watcher,
&up_count_watcher,
) {
log::warn!("{:?}", e)
}
let rs = scheduler.timeout(Duration::from_secs(10 * 60), move |x| {
up_status0(
x,
context,
current_device_info,
down_count_watcher,
up_count_watcher,
)
});
if !rs {
log::info!("定时任务停止");
}
}
fn send_up_status_packet(
context: &Context,
current_device_info: &AtomicCell<CurrentDeviceInfo>,
down_count_watcher: &WatchU64Adder,
up_count_watcher: &WatchSingleU64Adder,
) -> io::Result<()> {
let device_info = current_device_info.load();
if device_info.status.offline() {
return Ok(());
}
let routes = context.route_table.route_table_p2p();
if routes.is_empty() {
return Ok(());
}
let mut message = ClientStatusInfo::new();
message.source = device_info.virtual_ip.into();
for (ip, _) in routes {
let mut item = RouteItem::new();
item.next_ip = ip.into();
message.p2p_list.push(item);
}
message.up_stream = up_count_watcher.get();
message.down_stream = down_count_watcher.get();
message.nat_type = protobuf::EnumOrUnknown::new(if context.is_cone() {
PunchNatType::Cone
} else {
PunchNatType::Symmetric
});
let buf = message
.write_to_bytes()
.map_err(|e| io::Error::new(io::ErrorKind::Other, format!("up_status_packet {:?}", e)))?;
let mut net_packet =
NetPacket::new_encrypt(vec![0; HEAD_LEN + buf.len() + ENCRYPTION_RESERVED])?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol_into(service_packet::Protocol::ClientStatusInfo);
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(device_info.virtual_ip);
net_packet.set_destination(device_info.virtual_gateway);
net_packet.set_payload(&buf)?;
context.send_default(net_packet.buffer(), device_info.connect_server)?;
Ok(())
}
+112 -12
View File
@@ -1,12 +1,16 @@
use crossbeam_utils::atomic::AtomicCell;
use std::net::{Ipv4Addr, SocketAddr};
pub mod handshake_handler;
pub mod heartbeat_handler;
pub mod punch_handler;
pub mod recv_handler;
pub mod registration_handler;
pub mod callback;
pub mod handshaker;
pub mod maintain;
pub mod recv_data;
pub mod registrar;
pub mod tun_tap;
const SELF_IP: Ipv4Addr = Ipv4Addr::new(0, 0, 0, 2);
const GATEWAY_IP: Ipv4Addr = Ipv4Addr::new(0, 0, 0, 1);
pub fn now_time() -> u64 {
let now = std::time::SystemTime::now();
if let Ok(timestamp) = now.duration_since(std::time::UNIX_EPOCH) {
@@ -41,12 +45,48 @@ impl PeerDeviceInfo {
}
}
#[derive(Clone, Debug)]
pub struct BaseConfigInfo {
pub name: String,
pub token: String,
pub ip: Option<Ipv4Addr>,
pub client_secret: bool,
pub device_id: String,
pub server_addr: String,
}
impl BaseConfigInfo {
pub fn new(
name: String,
token: String,
ip: Option<Ipv4Addr>,
client_secret: bool,
device_id: String,
server_addr: String,
) -> Self {
Self {
name,
token,
ip,
client_secret,
device_id,
server_addr,
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd)]
pub enum PeerDeviceStatus {
Online,
Offline,
}
impl PeerDeviceStatus {
pub fn is_online(&self) -> bool {
self == &PeerDeviceStatus::Online
}
}
impl Into<u8> for PeerDeviceStatus {
fn into(self) -> u8 {
match self {
@@ -71,29 +111,43 @@ pub enum ConnectStatus {
Connected,
}
impl ConnectStatus {
pub fn online(&self) -> bool {
self == &ConnectStatus::Connected
}
pub fn offline(&self) -> bool {
self == &ConnectStatus::Connecting
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub struct CurrentDeviceInfo {
virtual_ip: Ipv4Addr,
pub virtual_gateway: Ipv4Addr,
//本机虚拟IP
pub virtual_ip: Ipv4Addr,
//子网掩码
pub virtual_netmask: Ipv4Addr,
//虚拟网关
pub virtual_gateway: Ipv4Addr,
//网络地址
pub virtual_network: Ipv4Addr,
//直接广播地址
pub broadcast_address: Ipv4Addr,
pub broadcast_ip: Ipv4Addr,
//链接的服务器地址
pub connect_server: SocketAddr,
//连接状态
pub status: ConnectStatus,
}
impl CurrentDeviceInfo {
pub fn new(
virtual_ip: Ipv4Addr,
virtual_gateway: Ipv4Addr,
virtual_netmask: Ipv4Addr,
virtual_gateway: Ipv4Addr,
connect_server: SocketAddr,
) -> Self {
let broadcast_address = (!u32::from_be_bytes(virtual_netmask.octets()))
let broadcast_ip = (!u32::from_be_bytes(virtual_netmask.octets()))
| u32::from_be_bytes(virtual_gateway.octets());
let broadcast_address = Ipv4Addr::from(broadcast_address);
let broadcast_ip = Ipv4Addr::from(broadcast_ip);
let virtual_network = u32::from_be_bytes(virtual_netmask.octets())
& u32::from_be_bytes(virtual_gateway.octets());
let virtual_network = Ipv4Addr::from(virtual_network);
@@ -102,10 +156,40 @@ impl CurrentDeviceInfo {
virtual_netmask,
virtual_gateway,
virtual_network,
broadcast_address,
broadcast_ip,
connect_server,
status: ConnectStatus::Connecting,
}
}
pub fn new0(connect_server: SocketAddr) -> Self {
Self {
virtual_ip: Ipv4Addr::UNSPECIFIED,
virtual_gateway: Ipv4Addr::UNSPECIFIED,
virtual_netmask: Ipv4Addr::UNSPECIFIED,
virtual_network: Ipv4Addr::UNSPECIFIED,
broadcast_ip: Ipv4Addr::UNSPECIFIED,
connect_server,
status: ConnectStatus::Connecting,
}
}
pub fn update(
&mut self,
virtual_ip: Ipv4Addr,
virtual_netmask: Ipv4Addr,
virtual_gateway: Ipv4Addr,
) {
let broadcast_ip = (!u32::from_be_bytes(virtual_netmask.octets()))
| u32::from_be_bytes(virtual_gateway.octets());
let broadcast_ip = Ipv4Addr::from(broadcast_ip);
let virtual_network = u32::from_be_bytes(virtual_netmask.octets())
& u32::from_be_bytes(virtual_gateway.octets());
let virtual_network = Ipv4Addr::from(virtual_network);
self.virtual_ip = virtual_ip;
self.virtual_netmask = virtual_netmask;
self.virtual_gateway = virtual_gateway;
self.broadcast_ip = broadcast_ip;
self.virtual_network = virtual_network;
}
#[inline]
pub fn virtual_ip(&self) -> Ipv4Addr {
self.virtual_ip
@@ -114,4 +198,20 @@ impl CurrentDeviceInfo {
pub fn virtual_gateway(&self) -> Ipv4Addr {
self.virtual_gateway
}
pub fn is_gateway(&self, ip: &Ipv4Addr) -> bool {
&self.virtual_gateway == ip || ip == &GATEWAY_IP
}
}
pub fn change_status(
current_device: &AtomicCell<CurrentDeviceInfo>,
connect_status: ConnectStatus,
) -> CurrentDeviceInfo {
loop {
let cur = current_device.load();
let mut new_info = cur;
new_info.status = connect_status;
if current_device.compare_exchange(cur, new_info).is_ok() {
return new_info;
}
}
}
-179
View File
@@ -1,179 +0,0 @@
use crate::channel::punch::{NatInfo, Punch};
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo};
use crate::nat::NatTest;
use crate::proto::message::{PunchInfo, PunchNatType};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{control_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use protobuf::Message;
use rand::prelude::SliceRandom;
use std::io;
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::mpsc::Receiver;
pub fn start(
mut worker: VntWorker,
receiver: Receiver<(Ipv4Addr, NatInfo)>,
punch: Punch,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) {
tokio::spawn(async move {
tokio::select! {
_=start0(receiver, punch, current_device,client_cipher)=>{}
_=worker.stop_wait()=>{
return;
}
}
worker.stop_all();
});
}
pub async fn start0(
mut receiver: Receiver<(Ipv4Addr, NatInfo)>,
mut punch: Punch,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) {
log::info!("启动打洞任务");
while let Some((peer_ip, nat_info)) = receiver.recv().await {
if let Err(e) = start_(
&client_cipher,
&mut punch,
&current_device,
peer_ip,
nat_info,
)
.await
{
log::warn!("网络打洞异常 {:?}", e);
}
}
}
async fn start_(
client_cipher: &Cipher,
punch: &mut Punch,
current_device: &Arc<AtomicCell<CurrentDeviceInfo>>,
peer_ip: Ipv4Addr,
nat_info: NatInfo,
) -> io::Result<()> {
let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED])?;
packet.set_version(Version::V1);
packet.first_set_ttl(1);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::PunchRequest.into());
packet.set_source(current_device.load().virtual_ip());
packet.set_destination(peer_ip);
log::info!("发起打洞,目标:{:?},{:?}", peer_ip, nat_info);
client_cipher.encrypt_ipv4(&mut packet)?;
punch.punch(packet.buffer(), peer_ip, nat_info).await
}
pub async fn start_punch(
mut worker: VntWorker,
nat_test: NatTest,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
sender: ChannelSender,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) {
let mut num = 0;
let sleep_time = [3, 5, 7, 11, 13, 17, 19, 23, 29];
log::info!("启动发起打洞请求任务");
loop {
if sender.is_close() {
break;
}
tokio::select! {
rs= start_punch_(Duration::from_secs(sleep_time[num % sleep_time.len()]),&nat_test, &device_list,
&sender, &current_device,&client_cipher)=>{
if let Err(e) = rs {
log::warn!("打洞处理任务异常 {:?}", e);
}
}
_=worker.stop_wait()=>{
break;
}
}
num += 1;
}
}
async fn start_punch_(
sleep_time: Duration,
nat_test: &NatTest,
device_list: &Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
sender: &ChannelSender,
current_device: &Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: &Cipher,
) -> crate::Result<()> {
let current_device = current_device.load();
let nat_info = nat_test.nat_info();
let mut list = device_list.lock().clone().1;
list.shuffle(&mut rand::thread_rng());
let mut count = 0;
for info in list {
if info.virtual_ip <= current_device.virtual_ip {
continue;
}
if !sender.need_punch(&info.virtual_ip) {
continue;
}
count += 1;
if count > 2 {
break;
}
let packet = punch_packet(
client_cipher,
current_device.virtual_ip(),
&nat_info,
info.virtual_ip,
)
.unwrap();
let _ = sender.send_main(packet.buffer(), current_device.connect_server);
}
tokio::time::sleep(sleep_time).await;
Ok(())
}
pub fn punch_packet(
client_cipher: &Cipher,
virtual_ip: Ipv4Addr,
nat_info: &NatInfo,
dest: Ipv4Addr,
) -> crate::Result<NetPacket<Vec<u8>>> {
let mut punch_reply = PunchInfo::new();
punch_reply.reply = false;
punch_reply.public_ip_list = nat_info
.public_ips
.iter()
.map(|ip| u32::from_be_bytes(ip.octets()))
.collect();
punch_reply.public_port = nat_info.public_port as u32;
punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.local_ip = u32::from_be_bytes(nat_info.local_ipv4_addr.ip().octets());
punch_reply.local_port = nat_info.local_ipv4_addr.port() as u32;
if !nat_info.ipv6_addr.ip().is_unspecified() {
punch_reply.ipv6_port = nat_info.ipv6_addr.port() as u32;
punch_reply.ipv6 = nat_info.ipv6_addr.ip().octets().to_vec();
}
punch_reply.nat_type = protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
let bytes = punch_reply.write_to_bytes()?;
let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::OtherTurn);
net_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip);
net_packet.set_destination(dest);
net_packet.set_payload(&bytes)?;
client_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
+337
View File
@@ -0,0 +1,337 @@
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;
use tun::device::IFace;
use tun::Device;
use crate::channel::context::Context;
use crate::channel::punch::NatInfo;
use crate::channel::{Route, RouteKey};
use crate::cipher::Cipher;
use crate::external_route::AllowExternalRoute;
use crate::handle::maintain::PunchSender;
use crate::handle::recv_data::PacketHandler;
use crate::handle::CurrentDeviceInfo;
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::{IpProxyMap, ProxyHandler};
use crate::nat::NatTest;
use crate::proto::message::{PunchInfo, PunchNatType};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::ControlPacket;
use crate::protocol::{
control_packet, ip_turn_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL,
};
/// 处理来源于客户端的包
#[derive(Clone)]
pub struct ClientPacketHandler {
device: Arc<Device>,
client_cipher: Cipher,
punch_sender: PunchSender,
peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>>,
nat_test: NatTest,
route: AllowExternalRoute,
#[cfg(feature = "ip_proxy")]
ip_proxy_map: Option<IpProxyMap>,
}
impl ClientPacketHandler {
pub fn new(
device: Arc<Device>,
client_cipher: Cipher,
punch_sender: PunchSender,
peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>>,
nat_test: NatTest,
route: AllowExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
) -> Self {
Self {
device,
client_cipher,
punch_sender,
peer_nat_info_map,
nat_test,
route,
#[cfg(feature = "ip_proxy")]
ip_proxy_map,
}
}
}
impl PacketHandler for ClientPacketHandler {
fn handle(
&self,
mut net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
context: &Context,
current_device: &CurrentDeviceInfo,
) -> io::Result<()> {
self.client_cipher.decrypt_ipv4(&mut net_packet)?;
context
.route_table
.update_read_time(&net_packet.source(), &route_key);
match net_packet.protocol() {
Protocol::Service => {}
Protocol::Error => {}
Protocol::Control => {
self.control(context, current_device, net_packet, route_key)?;
}
Protocol::IpTurn => {
self.ip_turn(net_packet, context, current_device, route_key)?;
}
Protocol::OtherTurn => {
self.other_turn(context, current_device, net_packet, route_key)?;
}
Protocol::Unknown(_) => {}
}
Ok(())
}
}
impl ClientPacketHandler {
fn ip_turn(
&self,
mut net_packet: NetPacket<&mut [u8]>,
context: &Context,
current_device: &CurrentDeviceInfo,
route_key: RouteKey,
) -> io::Result<()> {
let destination = net_packet.destination();
let source = net_packet.source();
match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) {
ip_turn_packet::Protocol::Ipv4 => {
let mut ipv4 = IpV4Packet::new(net_packet.payload_mut())?;
match ipv4.protocol() {
ipv4::protocol::Protocol::Icmp => {
if ipv4.destination_ip() == destination {
let mut icmp_packet = icmp::IcmpPacket::new(ipv4.payload_mut())?;
if icmp_packet.kind() == Kind::EchoRequest {
//开启ping
icmp_packet.set_kind(Kind::EchoReply);
icmp_packet.update_checksum();
ipv4.set_source_ip(destination);
ipv4.set_destination_ip(source);
ipv4.update_checksum();
net_packet.set_source(destination);
net_packet.set_destination(source);
//不管加不加密,和接收到的数据长度都一致
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.send_by_key(net_packet.buffer(), route_key)?;
return Ok(());
}
}
}
_ => {}
}
// ip代理只关心实际目标
let real_dest = ipv4.destination_ip();
if real_dest != destination
&& !(real_dest.is_broadcast()
|| real_dest.is_multicast()
|| real_dest == current_device.broadcast_ip
|| real_dest.is_unspecified())
{
if !self.route.allow(&real_dest) {
//拦截不符合的目标
return Ok(());
}
#[cfg(feature = "ip_proxy")]
if let Some(ip_proxy_map) = &self.ip_proxy_map {
if ip_proxy_map.recv_handle(&mut ipv4, source, destination)? {
return Ok(());
}
}
}
self.device.write(net_packet.payload())?;
}
ip_turn_packet::Protocol::Ipv4Broadcast => {
//客户端不帮忙转发广播包,所以不会出现这种类型的数据
}
ip_turn_packet::Protocol::Unknown(_) => {}
}
Ok(())
}
fn control(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
mut net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
let source = net_packet.source();
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PingPacket(_) => {
net_packet.set_transport_protocol(control_packet::Protocol::Pong.into());
net_packet.set_source(current_device.virtual_ip);
net_packet.set_destination(source);
net_packet.first_set_ttl(MAX_TTL);
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.send_by_key(net_packet.buffer(), route_key)?;
let route = Route::from_default_rt(route_key, metric);
context.route_table.add_route_if_absent(source, route);
}
ControlPacket::PongPacket(pong_packet) => {
let current_time = crate::handle::now_time() as u16;
if current_time < pong_packet.time() {
return Ok(());
}
let rt = (current_time - pong_packet.time()) as i64;
let route = Route::from(route_key, metric, rt);
context.route_table.add_route(source, route);
}
ControlPacket::PunchRequest => {
log::info!("PunchRequest={:?},source={}", route_key, source);
if context.use_channel_type().is_only_relay() {
return Ok(());
}
//回应
net_packet.set_transport_protocol(control_packet::Protocol::PunchResponse.into());
net_packet.set_source(current_device.virtual_ip);
net_packet.set_destination(source);
net_packet.first_set_ttl(1);
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.send_by_key(net_packet.buffer(), route_key)?;
let route = Route::from_default_rt(route_key, 1);
context.route_table.add_route_if_absent(source, route);
}
ControlPacket::PunchResponse => {
log::info!("PunchResponse={:?},source={}", route_key, source);
if context.use_channel_type().is_only_relay() {
return Ok(());
}
let route = Route::from_default_rt(route_key, 1);
context.route_table.add_route_if_absent(source, route);
}
ControlPacket::AddrRequest => match route_key.addr.ip() {
std::net::IpAddr::V4(ipv4) => {
let mut packet = NetPacket::new_encrypt([0; 12 + 6 + ENCRYPTION_RESERVED])?;
packet.set_version(Version::V1);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::AddrResponse.into());
packet.first_set_ttl(MAX_TTL);
packet.set_source(current_device.virtual_ip);
packet.set_destination(source);
let mut addr_packet = control_packet::AddrPacket::new(packet.payload_mut())?;
addr_packet.set_ipv4(ipv4);
addr_packet.set_port(route_key.addr.port());
self.client_cipher.encrypt_ipv4(&mut packet)?;
context.send_by_key(packet.buffer(), route_key)?;
}
std::net::IpAddr::V6(_) => {}
},
ControlPacket::AddrResponse(_) => {}
}
Ok(())
}
fn other_turn(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
if 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 public_ips = punch_info
.public_ip_list
.iter()
.map(|v| Ipv4Addr::from(v.to_be_bytes()))
.collect();
let local_ipv4 = Some(Ipv4Addr::from(punch_info.local_ip.to_be_bytes()));
let tcp_port = punch_info.tcp_port as u16;
let ipv6 = if punch_info.ipv6.len() == 16 {
let ipv6: [u8; 16] = punch_info.ipv6.try_into().unwrap();
Some(Ipv6Addr::from(ipv6))
} else {
None
};
//兼容旧版本
if punch_info.public_ports.is_empty() {
punch_info.public_ports.push(punch_info.public_port);
}
//兼容旧版本
if punch_info.udp_ports.is_empty() {
punch_info.udp_ports.push(punch_info.local_port);
}
let peer_nat_info = NatInfo::new(
public_ips,
punch_info.public_ports.iter().map(|e| *e as u16).collect(),
punch_info.public_port_range as u16,
local_ipv4,
ipv6,
punch_info.udp_ports.iter().map(|e| *e as u16).collect(),
tcp_port,
punch_info.nat_type.enum_value_or_default().into(),
);
{
let peer_nat_info = peer_nat_info.clone();
self.peer_nat_info_map.write().insert(source, peer_nat_info);
}
if !punch_info.reply {
let mut punch_reply = PunchInfo::new();
punch_reply.reply = true;
let nat_info = self.nat_test.nat_info();
punch_reply.public_ip_list = nat_info
.public_ips
.iter()
.map(|ip| u32::from_be_bytes(ip.octets()))
.collect();
punch_reply.public_port = nat_info.public_ports.get(0).map_or(0, |v| *v as u32);
punch_reply.public_ports =
nat_info.public_ports.iter().map(|e| *e as u32).collect();
punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.tcp_port = nat_info.tcp_port as u32;
punch_reply.nat_type =
protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
punch_reply.local_ip =
u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED));
punch_reply.local_port = nat_info.udp_ports[0] as u32;
punch_reply.udp_ports = nat_info.udp_ports.iter().map(|e| *e as u32).collect();
if let Some(ipv6) = nat_info.ipv6() {
punch_reply.ipv6 = ipv6.octets().to_vec();
punch_reply.ipv6_port = nat_info.udp_ports[0] as u32;
}
let bytes = punch_reply.write_to_bytes().map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("punch_reply {:?}", e))
})?;
let mut punch_packet =
NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
punch_packet.set_version(Version::V1);
punch_packet.set_protocol(Protocol::OtherTurn);
punch_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into());
punch_packet.first_set_ttl(MAX_TTL);
punch_packet.set_source(current_device.virtual_ip());
punch_packet.set_destination(source);
punch_packet.set_payload(&bytes)?;
self.client_cipher.encrypt_ipv4(&mut punch_packet)?;
if self.punch_sender.send(true, source, peer_nat_info) {
context.send_by_key(punch_packet.buffer(), route_key)?;
}
} else {
self.punch_sender.send(false, source, peer_nat_info);
}
}
other_turn_packet::Protocol::Unknown(e) => {
log::warn!("不支持的转发协议 {:?},source:{:?}", e, source);
}
}
Ok(())
}
}
+151
View File
@@ -0,0 +1,151 @@
use std::collections::HashMap;
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::{Mutex, RwLock};
use tun::Device;
use crate::channel::context::Context;
use crate::channel::handler::RecvChannelHandler;
use crate::channel::punch::NatInfo;
use crate::channel::RouteKey;
use crate::cipher::Cipher;
#[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher;
use crate::external_route::{AllowExternalRoute, ExternalRoute};
use crate::handle::callback::VntCallback;
use crate::handle::handshaker::Handshake;
use crate::handle::maintain::PunchSender;
use crate::handle::recv_data::client::ClientPacketHandler;
use crate::handle::recv_data::server::ServerPacketHandler;
use crate::handle::recv_data::turn::TurnPacketHandler;
use crate::handle::{BaseConfigInfo, CurrentDeviceInfo, PeerDeviceInfo, SELF_IP};
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::IpProxyMap;
use crate::nat::NatTest;
use crate::protocol::NetPacket;
use crate::util::U64Adder;
mod client;
mod server;
mod turn;
#[derive(Clone)]
pub struct RecvDataHandler<Call> {
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
turn: TurnPacketHandler,
client: ClientPacketHandler,
server: ServerPacketHandler<Call>,
counter: U64Adder,
}
impl<Call: VntCallback> RecvChannelHandler for RecvDataHandler<Call> {
fn handle(&mut self, buf: &mut [u8], route_key: RouteKey, context: &Context) {
if let Err(e) = self.handle0(buf, route_key, context) {
log::error!("[{}]-{:?}", thread::current().name().unwrap_or(""), e);
}
}
}
impl<Call: VntCallback> RecvDataHandler<Call> {
pub fn new(
#[cfg(feature = "server_encrypt")] rsa_cipher: Arc<Mutex<Option<RsaCipher>>>,
server_cipher: Cipher,
client_cipher: Cipher,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device: Arc<Device>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
config_info: BaseConfigInfo,
nat_test: NatTest,
callback: Call,
punch_sender: PunchSender,
peer_nat_info_map: Arc<RwLock<HashMap<Ipv4Addr, NatInfo>>>,
external_route: ExternalRoute,
route: AllowExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
counter: U64Adder,
handshake: Handshake,
) -> Self {
let server = ServerPacketHandler::new(
#[cfg(feature = "server_encrypt")]
rsa_cipher,
server_cipher,
current_device.clone(),
device.clone(),
device_list,
config_info,
nat_test.clone(),
callback,
external_route,
handshake,
);
let client = ClientPacketHandler::new(
device.clone(),
client_cipher,
punch_sender,
peer_nat_info_map,
nat_test,
route,
#[cfg(feature = "ip_proxy")]
ip_proxy_map,
);
let turn = TurnPacketHandler::new();
Self {
current_device,
turn,
client,
server,
counter,
}
}
fn handle0(
&mut self,
buf: &mut [u8],
route_key: RouteKey,
context: &Context,
) -> io::Result<()> {
// 统计流量
self.counter.add(buf.len() as _);
let net_packet = NetPacket::new(buf)?;
if net_packet.ttl() == 0 || net_packet.source_ttl() < net_packet.ttl() {
return Ok(());
}
let current_device = self.current_device.load();
let dest = net_packet.destination();
if dest == current_device.virtual_ip
|| dest.is_broadcast()
|| dest.is_multicast()
|| dest == SELF_IP
|| dest.is_unspecified()
|| dest == current_device.broadcast_ip
{
//发给自己的包
if net_packet.is_gateway() {
//服务端-客户端包
self.server
.handle(net_packet, route_key, context, &current_device)
} else {
//客户端-客户端包
self.client
.handle(net_packet, route_key, context, &current_device)
}
} else {
//转发包
self.turn
.handle(net_packet, route_key, context, &current_device)
}
}
}
pub trait PacketHandler {
fn handle(
&self,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
context: &Context,
current_device: &CurrentDeviceInfo,
) -> io::Result<()>;
}
+483
View File
@@ -0,0 +1,483 @@
use std::io;
use std::net::Ipv4Addr;
use std::sync::Arc;
#[cfg(feature = "server_encrypt")]
use std::time::{Duration, Instant};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use protobuf::Message;
use packet::icmp::{icmp, Kind};
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use tun::device::IFace;
use tun::Device;
use crate::channel::context::Context;
use crate::channel::{Route, RouteKey};
use crate::cipher::Cipher;
#[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher;
use crate::external_route::ExternalRoute;
use crate::handle::callback::{ErrorInfo, ErrorType, HandshakeInfo, RegisterInfo, VntCallback};
#[cfg(feature = "server_encrypt")]
use crate::handle::handshaker;
use crate::handle::handshaker::Handshake;
use crate::handle::recv_data::PacketHandler;
use crate::handle::{
registrar, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo, GATEWAY_IP,
};
use crate::nat::NatTest;
use crate::proto;
use crate::proto::message::{DeviceList, HandshakeResponse, RegistrationResponse};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::ControlPacket;
use crate::protocol::error_packet::InErrorPacket;
use crate::protocol::{ip_turn_packet, service_packet, NetPacket, Protocol, Version, MAX_TTL};
/// 处理来源于服务端的包
#[derive(Clone)]
pub struct ServerPacketHandler<Call> {
#[cfg(feature = "server_encrypt")]
rsa_cipher: Arc<Mutex<Option<RsaCipher>>>,
server_cipher: Cipher,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device: Arc<Device>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
config_info: BaseConfigInfo,
nat_test: NatTest,
callback: Call,
#[cfg(feature = "server_encrypt")]
up_key_time: Arc<AtomicCell<Instant>>,
route_record: Arc<Mutex<Vec<(Ipv4Addr, Ipv4Addr)>>>,
external_route: ExternalRoute,
handshake: Handshake,
}
impl<Call> ServerPacketHandler<Call> {
pub fn new(
#[cfg(feature = "server_encrypt")] rsa_cipher: Arc<Mutex<Option<RsaCipher>>>,
server_cipher: Cipher,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device: Arc<Device>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
config_info: BaseConfigInfo,
nat_test: NatTest,
callback: Call,
external_route: ExternalRoute,
handshake: Handshake,
) -> Self {
Self {
#[cfg(feature = "server_encrypt")]
rsa_cipher,
server_cipher,
current_device,
device,
device_list,
config_info,
nat_test,
callback,
#[cfg(feature = "server_encrypt")]
up_key_time: Arc::new(AtomicCell::new(Instant::now() - Duration::from_secs(60))),
route_record: Arc::new(Mutex::default()),
external_route,
handshake,
}
}
}
impl<Call: VntCallback> PacketHandler for ServerPacketHandler<Call> {
fn handle(
&self,
mut net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
context: &Context,
current_device: &CurrentDeviceInfo,
) -> io::Result<()> {
context
.route_table
.update_read_time(&net_packet.source(), &route_key);
if net_packet.protocol() == Protocol::Error
&& net_packet.transport_protocol()
== crate::protocol::error_packet::Protocol::NoKey.into()
{
//服务端通知客户端上传密钥
#[cfg(feature = "server_encrypt")]
{
let mutex_guard = self.rsa_cipher.lock();
if let Some(rsa_cipher) = mutex_guard.as_ref() {
let last = self.up_key_time.load();
if last.elapsed() < Duration::from_secs(1)
|| self
.up_key_time
.compare_exchange(last, Instant::now())
.is_err()
{
//短时间不重复上传服务端密钥
return Ok(());
}
if let Some(key) = self.server_cipher.key() {
log::info!("上传密钥到服务端:{:?}", route_key);
let packet = handshaker::secret_handshake_request_packet(
rsa_cipher,
self.config_info.token.clone(),
key,
)?;
context.send_by_key(packet.buffer(), route_key)?;
}
}
}
return Ok(());
} else if net_packet.protocol() == Protocol::Service
&& net_packet.transport_protocol() == service_packet::Protocol::HandshakeResponse.into()
{
let response =
HandshakeResponse::parse_from_bytes(net_packet.payload()).map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("HandshakeResponse {:?}", e))
})?;
//如果开启了加密,则发送加密握手请求
#[cfg(feature = "server_encrypt")]
if let Some(key) = self.server_cipher.key() {
let rsa_cipher = RsaCipher::new(&response.public_key)?;
let handshake_info = HandshakeInfo::new(
rsa_cipher.public_key()?.clone(),
rsa_cipher.finger()?,
response.version,
);
log::info!("加密握手请求:{:?}", handshake_info);
if self.callback.handshake(handshake_info) {
let packet = handshaker::secret_handshake_request_packet(
&rsa_cipher,
self.config_info.token.clone(),
key,
)?;
context.send_by_key(packet.buffer(), route_key)?;
self.rsa_cipher.lock().replace(rsa_cipher);
}
return Ok(());
}
let handshake_info = HandshakeInfo::new_no_secret(response.version);
if self.callback.handshake(handshake_info) {
//没有加密,则发送注册请求
self.register(current_device, context)?;
}
return Ok(());
}
//服务端数据解密
self.server_cipher.decrypt_ipv4(&mut net_packet)?;
match net_packet.protocol() {
Protocol::Service => {
self.service(context, current_device, net_packet, route_key)?;
}
Protocol::Error => {
self.error(context, current_device, net_packet, route_key)?;
}
Protocol::Control => {
self.control(context, current_device, net_packet, route_key)?;
}
Protocol::IpTurn => {
match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) {
ip_turn_packet::Protocol::Ipv4 => {
let ipv4 = IpV4Packet::new(net_packet.payload())?;
match ipv4.protocol() {
ipv4::protocol::Protocol::Icmp => {
if ipv4.destination_ip() == current_device.virtual_ip {
let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?;
if icmp_packet.kind() == Kind::EchoReply {
//网关ip ping的回应
self.device.write(net_packet.payload())?;
return Ok(());
}
}
}
_ => {}
}
}
ip_turn_packet::Protocol::Ipv4Broadcast => {}
ip_turn_packet::Protocol::Unknown(_) => {}
}
}
Protocol::OtherTurn => {}
Protocol::Unknown(_) => {}
}
Ok(())
}
}
impl<Call: VntCallback> ServerPacketHandler<Call> {
fn service(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
match service_packet::Protocol::from(net_packet.transport_protocol()) {
service_packet::Protocol::RegistrationResponse => {
let response = RegistrationResponse::parse_from_bytes(net_packet.payload())
.map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("RegistrationResponse {:?}", e),
)
})?;
let virtual_ip = Ipv4Addr::from(response.virtual_ip);
let virtual_netmask = Ipv4Addr::from(response.virtual_netmask);
let virtual_gateway = Ipv4Addr::from(response.virtual_gateway);
let virtual_network =
Ipv4Addr::from(response.virtual_ip & response.virtual_netmask);
let register_info = RegisterInfo::new(virtual_ip, virtual_netmask, virtual_gateway);
log::info!("注册成功:{:?}", register_info);
if self.callback.register(register_info) {
let route = Route::from_default_rt(route_key, 1);
context
.route_table
.add_route_if_absent(virtual_gateway, route);
let old = current_device;
let mut cur = *current_device;
loop {
let mut new_current_device = cur;
new_current_device.update(virtual_ip, virtual_netmask, virtual_gateway);
new_current_device.virtual_ip = virtual_ip;
new_current_device.virtual_netmask = virtual_netmask;
new_current_device.virtual_gateway = virtual_gateway;
new_current_device.status = crate::handle::ConnectStatus::Connected;
if let Err(c) = self
.current_device
.compare_exchange(cur, new_current_device)
{
cur = c;
} else {
break;
}
}
let public_ip = response.public_ip.into();
let public_port = response.public_port as u16;
self.nat_test
.update_addr(route_key.index(), public_ip, public_port);
if old.virtual_ip != virtual_ip
|| old.virtual_gateway != virtual_gateway
|| old.virtual_netmask != virtual_netmask
{
if old.virtual_ip != Ipv4Addr::UNSPECIFIED {
log::info!("ip发生变化,old:{:?},response={:?}", old, response);
}
if let Err(e) = self.device.set_ip(virtual_ip, virtual_netmask) {
log::error!("LocalIpExists {:?}", e);
self.callback.error(ErrorInfo::new_msg(
ErrorType::LocalIpExists,
format!("set_ip {:?}", e),
));
return Ok(());
}
let mut guard = self.route_record.lock();
for (dest, mask) in guard.drain(..) {
if let Err(e) = self.device.delete_route(dest, mask) {
log::warn!("删除路由失败 ={:?}", e);
}
}
if let Err(e) = self.device.add_route(virtual_network, virtual_netmask, 1) {
log::warn!("添加默认路由失败 ={:?}", e);
} else {
guard.push((virtual_network, virtual_netmask));
}
if let Err(e) =
self.device
.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, 1)
{
log::warn!("添加广播路由失败 ={:?}", e);
} else {
guard.push((Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST));
}
if let Err(e) = self.device.add_route(
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
1,
) {
log::warn!("添加组播路由失败 ={:?}", e);
} else {
guard.push((
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
));
}
for (dest, mask) in self.external_route.to_route() {
if let Err(e) = self.device.add_route(dest, mask, 1) {
log::warn!("添加路由失败 ={:?}", e);
} else {
guard.push((dest, mask));
}
}
self.callback.success();
}
self.set_device_info_list(response.device_info_list, response.epoch as _);
}
}
service_packet::Protocol::PushDeviceList => {
let response = DeviceList::parse_from_bytes(net_packet.payload()).map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("PushDeviceList {:?}", e))
})?;
self.set_device_info_list(response.device_info_list, response.epoch as _);
}
service_packet::Protocol::SecretHandshakeResponse => {
log::info!("SecretHandshakeResponse");
//加密握手结束,发送注册数据
self.register(current_device, context)?;
}
_ => {
log::warn!(
"service_packet::Protocol::Unknown = {:?}",
net_packet.head()
);
}
}
Ok(())
}
fn set_device_info_list(&self, device_info_list: Vec<proto::message::DeviceInfo>, epoch: u16) {
let ip_list: Vec<PeerDeviceInfo> = device_info_list
.into_iter()
.map(|info| {
PeerDeviceInfo::new(
Ipv4Addr::from(info.virtual_ip),
info.name,
info.device_status as u8,
info.client_secret,
)
})
.collect();
let mut dev = self.device_list.lock();
//这里可能会收到旧的消息,但是随着时间推移总会收到新的
dev.0 = epoch;
dev.1 = ip_list;
}
fn register(&self, current_device: &CurrentDeviceInfo, context: &Context) -> io::Result<()> {
if current_device.status.online() {
//已连接的不需要注册
return Ok(());
}
let token = self.config_info.token.clone();
let device_id = self.config_info.device_id.clone();
let name = self.config_info.name.clone();
let client_secret = self.config_info.client_secret;
let mut ip = self.config_info.ip;
if ip.is_none() {
ip = Some(current_device.virtual_ip)
}
let response = registrar::registration_request_packet(
&self.server_cipher,
token,
device_id,
name,
ip,
false,
false,
client_secret,
)?;
log::info!("发送注册请求,{:?}", self.config_info);
//注册请求只发送到默认通道
context.send_default(response.buffer(), current_device.connect_server)
}
fn error(
&self,
context: &Context,
_current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
InErrorPacket::TokenError => {
// token错误,可能是服务端设置了白名单
let err = ErrorInfo::new(ErrorType::TokenError);
self.callback.error(err);
}
InErrorPacket::Disconnect => {
crate::handle::change_status(&self.current_device, ConnectStatus::Connecting);
let err = ErrorInfo::new(ErrorType::Disconnect);
self.callback.error(err);
//掉线epoch要归零
{
let mut dev = self.device_list.lock();
dev.0 = 0;
drop(dev);
}
self.handshake
.send(context, self.config_info.client_secret, route_key.addr)?;
// self.register(current_device, context, route_key)?;
}
InErrorPacket::AddressExhausted => {
// 地址用尽
let err = ErrorInfo::new(ErrorType::AddressExhausted);
self.callback.error(err);
}
InErrorPacket::OtherError(e) => {
let err = ErrorInfo::new_msg(ErrorType::Unknown, e.message()?);
self.callback.error(err);
}
InErrorPacket::IpAlreadyExists => {
let err = ErrorInfo::new(ErrorType::IpAlreadyExists);
self.callback.error(err);
}
InErrorPacket::InvalidIp => {
let err = ErrorInfo::new(ErrorType::InvalidIp);
self.callback.error(err);
}
InErrorPacket::NoKey => {
//这个类型最开头已经处理过,这里忽略
}
}
Ok(())
}
fn control(
&self,
context: &Context,
current_device: &CurrentDeviceInfo,
net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
) -> io::Result<()> {
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PongPacket(pong_packet) => {
let current_time = crate::handle::now_time() as u16;
if current_time < pong_packet.time() {
return Ok(());
}
let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
let rt = (current_time - pong_packet.time()) as i64;
let route = Route::from(route_key, metric, rt);
context.route_table.add_route(net_packet.source(), route);
let epoch = self.device_list.lock().0;
if pong_packet.epoch() != epoch {
//纪元不一致,可能有新客户端连接,向服务端拉取客户端列表
let mut poll_device = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED])?;
poll_device.set_source(current_device.virtual_ip);
poll_device.set_destination(GATEWAY_IP);
poll_device.set_version(Version::V1);
poll_device.set_gateway_flag(true);
poll_device.first_set_ttl(MAX_TTL);
poll_device.set_protocol(Protocol::Service);
poll_device
.set_transport_protocol(service_packet::Protocol::PollDeviceList.into());
self.server_cipher.encrypt_ipv4(&mut poll_device)?;
//发送到默认服务端即可
context.send_default(poll_device.buffer(), current_device.connect_server)?;
}
}
ControlPacket::AddrResponse(addr_packet) => {
//更新本地公网ipv4
self.nat_test.update_addr(
route_key.index(),
addr_packet.ipv4(),
addr_packet.port(),
);
}
_ => {}
}
Ok(())
}
}
+44
View File
@@ -0,0 +1,44 @@
use crate::channel::context::Context;
use crate::channel::RouteKey;
use crate::handle::recv_data::PacketHandler;
use crate::handle::CurrentDeviceInfo;
use crate::protocol::NetPacket;
/// 处理客户端中转包
#[derive(Clone)]
pub struct TurnPacketHandler {}
impl TurnPacketHandler {
pub fn new() -> Self {
Self {}
}
}
impl PacketHandler for TurnPacketHandler {
fn handle(
&self,
mut net_packet: NetPacket<&mut [u8]>,
route_key: RouteKey,
context: &Context,
_current_device: &CurrentDeviceInfo,
) -> std::io::Result<()> {
// ttl减一
let ttl = net_packet.incr_ttl();
if ttl > 0 {
let destination = net_packet.destination();
if let Some(route) = context.route_table.route_one(&destination) {
if route.addr == route_key.addr {
//防止环路
log::warn!("来源和目标相同 {:?},{:?}", route_key, net_packet.head());
return Ok(());
}
if route.metric <= ttl {
context.send_by_key(net_packet.buffer(), route.route_key())?;
}
}
//其他没有路由的不转发
}
Ok(())
}
}
-789
View File
@@ -1,789 +0,0 @@
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6};
use std::sync::Arc;
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use parking_lot::Mutex;
use protobuf::Message;
use tokio::sync::mpsc::Sender;
use packet::icmp::{icmp, Kind};
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::channel::Context;
use crate::channel::punch::{NatInfo, NatType};
use crate::channel::{Route, RouteKey};
use crate::cipher::{Cipher, RsaCipher};
use crate::error::Error;
use crate::external_route::AllowExternalRoute;
use crate::handle::handshake_handler::secret_handshake_req;
use crate::handle::registration_handler::Register;
use crate::handle::{ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo, PeerDeviceStatus};
use crate::igmp_server::IgmpServer;
use crate::ip_proxy::IpProxyMap;
use crate::nat;
use crate::nat::NatTest;
use crate::proto::message::{DeviceList, PunchInfo, PunchNatType, RegistrationResponse};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::ControlPacket;
use crate::protocol::error_packet::InErrorPacket;
use crate::protocol::{
control_packet, ip_turn_packet, other_turn_packet, service_packet, NetPacket, Protocol,
Version, MAX_TTL,
};
use crate::tun_tap_device::DeviceWriter;
#[derive(Clone)]
pub struct ChannelDataHandler {
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
register: Arc<Register>,
nat_test: NatTest,
igmp_server: Option<IgmpServer>,
device_writer: DeviceWriter,
connect_status: Arc<AtomicCell<ConnectStatus>>,
peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>>,
ip_proxy_map: Option<IpProxyMap>,
out_external_route: AllowExternalRoute,
cone_sender: Sender<(Ipv4Addr, NatInfo)>,
symmetric_sender: Sender<(Ipv4Addr, NatInfo)>,
client_cipher: Cipher,
server_cipher: Cipher,
rsa_cipher: Option<RsaCipher>,
relay: bool,
token: String,
}
impl ChannelDataHandler {
pub fn new(
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
device_list: Arc<Mutex<(u16, Vec<PeerDeviceInfo>)>>,
register: Arc<Register>,
nat_test: NatTest,
igmp_server: Option<IgmpServer>,
device_writer: DeviceWriter,
connect_status: Arc<AtomicCell<ConnectStatus>>,
peer_nat_info_map: Arc<DashMap<Ipv4Addr, NatInfo>>,
ip_proxy_map: Option<IpProxyMap>,
out_external_route: AllowExternalRoute,
cone_sender: Sender<(Ipv4Addr, NatInfo)>,
symmetric_sender: Sender<(Ipv4Addr, NatInfo)>,
client_cipher: Cipher,
server_cipher: Cipher,
rsa_cipher: Option<RsaCipher>,
relay: bool,
token: String,
) -> Self {
Self {
current_device,
device_list,
register,
nat_test,
igmp_server,
device_writer,
connect_status,
peer_nat_info_map,
ip_proxy_map,
out_external_route,
cone_sender,
symmetric_sender,
client_cipher,
server_cipher,
rsa_cipher,
relay,
token,
}
}
}
impl ChannelDataHandler {
pub async fn handle(
&self,
buf: &mut [u8],
start: usize,
end: usize,
route_key: RouteKey,
context: &Context,
) {
assert_eq!(start, 14);
match self.handle0(&mut buf[..end], &route_key, context).await {
Ok(_) => {}
Err(e) => {
log::warn!("{:?}", e);
}
}
}
async fn handle0(
&self,
buf: &mut [u8],
route_key: &RouteKey,
context: &Context,
) -> crate::Result<()> {
let mut net_packet = NetPacket::new(&mut buf[14..])?;
if net_packet.ttl() == 0 || net_packet.source_ttl() < net_packet.ttl() {
return Ok(());
}
let source = net_packet.source();
context.update_read_time(&source, route_key);
let current_device = self.current_device.load();
let destination = net_packet.destination();
let not_broadcast = !destination.is_broadcast()
&& !destination.is_multicast()
&& destination != current_device.broadcast_address;
if current_device.virtual_ip() != destination
&& not_broadcast
&& !destination.is_unspecified()
{
//校验指纹,不需要解密
self.client_cipher.check_finger(&net_packet)?;
net_packet.set_ttl(net_packet.ttl() - 1);
let ttl = net_packet.ttl();
if ttl > 0 {
// 转发
if let Some(route) = context.route_one(&destination) {
if route.metric <= net_packet.ttl() {
context.try_send_by_key(net_packet.buffer(), &route.route_key())?;
}
} else if (ttl > 1 || destination == current_device.virtual_gateway())
&& source != current_device.virtual_gateway()
{
//网关默认要转发一次,生存时间不够的发到网关也会被丢弃
context.send_main(net_packet.buffer(), current_device.connect_server)?;
}
}
return Ok(());
}
if net_packet.is_gateway() {
if net_packet.protocol() == Protocol::Error
&& net_packet.transport_protocol()
== crate::protocol::error_packet::Protocol::NoKey.into()
{
if let Some(rsa_cipher) = &self.rsa_cipher {
secret_handshake_req(
context,
current_device.connect_server,
rsa_cipher,
&self.server_cipher,
self.token.clone(),
route_key,
)?;
}
} else {
//服务端解密
self.server_cipher.decrypt_ipv4(&mut net_packet)?;
let data_len = net_packet.data_len();
self.server_packet_handle(context, current_device, buf, data_len, route_key)
.await?;
}
return Ok(());
}
self.client_cipher.decrypt_ipv4(&mut net_packet)?;
match net_packet.protocol() {
Protocol::IpTurn => {
match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) {
ip_turn_packet::Protocol::Ipv4 => {
let mut ipv4 = IpV4Packet::new(net_packet.payload_mut())?;
match ipv4.protocol() {
ipv4::protocol::Protocol::Igmp => {
if let Some(igmp_server) = &self.igmp_server {
igmp_server.handle(ipv4.payload(), source)?;
}
return Ok(());
}
ipv4::protocol::Protocol::Icmp => {
if ipv4.destination_ip() == destination {
let mut icmp_packet =
icmp::IcmpPacket::new(ipv4.payload_mut())?;
if icmp_packet.kind() == Kind::EchoRequest {
//开启ping
icmp_packet.set_kind(Kind::EchoReply);
icmp_packet.update_checksum();
ipv4.set_source_ip(destination);
ipv4.set_destination_ip(source);
ipv4.update_checksum();
net_packet.set_source(destination);
net_packet.set_destination(source);
//不管加不加密,和接收到的数据长度都一致
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.try_send_by_key(net_packet.buffer(), route_key)?;
return Ok(());
}
}
}
_ => {}
}
if not_broadcast && ipv4.destination_ip() != destination {
if let Some(ip_proxy_map) = &self.ip_proxy_map {
if self.out_external_route.allow(&ipv4.destination_ip()) {
match ipv4.protocol() {
ipv4::protocol::Protocol::Tcp => {
let dest_ip = ipv4.destination_ip();
//转发到代理目标地址
let mut tcp_packet = packet::tcp::tcp::TcpPacket::new(
source,
destination,
ipv4.payload_mut(),
)?;
let source_port = tcp_packet.source_port();
let dest_port = tcp_packet.destination_port();
tcp_packet
.set_destination_port(ip_proxy_map.tcp_proxy_port);
tcp_packet.update_checksum();
ipv4.set_destination_ip(destination);
ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port);
//https://github.com/crossbeam-rs/crossbeam/issues/1023
ip_proxy_map
.tcp_proxy_map
.insert(key, SocketAddrV4::new(dest_ip, dest_port));
}
ipv4::protocol::Protocol::Udp => {
let dest_ip = ipv4.destination_ip();
//转发到代理目标地址
let mut udp_packet = packet::udp::udp::UdpPacket::new(
source,
destination,
ipv4.payload_mut(),
)?;
let source_port = udp_packet.source_port();
let dest_port = udp_packet.destination_port();
udp_packet
.set_destination_port(ip_proxy_map.udp_proxy_port);
udp_packet.update_checksum();
ipv4.set_destination_ip(destination);
ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port);
ip_proxy_map
.udp_proxy_map
.insert(key, SocketAddrV4::new(dest_ip, dest_port));
}
#[cfg(not(target_os = "android"))]
ipv4::protocol::Protocol::Icmp => {
let dest_ip = ipv4.destination_ip();
//转发到代理目标地址
let icmp_packet =
icmp::IcmpPacket::new(ipv4.payload())?;
match icmp_packet.header_other() {
icmp::HeaderOther::Identifier(id, seq) => {
ip_proxy_map
.icmp_proxy_map
.insert((dest_ip, id, seq), source);
ip_proxy_map
.send_icmp(ipv4.payload(), &dest_ip)?;
}
_ => {
log::warn!(
"不支持的ip代理Icmp协议:{}",
destination
);
return Err(Error::Warn(
"不支持的ip代理Icmp协议".to_string(),
));
}
}
}
_ => {
log::warn!("不支持的ip代理ipv4协议:{}", destination);
return Err(Error::Warn(
"不支持的ip代理ipv4协议".to_string(),
));
}
}
} else {
log::warn!("没有ip代理规则:{}", destination);
return Err(Error::Warn("没有ip代理规则".to_string()));
}
} else {
log::warn!("不支持ip代理:{}", destination);
return Err(Error::Warn("不支持ip代理".to_string()));
}
}
//传输协议12字节
self.device_writer.write_ipv4(&mut buf[12..])?;
return Ok(());
}
ip_turn_packet::Protocol::Ipv4Broadcast => {
//客户端不帮忙转发广播包,所以不会出现这种类型的数据
}
ip_turn_packet::Protocol::Unknown(_) => {}
}
}
Protocol::Service => {}
Protocol::Error => {}
Protocol::Control => {
self.control(context, current_device, source, net_packet, route_key)
.await?;
}
Protocol::OtherTurn => {
self.other_turn(context, current_device, source, net_packet, route_key)
.await?;
}
Protocol::UnKnow(e) => {
log::info!("不支持的协议:{}", e);
}
}
Ok(())
}
async fn pong_packet(
&self,
gateway: bool,
metric: u8,
context: &Context,
current_device: CurrentDeviceInfo,
source: Ipv4Addr,
pong_packet: control_packet::PongPacket<&[u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
let current_time = crate::handle::now_time() as u16;
if current_time < pong_packet.time() {
return Ok(());
}
let rt = (current_time - pong_packet.time()) as i64;
let route = Route::from(*route_key, metric, rt);
context.add_route(source, route);
if gateway {
let epoch = self.device_list.lock().0;
if pong_packet.epoch() != epoch {
let mut poll_device = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED])?;
poll_device.set_source(current_device.virtual_ip());
poll_device.set_destination(source);
poll_device.set_version(Version::V1);
poll_device.set_gateway_flag(true);
poll_device.first_set_ttl(MAX_TTL);
poll_device.set_protocol(Protocol::Service);
poll_device.set_transport_protocol(service_packet::Protocol::PollDeviceList.into());
self.server_cipher.encrypt_ipv4(&mut poll_device)?;
context.send_main(poll_device.buffer(), current_device.connect_server)?;
}
}
Ok(())
}
async fn control(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
source: Ipv4Addr,
mut net_packet: NetPacket<&mut [u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PingPacket(_) => {
net_packet.set_transport_protocol(control_packet::Protocol::Pong.into());
net_packet.set_source(current_device.virtual_ip());
net_packet.set_destination(source);
net_packet.first_set_ttl(MAX_TTL);
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.try_send_by_key(net_packet.buffer(), route_key)?;
let route = Route::from(*route_key, metric, 199);
context.add_route_if_absent(source, route);
}
ControlPacket::PongPacket(pong_packet) => {
self.pong_packet(
false,
metric,
context,
current_device,
source,
pong_packet,
route_key,
)
.await?;
}
ControlPacket::PunchRequest => {
if self.relay {
return Ok(());
}
//回应
net_packet.set_transport_protocol(control_packet::Protocol::PunchResponse.into());
net_packet.set_source(current_device.virtual_ip());
net_packet.set_destination(source);
net_packet.first_set_ttl(1);
self.client_cipher.encrypt_ipv4(&mut net_packet)?;
context.try_send_by_key(net_packet.buffer(), route_key)?;
let route = Route::from(*route_key, 1, 199);
context.add_route_if_absent(source, route);
}
ControlPacket::PunchResponse => {
if self.relay {
return Ok(());
}
let route = Route::from(*route_key, 1, 199);
context.add_route_if_absent(source, route);
}
ControlPacket::AddrRequest => match route_key.addr.ip() {
std::net::IpAddr::V4(ipv4) => {
let mut packet = NetPacket::new_encrypt([0; 12 + 6 + ENCRYPTION_RESERVED])?;
packet.set_version(Version::V1);
packet.set_protocol(Protocol::Control);
packet.set_transport_protocol(control_packet::Protocol::AddrResponse.into());
packet.first_set_ttl(MAX_TTL);
packet.set_source(current_device.virtual_ip());
packet.set_destination(source);
let mut addr_packet = control_packet::AddrPacket::new(packet.payload_mut())?;
addr_packet.set_ipv4(ipv4);
addr_packet.set_port(route_key.addr.port());
self.client_cipher.encrypt_ipv4(&mut packet)?;
context.try_send_by_key(packet.buffer(), route_key)?;
}
std::net::IpAddr::V6(_) => {}
},
ControlPacket::AddrResponse(addr_packet) => {
if !addr_packet.ipv4().is_multicast()
&& !addr_packet.ipv4().is_broadcast()
&& !addr_packet.ipv4().is_unspecified()
&& !addr_packet.ipv4().is_loopback()
&& !addr_packet.ipv4().is_private()
&& addr_packet.port() != 0
{
self.nat_test
.update_addr(addr_packet.ipv4(), addr_packet.port())
}
}
}
Ok(())
}
async fn other_turn(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
source: Ipv4Addr,
net_packet: NetPacket<&mut [u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
if self.relay {
return Ok(());
}
match other_turn_packet::Protocol::from(net_packet.transport_protocol()) {
other_turn_packet::Protocol::Punch => {
let punch_info = PunchInfo::parse_from_bytes(net_packet.payload())?;
let public_ips = punch_info
.public_ip_list
.iter()
.map(|v| Ipv4Addr::from(v.to_be_bytes()))
.collect();
let local_ipv4_addr = SocketAddrV4::new(
Ipv4Addr::from(punch_info.local_ip.to_be_bytes()),
punch_info.local_port as u16,
);
let ipv6_addr = if punch_info.ipv6.len() == 16 {
let ipv6: [u8; 16] = punch_info.ipv6.try_into().unwrap();
SocketAddrV6::new(Ipv6Addr::from(ipv6), punch_info.ipv6_port as u16, 0, 0)
} else {
SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, 0, 0, 0)
};
let peer_nat_info = NatInfo::new(
public_ips,
punch_info.public_port as u16,
punch_info.public_port_range as u16,
local_ipv4_addr,
ipv6_addr,
punch_info.nat_type.enum_value_or_default().into(),
);
self.peer_nat_info_map.insert(source, peer_nat_info.clone());
if !punch_info.reply {
let mut punch_reply = PunchInfo::new();
punch_reply.reply = true;
let nat_info = self.nat_test.nat_info();
punch_reply.public_ip_list = nat_info
.public_ips
.iter()
.map(|ip| u32::from_be_bytes(ip.octets()))
.collect();
punch_reply.public_port = nat_info.public_port as u32;
punch_reply.public_port_range = nat_info.public_port_range as u32;
punch_reply.nat_type =
protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type));
punch_reply.local_ip =
u32::from_be_bytes(nat_info.local_ipv4_addr.ip().octets());
punch_reply.local_port = nat_info.local_ipv4_addr.port() as u32;
if !nat_info.ipv6_addr.ip().is_unspecified() {
punch_reply.ipv6 = nat_info.ipv6_addr.ip().octets().to_vec();
punch_reply.ipv6_port = nat_info.ipv6_addr.port() as u32;
}
let bytes = punch_reply.write_to_bytes()?;
let mut punch_packet =
NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?;
punch_packet.set_version(Version::V1);
punch_packet.set_protocol(Protocol::OtherTurn);
punch_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into());
punch_packet.first_set_ttl(MAX_TTL);
punch_packet.set_source(current_device.virtual_ip());
punch_packet.set_destination(source);
punch_packet.set_payload(&bytes)?;
// if !peer_nat_info.local_ip.is_unspecified() && peer_nat_info.local_port != 0 {
// let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED])?;
// packet.set_version(Version::V1);
// packet.first_set_ttl(1);
// packet.set_protocol(Protocol::Control);
// packet.set_transport_protocol(control_packet::Protocol::PunchRequest.into());
// packet.set_source(current_device.virtual_ip());
// packet.set_destination(source);
// self.client_cipher.encrypt_ipv4(&mut packet)?;
// let _ = context.try_send_main_udp(packet.buffer(),
// SocketAddr::V4(SocketAddrV4::new(peer_nat_info.local_ip, peer_nat_info.local_port)));
// }
if self.punch(source, peer_nat_info).await {
self.client_cipher.encrypt_ipv4(&mut punch_packet)?;
context.try_send_by_key(punch_packet.buffer(), route_key)?;
}
} else {
self.punch(source, peer_nat_info).await;
}
}
other_turn_packet::Protocol::Unknown(e) => {
log::warn!("不支持的转发协议 {:?},source:{:?}", e, source);
}
}
Ok(())
}
async fn punch(&self, peer_ip: Ipv4Addr, peer_nat_info: NatInfo) -> bool {
match peer_nat_info.nat_type {
NatType::Symmetric => self
.symmetric_sender
.try_send((peer_ip, peer_nat_info))
.is_ok(),
NatType::Cone => self.cone_sender.try_send((peer_ip, peer_nat_info)).is_ok(),
}
}
}
/// 处理服务端数据
impl ChannelDataHandler {
async fn server_packet_handle(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
buf: &mut [u8],
data_len: usize,
route_key: &RouteKey,
) -> crate::Result<()> {
let net_packet = NetPacket::new0(data_len, &buf[14..])?;
let source = net_packet.source();
match net_packet.protocol() {
Protocol::Service => {
self.service(context, current_device, net_packet, route_key)
.await?;
}
Protocol::Error => {
self.error(context, current_device, source, net_packet, route_key)
.await?;
}
Protocol::Control => {
self.control_gateway(context, current_device, net_packet, route_key)
.await?;
}
Protocol::IpTurn => {
match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) {
ip_turn_packet::Protocol::Ipv4 => {
let ipv4 = IpV4Packet::new(net_packet.payload())?;
match ipv4.protocol() {
ipv4::protocol::Protocol::Igmp => {
if let Some(igmp_server) = &self.igmp_server {
igmp_server.handle(ipv4.payload(), source)?;
}
return Ok(());
}
ipv4::protocol::Protocol::Icmp => {
if ipv4.destination_ip() == current_device.virtual_ip {
let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?;
if icmp_packet.kind() == Kind::EchoReply {
self.device_writer.write_ipv4(&mut buf[12..])?;
return Ok(());
}
}
}
_ => {}
}
}
ip_turn_packet::Protocol::Ipv4Broadcast => {}
ip_turn_packet::Protocol::Unknown(_) => {}
}
}
Protocol::OtherTurn => {}
Protocol::UnKnow(_) => {}
}
return Ok(());
}
async fn control_gateway(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
net_packet: NetPacket<&[u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
ControlPacket::PongPacket(pong_packet) => {
let metric = net_packet.source_ttl() - net_packet.ttl() + 1;
self.pong_packet(
true,
metric,
context,
current_device,
net_packet.source(),
pong_packet,
route_key,
)
.await?;
}
ControlPacket::AddrResponse(addr_packet) => {
if addr_packet.port() != 0
&& !addr_packet.ipv4().is_multicast()
&& !addr_packet.ipv4().is_broadcast()
&& !addr_packet.ipv4().is_unspecified()
&& !addr_packet.ipv4().is_loopback()
&& !addr_packet.ipv4().is_private()
{
self.nat_test
.update_addr(addr_packet.ipv4(), addr_packet.port())
}
}
_ => {}
}
Ok(())
}
async fn service(
&self,
context: &Context,
current_device: CurrentDeviceInfo,
net_packet: NetPacket<&[u8]>,
route_key: &RouteKey,
) -> crate::Result<()> {
match service_packet::Protocol::from(net_packet.transport_protocol()) {
service_packet::Protocol::RegistrationRequest => {}
service_packet::Protocol::RegistrationResponse => {
let response = RegistrationResponse::parse_from_bytes(net_packet.payload())?;
{
let context = context.clone();
let nat_test = self.nat_test.clone();
tokio::spawn(async move {
let local_port = context.main_local_ipv4_port().unwrap_or(0);
let local_ipv4_addr = nat::local_ipv4_addr(local_port);
let local_port = context.main_local_ipv6_port().unwrap_or(0);
let ipv6_addr = nat::local_ipv6_addr(local_port);
let nat_info = nat_test
.re_test(
Ipv4Addr::from(response.public_ip),
response.public_port as u16,
local_ipv4_addr,
ipv6_addr,
)
.await;
context.switch(nat_info.nat_type);
});
}
let new_ip = Ipv4Addr::from(response.virtual_ip);
let current_ip = current_device.virtual_ip();
if current_ip != new_ip {
// ip发生变化
log::info!("ip发生变化,old_ip:{:?},new_ip:{:?}", current_ip, new_ip);
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
let old_netmask = current_device.virtual_netmask;
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
let old_gateway = current_device.virtual_gateway();
let virtual_ip = Ipv4Addr::from(response.virtual_ip);
let virtual_gateway = Ipv4Addr::from(response.virtual_gateway);
let virtual_netmask = Ipv4Addr::from(response.virtual_netmask);
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
self.device_writer.change_ip(
virtual_ip,
virtual_netmask,
virtual_gateway,
old_netmask,
old_gateway,
)?;
let new_current_device = CurrentDeviceInfo::new(
virtual_ip,
virtual_gateway,
virtual_netmask,
current_device.connect_server,
);
if let Err(e) = self
.current_device
.compare_exchange(current_device, new_current_device)
{
log::warn!("替换失败:{:?}", e);
}
}
self.connect_status.store(ConnectStatus::Connected);
}
service_packet::Protocol::PollDeviceList => {}
service_packet::Protocol::PushDeviceList => {
let device_list_t = DeviceList::parse_from_bytes(net_packet.payload())?;
let ip_list: Vec<PeerDeviceInfo> = device_list_t
.device_info_list
.into_iter()
.map(|info| {
PeerDeviceInfo::new(
Ipv4Addr::from(info.virtual_ip),
info.name,
info.device_status as u8,
info.client_secret,
)
})
.collect();
let route = Route::from(*route_key, 2, 199);
for x in &ip_list {
if x.status == PeerDeviceStatus::Online {
context.add_route_if_absent(x.virtual_ip, route);
}
}
let mut dev = self.device_list.lock();
if dev.0 != device_list_t.epoch as u16 {
dev.0 = device_list_t.epoch as u16;
dev.1 = ip_list;
}
}
service_packet::Protocol::Unknown(u) => {
log::warn!("未知服务协议:{}", u);
}
_ => {}
}
Ok(())
}
async fn error(
&self,
_context: &Context,
current_device: CurrentDeviceInfo,
_source: Ipv4Addr,
net_packet: NetPacket<&[u8]>,
_route_key: &RouteKey,
) -> crate::Result<()> {
log::info!("current_device:{:?}", current_device);
match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload())? {
InErrorPacket::TokenError => {
return Err(Error::Stop("Token error".to_string()));
}
InErrorPacket::Disconnect => {
{
//掉线epoch要归零
let mut dev = self.device_list.lock();
dev.0 = 0;
}
self.connect_status.store(ConnectStatus::Connecting);
self.register.fast_register(current_device.virtual_ip)?;
}
InErrorPacket::AddressExhausted => {
//地址用尽
return Err(Error::Stop("IP address has been exhausted".to_string()));
}
InErrorPacket::OtherError(e) => {
log::error!("OtherError {:?}", e.message());
}
InErrorPacket::IpAlreadyExists => {
log::error!("IpAlreadyExists");
}
InErrorPacket::InvalidIp => {
log::error!("InvalidIp");
}
InErrorPacket::NoKey => {}
}
Ok(())
}
}
+49
View File
@@ -0,0 +1,49 @@
use std::io;
use std::net::Ipv4Addr;
use protobuf::Message;
use crate::cipher::Cipher;
use crate::handle::{GATEWAY_IP, SELF_IP};
use crate::proto::message::RegistrationRequest;
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
/// 注册数据
pub fn registration_request_packet(
server_cipher: &Cipher,
token: String,
device_id: String,
name: String,
ip: Option<Ipv4Addr>,
is_fast: bool,
allow_ip_change: bool,
client_secret: bool,
) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = RegistrationRequest::new();
request.token = token;
request.device_id = device_id;
request.name = name;
if let Some(ip) = ip {
request.virtual_ip = ip.into();
}
request.allow_ip_change = allow_ip_change;
request.is_fast = is_fast;
request.version = crate::VNT_VERSION.to_string();
request.client_secret = client_secret;
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("RegistrationRequest {:?}", e))
})?;
let buf = vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED];
let mut net_packet = NetPacket::new_encrypt(buf)?;
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::RegistrationRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
server_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
-253
View File
@@ -1,253 +0,0 @@
use crossbeam_utils::atomic::AtomicCell;
use std::net::{Ipv4Addr, SocketAddr};
use std::time::{Duration, Instant};
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::handle::PeerDeviceInfo;
use protobuf::Message;
use std::net::UdpSocket;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use crate::proto::message::{RegistrationRequest, RegistrationResponse};
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::error_packet::InErrorPacket;
use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL};
pub enum ReqEnum {
TokenError,
AddressExhausted,
IpAlreadyExists,
InvalidIp,
Timeout,
ServerError(String),
Other(String),
}
#[derive(Clone, Debug)]
pub struct RegResponse {
pub virtual_ip: Ipv4Addr,
pub virtual_gateway: Ipv4Addr,
pub virtual_netmask: Ipv4Addr,
pub epoch: u16,
pub device_info_list: Vec<PeerDeviceInfo>,
pub public_ip: Ipv4Addr,
pub public_port: u16,
}
///向中继服务器注册,token标识一个虚拟网关,device_id防止多次注册时得到的ip不一致
pub async fn registration(
main_channel: &UdpSocket,
main_tcp_channel: Option<&mut TcpStream>,
server_cipher: &Cipher,
server_address: SocketAddr,
token: String,
device_id: String,
name: String,
ip: Ipv4Addr,
client_secret: bool,
) -> Result<RegResponse, ReqEnum> {
let request_packet = registration_request_packet(
server_cipher,
token.clone(),
device_id.clone(),
name.clone(),
ip,
false,
false,
client_secret,
)
.unwrap();
let buf = request_packet.buffer();
let mut recv_buf = [0u8; 10240];
let recv_buf = if let Some(main_tcp_channel) = main_tcp_channel {
let mut vec = vec![0; 4 + buf.len()];
let len = buf.len();
vec[2] = (len >> 8) as u8;
vec[3] = (len & 0xFF) as u8;
vec[4..].copy_from_slice(buf);
if let Err(e) = main_tcp_channel.write_all(&vec).await {
return Err(ReqEnum::Other(format!("send error:{}", e)));
}
if let Err(e) = main_tcp_channel.read_exact(&mut recv_buf[..4]).await {
return Err(ReqEnum::Other(format!("read error:{}", e)));
}
let len = 4 + (((recv_buf[2] as u16) << 8) | recv_buf[3] as u16) as usize;
if len > recv_buf.len() {
return Err(ReqEnum::Other("too long".to_string()));
}
if let Err(e) = main_tcp_channel.read_exact(&mut recv_buf[4..len]).await {
return Err(ReqEnum::Other(format!("read error:{}", e)));
}
&mut recv_buf[4..len]
} else {
if let Err(e) = main_channel.send_to(buf, server_address) {
return Err(ReqEnum::Other(format!("send error:{}", e)));
}
match main_channel.recv_from(&mut recv_buf) {
Ok((len, addr)) => {
if server_address != addr {
return Err(ReqEnum::Other(format!("invalid data,from {}", addr)));
}
&mut recv_buf[..len]
}
Err(e) => {
return Err(ReqEnum::Other(format!("receiver error:{}", e)));
}
}
};
let mut net_packet = match NetPacket::new(recv_buf) {
Ok(net_packet) => net_packet,
Err(e) => {
return Err(ReqEnum::ServerError(format!("{}", e)));
}
};
if let Err(e) = server_cipher.decrypt_ipv4(&mut net_packet) {
return Err(ReqEnum::ServerError(format!("decrypt_ipv4 {}", e)));
}
match net_packet.protocol() {
Protocol::Service => {
match service_packet::Protocol::from(net_packet.transport_protocol()) {
service_packet::Protocol::RegistrationResponse => {
match RegistrationResponse::parse_from_bytes(net_packet.payload()) {
Ok(response) => {
let device_info_list: Vec<PeerDeviceInfo> = response
.device_info_list
.into_iter()
.map(|info| {
PeerDeviceInfo::new(
Ipv4Addr::from(info.virtual_ip),
info.name,
info.device_status as u8,
info.client_secret,
)
})
.collect();
Ok(RegResponse {
virtual_ip: Ipv4Addr::from(response.virtual_ip),
virtual_gateway: Ipv4Addr::from(response.virtual_gateway),
virtual_netmask: Ipv4Addr::from(response.virtual_netmask),
epoch: response.epoch as u16,
device_info_list,
public_ip: Ipv4Addr::from(response.public_ip),
public_port: response.public_port as u16,
})
}
Err(_) => Err(ReqEnum::ServerError("invalid data".to_string())),
}
}
_ => Err(ReqEnum::ServerError("invalid data".to_string())),
}
}
Protocol::Error => {
match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload()) {
Ok(e) => match e {
InErrorPacket::TokenError => Err(ReqEnum::TokenError),
InErrorPacket::Disconnect => {
Err(ReqEnum::ServerError("disconnect".to_string()))
}
InErrorPacket::AddressExhausted => Err(ReqEnum::AddressExhausted),
InErrorPacket::OtherError(e) => match e.message() {
Ok(str) => Err(ReqEnum::ServerError(str)),
Err(e) => Err(ReqEnum::Other(format!("{}", e))),
},
InErrorPacket::IpAlreadyExists => Err(ReqEnum::IpAlreadyExists),
InErrorPacket::InvalidIp => Err(ReqEnum::InvalidIp),
InErrorPacket::NoKey => Err(ReqEnum::ServerError("no key".to_string())),
},
Err(e) => Err(ReqEnum::Other(format!("{}", e))),
}
}
_ => Err(ReqEnum::ServerError("invalid data".to_string())),
}
}
fn registration_request_packet(
server_cipher: &Cipher,
token: String,
device_id: String,
name: String,
ip: Ipv4Addr,
is_fast: bool,
allow_ip_change: bool,
client_secret: bool,
) -> crate::Result<NetPacket<Vec<u8>>> {
let mut request = RegistrationRequest::new();
request.token = token;
request.device_id = device_id;
request.name = name;
request.virtual_ip = ip.into();
request.allow_ip_change = allow_ip_change;
request.is_fast = is_fast;
request.version = crate::VNT_VERSION.to_string();
request.client_secret = client_secret;
let bytes = request.write_to_bytes()?;
let buf = vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED];
let mut net_packet = NetPacket::new_encrypt(buf)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::RegistrationRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
server_cipher.encrypt_ipv4(&mut net_packet)?;
Ok(net_packet)
}
pub struct Register {
server_cipher: Cipher,
sender: ChannelSender,
server_address: SocketAddr,
token: String,
device_id: String,
name: String,
time: AtomicCell<Instant>,
client_secret: bool,
}
impl Register {
pub fn new(
server_cipher: Cipher,
sender: ChannelSender,
server_address: SocketAddr,
token: String,
device_id: String,
name: String,
client_secret: bool,
) -> Self {
Self {
server_cipher,
sender,
server_address,
token,
device_id,
name,
time: AtomicCell::new(Instant::now()),
client_secret,
}
}
pub fn fast_register(&self, ip: Ipv4Addr) -> crate::Result<()> {
let last = self.time.load();
if last.elapsed() < Duration::from_secs(2)
|| self.time.compare_exchange(last, Instant::now()).is_err()
{
//短时间不重复注册
return Ok(());
}
log::info!("重新连接");
let request_packet = registration_request_packet(
&self.server_cipher,
self.token.clone(),
self.device_id.clone(),
self.name.clone(),
ip,
false,
true,
self.client_secret,
)?;
let buf = request_packet.buffer();
self.sender.send_main(buf, self.server_address)?;
Ok(())
}
}
+24 -26
View File
@@ -1,32 +1,30 @@
use byte_pool::Block;
use std::sync::mpsc::{sync_channel, Receiver, SendError, SyncSender};
#[derive(Clone)]
pub struct BufSenderGroup(
usize,
Vec<std::sync::mpsc::SyncSender<(Block<'static>, usize, usize)>>,
);
pub struct BufReceiverGroup(pub Vec<std::sync::mpsc::Receiver<(Block<'static>, usize, usize)>>);
impl BufSenderGroup {
pub fn send(&mut self, val: (Block<'static>, usize, usize)) -> bool {
let index = self.0 % self.1.len();
self.0 = self.0.wrapping_add(1);
self.1[index].send(val).is_ok()
}
}
pub fn buf_channel_group(size: usize) -> (BufSenderGroup, BufReceiverGroup) {
let mut buf_sender_group = Vec::with_capacity(size);
let mut buf_receiver_group = Vec::with_capacity(size);
pub fn channel_group<T>(size: usize, bound: usize) -> (GroupSyncSender<T>, Vec<Receiver<T>>) {
let mut senders = Vec::with_capacity(size);
let mut receivers = Vec::with_capacity(size);
for _ in 0..size {
let (buf_sender, buf_receiver) =
std::sync::mpsc::sync_channel::<(Block<'static>, usize, usize)>(1);
buf_sender_group.push(buf_sender);
buf_receiver_group.push(buf_receiver);
let (s, r) = sync_channel(bound);
senders.push(s);
receivers.push(r);
}
(
BufSenderGroup(0, buf_sender_group),
BufReceiverGroup(buf_receiver_group),
GroupSyncSender {
count: 0,
base: senders,
},
receivers,
)
}
pub struct GroupSyncSender<T> {
count: usize,
base: Vec<SyncSender<T>>,
}
impl<T> GroupSyncSender<T> {
pub fn send(&mut self, t: T) -> Result<(), SendError<T>> {
self.count += 1;
self.base[self.count % self.base.len()].send(t)
}
}
+49 -151
View File
@@ -1,37 +1,31 @@
use crate::channel::sender::ChannelSender;
use std::io;
use std::net::Ipv4Addr;
use crate::channel::context::Context;
use packet::ip::ipv4::packet::IpV4Packet;
use packet::ip::ipv4::protocol::Protocol;
use crate::cipher::Cipher;
use crate::error::*;
use crate::external_route::ExternalRoute;
use crate::handle::{check_dest, CurrentDeviceInfo};
use crate::igmp_server::{IgmpServer, Multicast};
use crate::ip_proxy::IpProxyMap;
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::{IpProxyMap, ProxyHandler};
use crate::protocol;
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::ip_turn_packet::BroadcastPacket;
use crate::protocol::{ip_turn_packet, NetPacket, Version, MAX_TTL};
use packet::ip::ipv4::packet::IpV4Packet;
use packet::ip::ipv4::protocol::Protocol;
use packet::tcp::tcp::TcpPacket;
use packet::udp::udp::UdpPacket;
use parking_lot::RwLock;
use std::io;
use std::net::{Ipv4Addr, SocketAddrV4};
use std::sync::Arc;
pub mod channel_group;
#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))]
pub mod tap_handler;
mod channel_group;
pub mod tun_handler;
fn broadcast(
server_cipher: &Cipher,
multicast_members: Option<Arc<RwLock<Multicast>>>,
sender: &ChannelSender,
sender: &Context,
net_packet: &mut NetPacket<&mut [u8]>,
current_device: &CurrentDeviceInfo,
) -> Result<()> {
) -> io::Result<()> {
let mut peer_ips = Vec::with_capacity(8);
let vec = sender.route_table_one();
let vec = sender.route_table.route_table_one();
let mut relay_count = 0;
const MAX_COUNT: usize = 8;
for (peer_ip, route) in vec {
@@ -39,16 +33,12 @@ fn broadcast(
continue;
}
if peer_ips.len() == MAX_COUNT {
relay_count += 1;
break;
}
if let Some(members) = &multicast_members {
if !members.read().is_send(&peer_ip) {
continue;
}
}
if route.is_p2p()
&& sender
.try_send_by_key(net_packet.buffer(), &route.route_key())
.send_by_key(net_packet.buffer(), route.route_key())
.is_ok()
{
peer_ips.push(peer_ip);
@@ -56,18 +46,16 @@ fn broadcast(
relay_count += 1;
}
}
if relay_count == 0 && !peer_ips.is_empty() && peer_ips.len() != MAX_COUNT {
if (relay_count == 0 && !peer_ips.is_empty()) || current_device.status.offline() {
//不需要转发
return Ok(());
}
//转发到服务端的可选择广播,还要进行服务端加密
if peer_ips.is_empty() {
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
sender.send_default(net_packet.buffer(), current_device.connect_server)?;
} else {
let buf = vec![
0 as u8;
12 + 1 + peer_ips.len() * 4 + net_packet.data_len() + ENCRYPTION_RESERVED
];
let buf =
vec![0u8; 12 + 1 + peer_ips.len() * 4 + net_packet.data_len() + ENCRYPTION_RESERVED];
//剩余的发送到服务端,需要告知哪些已发送过
let mut server_packet = NetPacket::new_encrypt(buf)?;
server_packet.set_version(Version::V1);
@@ -83,7 +71,7 @@ fn broadcast(
broadcast.set_address(&peer_ips)?;
broadcast.set_data(net_packet.buffer())?;
server_cipher.encrypt_ipv4(&mut server_packet)?;
sender.send_main(server_packet.buffer(), current_device.connect_server)?;
sender.send_default(server_packet.buffer(), current_device.connect_server)?;
}
Ok(())
}
@@ -93,81 +81,44 @@ fn broadcast(
///
#[inline]
pub fn base_handle(
sender: &ChannelSender,
context: &Context,
buf: &mut [u8],
data_len: usize, //数据总长度=12+ip包长度
igmp_server: &Option<IgmpServer>,
current_device: CurrentDeviceInfo,
ip_route: &Option<ExternalRoute>,
proxy_map: &Option<IpProxyMap>,
ip_route: &ExternalRoute,
#[cfg(feature = "ip_proxy")] proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> Result<()> {
) -> io::Result<()> {
let ipv4_packet = IpV4Packet::new(&buf[12..data_len])?;
let protocol = ipv4_packet.protocol();
let ip_head_len = ipv4_packet.header_len() as usize * 4;
if 12 + ip_head_len >= data_len {
Err(io::Error::new(io::ErrorKind::Other, "ip_head_len err"))?
}
let src_ip = ipv4_packet.source_ip();
let mut dest_ip = ipv4_packet.destination_ip();
let mut net_packet = NetPacket::new0(data_len, buf)?;
net_packet.set_version(Version::V1);
net_packet.set_protocol(protocol::Protocol::IpTurn);
net_packet.set_transport_protocol(ip_turn_packet::Protocol::Ipv4.into());
net_packet.first_set_ttl(3);
net_packet.first_set_ttl(6);
net_packet.set_source(src_ip);
net_packet.set_destination(dest_ip);
if dest_ip == current_device.virtual_gateway {
// 发到网关的加密方式不一样,要单独处理
if protocol == Protocol::Icmp {
net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?;
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
context.send_default(net_packet.buffer(), current_device.connect_server)?;
}
return Ok(());
}
if dest_ip.is_multicast() {
match protocol {
Protocol::Igmp => {
if igmp_server.is_some() {
//发送到服务端
net_packet.set_destination(current_device.virtual_gateway);
net_packet.set_gateway_flag(true);
server_cipher.encrypt_ipv4(&mut net_packet)?;
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
}
}
Protocol::Udp => {
let multicast_members = if let Some(igmp_server) = igmp_server {
igmp_server.load(&dest_ip)
} else {
//当作广播处理
net_packet.set_destination(Ipv4Addr::BROADCAST);
None
};
client_cipher.encrypt_ipv4(&mut net_packet)?;
broadcast(
server_cipher,
multicast_members,
sender,
&mut net_packet,
&current_device,
)?;
}
_ => {}
}
return Ok(());
//当作广播处理
dest_ip = Ipv4Addr::BROADCAST;
net_packet.set_destination(Ipv4Addr::BROADCAST);
}
if dest_ip.is_broadcast() || current_device.broadcast_address == dest_ip {
if dest_ip.is_broadcast() || current_device.broadcast_ip == dest_ip {
// 广播 发送到直连目标
client_cipher.encrypt_ipv4(&mut net_packet)?;
broadcast(
server_cipher,
None,
sender,
&mut net_packet,
&current_device,
)?;
broadcast(server_cipher, context, &mut net_packet, &current_device)?;
return Ok(());
}
if !check_dest(
@@ -175,81 +126,28 @@ pub fn base_handle(
current_device.virtual_netmask,
current_device.virtual_network,
) {
if let Some(ip_route) = ip_route {
if let Some(r_dest_ip) = ip_route.route(&dest_ip) {
//路由的目标不能是自己
if r_dest_ip == src_ip {
return Ok(());
}
//需要修改目的地址
dest_ip = r_dest_ip;
net_packet.set_destination(r_dest_ip);
} else {
if let Some(r_dest_ip) = ip_route.route(&dest_ip) {
//路由的目标不能是自己
if r_dest_ip == src_ip {
return Ok(());
}
//需要修改目的地址
dest_ip = r_dest_ip;
net_packet.set_destination(r_dest_ip);
} else {
return Ok(());
}
} else if let Some(proxy_map) = proxy_map {
match protocol {
Protocol::Tcp => {
let dest_addr = {
let tcp_packet = TcpPacket::new(
src_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
SocketAddrV4::new(dest_ip, tcp_packet.destination_port())
};
if let Some(entry) = proxy_map.tcp_proxy_map.get(&dest_addr) {
let source_addr = entry.value();
let source_ip = *source_addr.ip();
let mut tcp_packet = TcpPacket::new(
source_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
tcp_packet.set_source_port(source_addr.port());
tcp_packet.update_checksum();
let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?;
ipv4_packet.set_source_ip(source_ip);
ipv4_packet.update_checksum();
}
}
Protocol::Udp => {
let dest_addr = {
let udp_packet = UdpPacket::new(
src_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
SocketAddrV4::new(dest_ip, udp_packet.destination_port())
};
if let Some(entry) = proxy_map.udp_proxy_map.get(&dest_addr) {
let source_addr = entry.value();
let source_ip = *source_addr.ip();
let mut udp_packet = UdpPacket::new(
source_ip,
dest_ip,
&mut net_packet.payload_mut()[ip_head_len..],
)?;
udp_packet.set_source_port(source_addr.port());
udp_packet.update_checksum();
let mut ipv4_packet = IpV4Packet::new(net_packet.payload_mut())?;
ipv4_packet.set_source_ip(source_ip);
ipv4_packet.update_checksum();
}
}
_ => {}
}
}
#[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)?;
//优先发到直连到地址
if sender
.try_send_by_id(net_packet.buffer(), &dest_ip)
.is_err()
{
sender.send_main(net_packet.buffer(), current_device.connect_server)?;
}
return Ok(());
context.send_ipv4_by_id(
net_packet.buffer(),
&dest_ip,
current_device.connect_server,
current_device.status.online(),
)
}
-259
View File
@@ -1,259 +0,0 @@
use byte_pool::BytePool;
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use lazy_static::lazy_static;
use packet::arp::arp::ArpPacket;
use packet::ethernet;
use packet::ethernet::packet::EthernetPacket;
use packet::icmp::icmp::IcmpPacket;
use packet::icmp::Kind;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use crate::external_route::ExternalRoute;
use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup};
use crate::handle::CurrentDeviceInfo;
use crate::igmp_server::IgmpServer;
use crate::ip_proxy::IpProxyMap;
use crate::tun_tap_device::{DeviceReader, DeviceWriter};
lazy_static! {
static ref POOL: BytePool<Vec<u8>> = BytePool::<Vec<u8>>::new();
}
pub fn start(
worker: VntWorker,
sender: ChannelSender,
device_reader: DeviceReader,
device_writer: DeviceWriter,
igmp_server: Option<IgmpServer>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_route: Option<ExternalRoute>,
ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
parallel: usize,
) {
if parallel == 1 {
thread::Builder::new()
.name("tap_handler".into())
.spawn(move || {
if let Err(e) = start_simple(
&sender,
device_reader,
&device_writer,
igmp_server,
current_device,
ip_route,
ip_proxy_map,
client_cipher,
server_cipher,
) {
log::warn!("tap:{:?}", e);
}
let _ = sender.close();
let _ = device_writer.close();
worker.stop_all();
})
.unwrap();
} else {
let (buf_sender, buf_receiver) = buf_channel_group(parallel);
for buf_receiver in buf_receiver.0 {
let sender = sender.clone();
let device_writer = device_writer.clone();
let igmp_server = igmp_server.clone();
let current_device = current_device.clone();
let ip_route = ip_route.clone();
let ip_proxy_map = ip_proxy_map.clone();
let client_cipher = client_cipher.clone();
let server_cipher = server_cipher.clone();
thread::spawn(move || {
while let Ok((mut buf, _, len)) = buf_receiver.recv() {
match handle(
&mut buf,
len,
&igmp_server,
&current_device,
&device_writer,
&sender,
&ip_route,
&ip_proxy_map,
&client_cipher,
&server_cipher,
) {
Ok(_) => {}
Err(e) => {
log::warn!("{:?}", e)
}
}
}
let _ = sender.close();
let _ = device_writer.close();
});
}
thread::Builder::new()
.name("tap_handler".into())
.spawn(move || {
if let Err(e) = start_(&sender, device_reader, buf_sender) {
log::warn!("tap:{:?}", e);
}
let _ = sender.close();
let _ = device_writer.close();
worker.stop_all();
})
.unwrap();
}
}
fn start_(
sender: &ChannelSender,
device_reader: DeviceReader,
mut buf_sender: BufSenderGroup,
) -> io::Result<()> {
loop {
let mut buf = POOL.alloc(4096);
if sender.is_close() {
return Ok(());
}
let start = 0;
let len = device_reader.read(&mut buf)?;
if !buf_sender.send((buf, start, len)) {
return Err(io::Error::new(
io::ErrorKind::Other,
"tap buf_sender发送失败",
));
}
}
}
fn start_simple(
sender: &ChannelSender,
device_reader: DeviceReader,
device_writer: &DeviceWriter,
igmp_server: Option<IgmpServer>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_route: Option<ExternalRoute>,
ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
) -> io::Result<()> {
let mut buf = [0; 4096];
loop {
let len = device_reader.read(&mut buf)?;
if let Err(e) = handle(
&mut buf,
len,
&igmp_server,
&current_device,
device_writer,
sender,
&ip_route,
&ip_proxy_map,
&client_cipher,
&server_cipher,
) {
log::warn!("tap handle{:?}", e);
}
}
}
fn handle(
buf: &mut [u8],
len: usize,
igmp_server: &Option<IgmpServer>,
current_device: &AtomicCell<CurrentDeviceInfo>,
device_writer: &DeviceWriter,
sender: &ChannelSender,
ip_route: &Option<ExternalRoute>,
proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> crate::Result<()> {
let mut ethernet_packet = EthernetPacket::new(&mut buf[..len])?;
let current_device = current_device.load();
match ethernet_packet.protocol() {
ethernet::protocol::Protocol::Arp => {
let mut out_ethernet_packet =
EthernetPacket::unchecked(ethernet_packet.buffer.to_vec());
let arp_packet = ArpPacket::unchecked(ethernet_packet.payload());
let mut out_arp_packet = ArpPacket::unchecked(out_ethernet_packet.payload_mut());
let sender_h = arp_packet.sender_hardware_addr();
let sender_p = arp_packet.sender_protocol_addr();
let target_p = arp_packet.target_protocol_addr();
if target_p == &[0, 0, 0, 0] || sender_p == &[0, 0, 0, 0] || target_p == sender_p {
return Ok(());
}
//回复一个虚假的MAC地址
out_arp_packet.set_sender_hardware_addr(&[
target_p[0],
target_p[1],
target_p[2],
target_p[3],
!sender_h[5],
234,
]);
out_arp_packet.set_sender_protocol_addr(target_p);
out_arp_packet.set_target_hardware_addr(sender_h);
out_arp_packet.set_target_protocol_addr(sender_p);
out_arp_packet.set_op_code(2);
out_ethernet_packet.set_source(&[
target_p[0],
target_p[1],
target_p[2],
target_p[3],
!sender_h[5],
234,
]);
out_ethernet_packet.set_destination(sender_h);
device_writer.write_ethernet_tap(&out_ethernet_packet.buffer)?;
}
ethernet::protocol::Protocol::Ipv4 => {
let mut ipv4_packet = IpV4Packet::unchecked(ethernet_packet.payload_mut());
let src_ip = ipv4_packet.source_ip();
if src_ip != current_device.virtual_ip() {
return Ok(());
}
let dest_ip = ipv4_packet.destination_ip();
let protocol = ipv4_packet.protocol();
if src_ip == dest_ip {
if protocol == ipv4::protocol::Protocol::Icmp {
let mut icmp = IcmpPacket::new(ipv4_packet.payload_mut())?;
if icmp.kind() == Kind::EchoRequest {
icmp.set_kind(Kind::EchoReply);
icmp.update_checksum();
ipv4_packet.set_source_ip(dest_ip);
ipv4_packet.set_destination_ip(src_ip);
ipv4_packet.update_checksum();
let source = ethernet_packet.source().to_vec();
let dest = ethernet_packet.destination().to_vec();
ethernet_packet.set_source(&dest);
ethernet_packet.set_destination(&source);
device_writer.write_ethernet_tap(&ethernet_packet.buffer)?;
}
}
return Ok(());
}
// 以太网帧头部14字节,预留12字节
return crate::handle::tun_tap::base_handle(
sender,
&mut buf[2..],
len - 2,
igmp_server,
current_device,
ip_route,
proxy_map,
client_cipher,
server_cipher,
);
}
_ => {
// log::warn!("不支持的二层协议:{:?}",p)
}
}
Ok(())
}
+135 -129
View File
@@ -1,28 +1,25 @@
use byte_pool::BytePool;
use std::sync::Arc;
use std::{io, thread};
use crossbeam_utils::atomic::AtomicCell;
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::core::status::VntWorker;
use packet::icmp::icmp::IcmpPacket;
use packet::icmp::Kind;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use tun::device::IFace;
use tun::Device;
use crate::error::*;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::external_route::ExternalRoute;
use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup};
use crate::handle::tun_tap::channel_group::{channel_group, GroupSyncSender};
use crate::handle::CurrentDeviceInfo;
use crate::igmp_server::IgmpServer;
#[cfg(feature = "ip_proxy")]
use crate::ip_proxy::IpProxyMap;
use crate::tun_tap_device::{DeviceReader, DeviceWriter};
lazy_static::lazy_static! {
static ref POOL:BytePool<Vec<u8>> = BytePool::<Vec<u8>>::new();
}
fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> Result<()> {
use crate::util::{SingleU64Adder, StopManager};
fn icmp(device_writer: &Device, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> io::Result<()> {
if ipv4_packet.protocol() == ipv4::protocol::Protocol::Icmp {
let mut icmp = IcmpPacket::new(ipv4_packet.payload_mut())?;
if icmp.kind() == Kind::EchoRequest {
@@ -32,46 +29,41 @@ fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) ->
ipv4_packet.set_source_ip(ipv4_packet.destination_ip());
ipv4_packet.set_destination_ip(src);
ipv4_packet.update_checksum();
device_writer.write_ipv4_tun(ipv4_packet.buffer)?;
device_writer.write(ipv4_packet.buffer)?;
}
}
Ok(())
}
/// 接收tun数据,并且转发到udp上
#[inline]
fn handle(
sender: &ChannelSender,
context: &Context,
data: &mut [u8],
len: usize,
device_writer: &DeviceWriter,
igmp_server: &Option<IgmpServer>,
device_writer: &Device,
current_device: CurrentDeviceInfo,
ip_route: &Option<ExternalRoute>,
proxy_map: &Option<IpProxyMap>,
ip_route: &ExternalRoute,
#[cfg(feature = "ip_proxy")] proxy_map: &Option<IpProxyMap>,
client_cipher: &Cipher,
server_cipher: &Cipher,
) -> Result<()> {
let ipv4_packet = if let Ok(ipv4_packet) = IpV4Packet::new(&mut data[12..len]) {
ipv4_packet
} else {
return Ok(());
) -> 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 != current_device.virtual_ip() {
return Ok(());
}
if src_ip == dest_ip {
return icmp(&device_writer, ipv4_packet);
}
return crate::handle::tun_tap::base_handle(
sender,
context,
data,
len,
igmp_server,
current_device,
ip_route,
#[cfg(feature = "ip_proxy")]
proxy_map,
client_cipher,
server_cipher,
@@ -79,142 +71,136 @@ fn handle(
}
pub fn start(
worker: VntWorker,
sender: ChannelSender,
device_reader: DeviceReader,
device_writer: DeviceWriter,
igmp_server: Option<IgmpServer>,
stop_manager: StopManager,
context: Context,
device: Arc<Device>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_route: Option<ExternalRoute>,
ip_proxy_map: Option<IpProxyMap>,
ip_route: ExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
parallel: usize,
) {
if parallel == 1 {
thread::Builder::new()
.name("tun_handler".into())
.spawn(move || {
if let Err(e) = start_simple(
&sender,
device_reader,
&device_writer,
igmp_server,
current_device,
ip_route,
ip_proxy_map,
client_cipher,
server_cipher,
) {
log::warn!("stop:{}", e);
mut up_counter: SingleU64Adder,
) -> io::Result<()> {
let worker = {
#[cfg(target_os = "macos")]
let current_device = current_device.clone();
let device = device.clone();
stop_manager.add_listener("tun_device".into(), move || {
if let Err(e) = device.shutdown() {
log::warn!("{:?}", e);
}
#[cfg(target_os = "macos")]
{
let ip = current_device.load().virtual_ip;
if let Ok(udp) = std::net::UdpSocket::bind("0.0.0.0:0") {
let _ = udp.send_to(b"stop", format!("{:?}:1234", ip));
}
let _ = sender.close();
let _ = device_writer.close();
worker.stop_all();
})
.unwrap();
} else {
let (buf_sender, buf_receiver) = buf_channel_group(parallel);
for buf_receiver in buf_receiver.0 {
let sender = sender.clone();
let device_writer = device_writer.clone();
let igmp_server = igmp_server.clone();
}
})?
};
if parallel > 1 {
let (sender, receivers) = channel_group::<(Vec<u8>, usize)>(parallel, 16);
for (index, receiver) in receivers.into_iter().enumerate() {
let context = context.clone();
let device = device.clone();
let current_device = current_device.clone();
let ip_route = ip_route.clone();
#[cfg(feature = "ip_proxy")]
let ip_proxy_map = ip_proxy_map.clone();
let client_cipher = client_cipher.clone();
let server_cipher = server_cipher.clone();
thread::spawn(move || {
while let Ok((mut buf, start, len)) = buf_receiver.recv() {
match handle(
&sender,
&mut buf[start..],
len,
&device_writer,
&igmp_server,
current_device.load(),
&ip_route,
&ip_proxy_map,
&client_cipher,
&server_cipher,
) {
Ok(_) => {}
Err(e) => {
log::warn!("{:?}", e)
thread::Builder::new()
.name(format!("tunHandler-{}", index))
.spawn(move || {
while let Ok((mut buf, len)) = receiver.recv() {
#[cfg(not(target_os = "macos"))]
let start = 0;
#[cfg(target_os = "macos")]
let start = 4;
match handle(
&context,
&mut buf[start..],
len,
&device,
current_device.load(),
&ip_route,
#[cfg(feature = "ip_proxy")]
&ip_proxy_map,
&client_cipher,
&server_cipher,
) {
Ok(_) => {}
Err(e) => {
log::warn!("{:?}", e)
}
}
}
}
let _ = sender.close();
let _ = device_writer.close();
});
})?;
}
thread::Builder::new()
.name("tun_handler".into())
.name("tunHandlerM".into())
.spawn(move || {
if let Err(e) = start_(&sender, device_reader, buf_sender) {
if let Err(e) = start_multi(stop_manager, device, sender, &mut up_counter) {
log::warn!("stop:{}", e);
}
let _ = sender.close();
let _ = device_writer.close();
worker.stop_all();
})
.unwrap();
}
}
fn start_(
sender: &ChannelSender,
device_reader: DeviceReader,
mut buf_sender: BufSenderGroup,
) -> io::Result<()> {
loop {
let mut buf = POOL.alloc(4096);
buf[..12].fill(0);
if sender.is_close() {
return Ok(());
}
let start = 0;
let len = device_reader.read(&mut buf[12..])? + 12;
#[cfg(any(target_os = "macos"))]
let start = 4;
if !buf_sender.send((buf, start, len)) {
return Err(io::Error::new(
io::ErrorKind::Other,
"tun buf_sender发送失败",
));
}
})?;
} else {
thread::Builder::new()
.name("tunHandlerS".into())
.spawn(move || {
if let Err(e) = start_simple(
stop_manager,
&context,
device,
current_device,
ip_route,
#[cfg(feature = "ip_proxy")]
ip_proxy_map,
client_cipher,
server_cipher,
&mut up_counter,
) {
log::warn!("stop:{}", e);
}
worker.stop_all();
})?;
}
Ok(())
}
fn start_simple(
sender: &ChannelSender,
device_reader: DeviceReader,
device_writer: &DeviceWriter,
igmp_server: Option<IgmpServer>,
stop_manager: StopManager,
context: &Context,
device: Arc<Device>,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
ip_route: Option<ExternalRoute>,
ip_proxy_map: Option<IpProxyMap>,
ip_route: ExternalRoute,
#[cfg(feature = "ip_proxy")] ip_proxy_map: Option<IpProxyMap>,
client_cipher: Cipher,
server_cipher: Cipher,
up_counter: &mut SingleU64Adder,
) -> io::Result<()> {
let mut buf = [0; 4096];
let mut buf = [0; 1024 * 16];
loop {
if sender.is_close() {
if stop_manager.is_stop() {
return Ok(());
}
buf[..12].fill(0);
let len = device_reader.read(&mut buf[12..])? + 12;
let len = device.read(&mut buf[12..])? + 12;
//单线程的
up_counter.add(len as u64);
#[cfg(any(target_os = "macos"))]
let mut buf = &mut buf[4..];
// buf是重复利用的,需要重置头部
buf[..12].fill(0);
match handle(
sender,
context,
&mut buf,
len,
device_writer,
&igmp_server,
&device,
current_device.load(),
&ip_route,
#[cfg(feature = "ip_proxy")]
&ip_proxy_map,
&client_cipher,
&server_cipher,
@@ -226,3 +212,23 @@ fn start_simple(
}
}
}
fn start_multi(
stop_manager: StopManager,
device: Arc<Device>,
mut group_sync_sender: GroupSyncSender<(Vec<u8>, usize)>,
up_counter: &mut SingleU64Adder,
) -> io::Result<()> {
loop {
if stop_manager.is_stop() {
return Ok(());
}
let mut buf = vec![0; 1024 * 16];
let len = device.read(&mut buf[12..])? + 12;
//单线程的
up_counter.add(len as u64);
if group_sync_sender.send((buf, len)).is_err() {
return Ok(());
}
}
}
-243
View File
@@ -1,243 +0,0 @@
use std::collections::{HashMap, HashSet};
use std::net::Ipv4Addr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use dashmap::DashMap;
use parking_lot::RwLock;
use packet::igmp::igmp_v2::IgmpV2Packet;
use packet::igmp::igmp_v3::{IgmpV3QueryPacket, IgmpV3RecordType, IgmpV3ReportPacket};
use packet::igmp::IgmpType;
use packet::ip::ipv4::protocol::Protocol;
use crate::ip_proxy::DashMapNew;
use crate::tun_tap_device::DeviceWriter;
//1. 定时发送query,启动时20秒一次,连发3次,之后8分钟一次
//2. 接收网关的igmp report 维护组播源信息
#[derive(Clone, Debug)]
pub struct Multicast {
//成员虚拟ip
members: HashMap<Ipv4Addr, Instant>,
//是否是过滤模式
//成员过滤或包含的源ip
map: HashMap<Ipv4Addr, (bool, HashSet<Ipv4Addr>)>,
}
impl Multicast {
pub fn new() -> Self {
Self {
members: Default::default(),
map: Default::default(),
}
}
pub fn is_send(&self, ip: &Ipv4Addr) -> bool {
if self.members.contains_key(ip) {
if let Some((is_include, set)) = self.map.get(ip) {
if *is_include {
set.contains(ip)
} else {
!set.contains(ip)
}
} else {
true
}
} else {
false
}
}
}
#[derive(Clone)]
pub struct IgmpServer {
multicast: Arc<DashMap<Ipv4Addr, Arc<RwLock<Multicast>>>>,
}
impl IgmpServer {
pub fn new(device_writer: DeviceWriter) -> Self {
let multicast: Arc<DashMap<Ipv4Addr, Arc<RwLock<Multicast>>>> = Arc::new(DashMap::new0());
std::thread::spawn(move || {
//预留以太网帧头和ip头
let mut buf = [0; 14 + 24 + 12];
let dest = Ipv4Addr::new(224, 0, 0, 1);
let src = Ipv4Addr::new(10, 26, 0, 1);
{
let buf = &mut buf[14..];
let len = buf.len();
// ipv4 头部20字节
buf[0] = 0b0100_0110;
//写入总长度
buf[2..4].copy_from_slice(&(len as u16).to_be_bytes());
//ttl
buf[8] = 1;
buf[20] = 0x94;
buf[21] = 0x04;
let mut ipv4 = packet::ip::ipv4::packet::IpV4Packet::unchecked(buf);
ipv4.set_flags(2);
ipv4.set_protocol(Protocol::Igmp);
ipv4.set_source_ip(src);
ipv4.set_destination_ip(dest);
ipv4.update_checksum();
}
{
let mut igmp_query = IgmpV3QueryPacket::unchecked(&mut buf[14 + 24..]);
igmp_query.set_igmp_type();
igmp_query.set_max_resp_code(50);
igmp_query.set_group_address(Ipv4Addr::UNSPECIFIED);
igmp_query.set_qrv(2);
igmp_query.set_qqic(10);
igmp_query.update_checksum();
}
loop {
let _ = device_writer.write_ipv4(&mut buf);
std::thread::sleep(Duration::from_secs(20))
}
});
Self { multicast }
}
pub fn load(&self, multicast_addr: &Ipv4Addr) -> Option<Arc<RwLock<Multicast>>> {
if let Some(entry) = self.multicast.get(multicast_addr) {
Some(entry.value().clone())
} else {
None
}
}
pub fn handle(&self, buf: &[u8], source: Ipv4Addr) -> crate::Result<()> {
for x in self.multicast.iter() {
let mut list = Vec::new();
let mut write_guard = x.value().write();
for (ip, time) in &write_guard.members {
if time.elapsed() > Duration::from_secs(30) {
list.push(*ip);
}
}
for ip in list {
write_guard.members.remove(&ip);
write_guard.map.remove(&ip);
}
}
match IgmpType::from(buf[0]) {
IgmpType::Query => {}
IgmpType::ReportV1 | IgmpType::ReportV2 => {
//加入组播,v1和v2差不多
let report = IgmpV2Packet::new(buf)?;
let multicast_addr = report.group_address();
if !multicast_addr.is_multicast() {
return Ok(());
}
let multi = {
self.multicast
.entry(multicast_addr)
.or_insert_with(|| Arc::new(RwLock::new(Multicast::new())))
.value()
.clone()
};
let mut guard = multi.write();
guard.members.insert(source, Instant::now());
}
IgmpType::LeaveV2 => {
//退出组播
let leave = IgmpV2Packet::new(buf)?;
let multicast_addr = leave.group_address();
if !multicast_addr.is_multicast() {
return Ok(());
}
if let Some(entry) = self.multicast.get(&multicast_addr) {
let mut guard = entry.value().write();
guard.map.remove(&source);
guard.members.remove(&source);
}
}
IgmpType::ReportV3 => {
let report = IgmpV3ReportPacket::new(buf)?;
if let Some(group_records) = report.group_records() {
for group_record in group_records {
let multicast_addr = group_record.multicast_address();
if !multicast_addr.is_multicast() {
return Ok(());
}
let multi = self
.multicast
.entry(multicast_addr)
.or_insert_with(|| Arc::new(RwLock::new(Multicast::new())))
.value()
.clone();
let mut guard = multi.write();
match group_record.record_type() {
IgmpV3RecordType::ModeIsInclude
| IgmpV3RecordType::ChangeToIncludeMode => {
match group_record.source_addresses() {
None => {
//不接收所有
guard.members.remove(&source);
guard.map.remove(&source);
}
Some(src) => {
guard.members.insert(source, Instant::now());
guard.map.insert(source, (true, HashSet::from_iter(src)));
}
}
}
IgmpV3RecordType::ModeIsExclude
| IgmpV3RecordType::ChangeToExcludeMode => {
match group_record.source_addresses() {
None => {
//接收所有
guard.members.insert(source, Instant::now());
guard.map.remove(&source);
}
Some(src) => {
guard.members.insert(source, Instant::now());
guard.map.insert(source, (false, HashSet::from_iter(src)));
}
}
}
IgmpV3RecordType::AllowNewSources => {
//在已有源的基础上,接收目标源,如果是排除模式,则删除;是包含模式则添加
match group_record.source_addresses() {
None => {}
Some(src) => match guard.map.get_mut(&source) {
None => {}
Some((is_include, set)) => {
for ip in src {
if *is_include {
set.insert(ip);
} else {
set.remove(&ip);
}
}
}
},
}
}
IgmpV3RecordType::BlockOldSources => {
//在已有源的基础上,不接收目标源
match group_record.source_addresses() {
None => {}
Some(src) => match guard.map.get_mut(&source) {
None => {}
Some((is_include, set)) => {
for ip in src {
if *is_include {
set.remove(&ip);
} else {
set.insert(ip);
}
}
}
},
}
}
IgmpV3RecordType::Unknown(_) => {}
}
}
}
}
IgmpType::Unknown(_) => {}
}
Ok(())
}
}
+218 -125
View File
@@ -1,150 +1,243 @@
use crossbeam_utils::atomic::AtomicCell;
use dashmap::DashMap;
use std::io;
use std::mem::MaybeUninit;
use std::net::{IpAddr, Ipv4Addr, SocketAddrV4};
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::Arc;
use std::{io, thread};
use socket2::{Domain, SockAddr, Socket, Type};
use crossbeam_utils::atomic::AtomicCell;
use mio::net::UdpSocket;
use mio::{Events, Interest, Poll, Token, Waker};
use parking_lot::Mutex;
use crate::channel::sender::ChannelSender;
use crate::cipher::Cipher;
use crate::handle::CurrentDeviceInfo;
use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::{NetPacket, Protocol, Version, MAX_TTL};
use packet::icmp::icmp;
use packet::icmp::icmp::HeaderOther;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::CurrentDeviceInfo;
use crate::ip_proxy::ProxyHandler;
use crate::protocol;
use crate::protocol::{NetPacket, Version, MAX_TTL};
use crate::util::StopManager;
#[derive(Clone)]
pub struct IcmpProxy {
icmp_socket: Arc<Socket>,
icmp_socket: Arc<std::net::UdpSocket>,
// 对端-> 真实来源
icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
sender: ChannelSender,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
nat_map: Arc<Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>>,
}
impl IcmpProxy {
pub fn new(
addr: SocketAddrV4,
icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
sender: ChannelSender,
context: Context,
stop_manager: StopManager,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<IcmpProxy> {
let icmp_socket = Arc::new(Socket::new(
Domain::IPV4,
Type::RAW,
) -> io::Result<Self> {
let icmp_socket = socket2::Socket::new(
socket2::Domain::IPV4,
socket2::Type::RAW,
Some(socket2::Protocol::ICMPV4),
)?);
icmp_socket.bind(&SockAddr::from(addr))?;
Ok(IcmpProxy {
icmp_socket,
icmp_proxy_map,
sender,
current_device,
client_cipher,
)?;
let addr: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
icmp_socket.bind(&socket2::SockAddr::from(addr))?;
icmp_socket.set_nonblocking(true)?;
let std_socket: std::net::UdpSocket = icmp_socket.into();
let mio_icmp_socket = UdpSocket::from_std(std_socket.try_clone()?);
let nat_map: Arc<Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>> =
Arc::new(Mutex::new(HashMap::with_capacity(16)));
{
let nat_map = nat_map.clone();
thread::Builder::new()
.name("icmpProxy".into())
.spawn(move || {
if let Err(e) = icmp_proxy(
mio_icmp_socket,
nat_map,
context,
stop_manager,
current_device,
client_cipher,
) {
log::warn!("icmp_proxy:{:?}", e);
}
})
.expect("icmpProxy");
}
Ok(Self {
icmp_socket: Arc::new(std_socket),
nat_map,
})
}
pub fn icmp_socket(&self) -> Arc<Socket> {
self.icmp_socket.clone()
}
pub fn start(self) {
let mut buf = [0 as u8; 1500];
let data: &mut [MaybeUninit<u8>] = unsafe { std::mem::transmute(&mut buf[..]) };
}
loop {
match self.recv(data) {
Ok((len, peer_ip)) => {
match peer_ip {
IpAddr::V4(peer_ip) => {
match ipv4::packet::IpV4Packet::new(&mut buf[..len]) {
Ok(mut ipv4_packet) => {
match icmp::IcmpPacket::new(ipv4_packet.payload()) {
Ok(icmp_packet) => {
match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => {
if let Some(entry) =
self.icmp_proxy_map.get(&(peer_ip, id, seq))
{
//将数据发送到真实的来源
let dest_ip = *entry.value();
drop(entry);
ipv4_packet.set_destination_ip(dest_ip);
ipv4_packet.update_checksum();
let current_device =
self.current_device.load();
let virtual_ip =
current_device.virtual_ip();
let connect_server =
current_device.connect_server;
let mut net_packet =
NetPacket::new_encrypt(vec![
0u8;
12 + len + ENCRYPTION_RESERVED
])
.unwrap();
net_packet.set_version(Version::V1);
net_packet.set_protocol(Protocol::IpTurn);
net_packet.set_transport_protocol(crate::protocol::ip_turn_packet::Protocol::Ipv4.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip);
net_packet.set_destination(dest_ip);
net_packet
.set_payload(ipv4_packet.buffer)
.unwrap();
if let Err(e) = self
.client_cipher
.encrypt_ipv4(&mut net_packet)
{
log::warn!("加密失败:{}", e);
continue;
}
if self
.sender
.try_send_by_id(
net_packet.buffer(),
&dest_ip,
)
.is_err()
{
let _ = self.sender.send_main(
net_packet.buffer(),
connect_server,
);
}
}
}
_ => {
continue;
}
}
}
Err(_) => {}
};
}
Err(_) => {}
}
}
IpAddr::V6(_) => {}
}
}
Err(e) => {
log::warn!("icmp代理异常:{:?}", e);
const SERVER_VAL: usize = 0;
const SERVER: Token = Token(SERVER_VAL);
const NOTIFY_VAL: usize = 1;
const NOTIFY: Token = Token(NOTIFY_VAL);
fn icmp_proxy(
mut icmp_socket: UdpSocket,
// 对端-> 真实来源
nat_map: Arc<Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>>,
context: Context,
stop_manager: StopManager,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<()> {
let mut poll = Poll::new()?;
poll.registry()
.register(&mut icmp_socket, SERVER, Interest::READABLE)?;
let mut events = Events::with_capacity(32);
let stop = Arc::new(Waker::new(poll.registry(), NOTIFY)?);
let _stop = stop.clone();
let _worker = stop_manager.add_listener("icmp_proxy".into(), move || {
if let Err(e) = stop.wake() {
log::warn!("stop icmp_proxy:{:?}", e);
}
})?;
let mut buf = [0u8; 65535 - 20 - 8];
loop {
poll.poll(&mut events, None)?;
if stop_manager.is_stop() {
return Ok(());
}
for event in events.iter() {
match event.token() {
SERVER => readable_handle(
&icmp_socket,
&mut buf,
&nat_map,
&context,
&current_device,
&client_cipher,
),
NOTIFY => {
return Ok(());
}
_ => {}
}
}
}
fn recv(&self, buf: &mut [MaybeUninit<u8>]) -> io::Result<(usize, IpAddr)> {
let (size, addr) = self.icmp_socket.recv_from(buf)?;
let addr = match addr.as_socket() {
None => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
Some(add) => add.ip(),
}
fn readable_handle(
icmp_socket: &UdpSocket,
buf: &mut [u8],
nat_map: &Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
client_cipher: &Cipher,
) {
loop {
let (len, addr) = match icmp_socket.recv_from(&mut buf[12..]) {
Ok(rs) => rs,
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::warn!("icmp_socket {:?}", e);
return;
}
};
Ok((size, addr))
if let IpAddr::V4(peer_ip) = addr.ip() {
recv_handle(
buf,
12 + len,
peer_ip,
&nat_map,
&context,
&current_device,
&client_cipher,
);
}
}
// fn send_to(&self, buf: &[u8], addr: SocketAddrV4) -> io::Result<usize> {
// self.icmp_socket.send_to(buf, &SockAddr::from(addr))
// }
}
fn recv_handle(
buf: &mut [u8],
data_len: usize,
peer_ip: Ipv4Addr,
nat_map: &Mutex<HashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
context: &Context,
current_device: &AtomicCell<CurrentDeviceInfo>,
client_cipher: &Cipher,
) {
match IpV4Packet::new(&mut buf[12..data_len]) {
Ok(mut ipv4_packet) => match icmp::IcmpPacket::new(ipv4_packet.payload()) {
Ok(icmp_packet) => match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => {
if let Some(dest_ip) = nat_map.lock().get(&(peer_ip, id, seq)).cloned() {
ipv4_packet.set_destination_ip(dest_ip);
ipv4_packet.update_checksum();
let current_device = current_device.load();
let virtual_ip = current_device.virtual_ip();
let mut net_packet = NetPacket::new0(data_len, buf).unwrap();
net_packet.set_version(Version::V1);
net_packet.set_protocol(protocol::Protocol::IpTurn);
net_packet.set_transport_protocol(
protocol::ip_turn_packet::Protocol::Ipv4.into(),
);
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_source(virtual_ip);
net_packet.set_destination(dest_ip);
if let Err(e) = client_cipher.encrypt_ipv4(&mut net_packet) {
log::warn!("加密失败:{}", e);
return;
}
if let Err(e) = context.send_ipv4_by_id(
net_packet.buffer(),
&dest_ip,
current_device.connect_server,
current_device.status.online(),
) {
log::warn!("发送到目标失败:{}", e);
}
}
}
_ => {}
},
Err(_) => {}
},
Err(_) => {}
}
}
/// icmp用Identifier来区分,没有Identifier的一律不转发
impl ProxyHandler for IcmpProxy {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
source: Ipv4Addr,
destination: Ipv4Addr,
) -> io::Result<bool> {
if ipv4.offset() != 0 || ipv4.flags() & 1 == 1 {
// ip分片的直接丢弃
return Ok(true);
}
let dest_ip = ipv4.destination_ip();
//转发到代理目标地址
let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?;
match icmp_packet.header_other() {
HeaderOther::Identifier(id, seq) => {
self.nat_map.lock().insert((dest_ip, id, seq), source);
self.icmp_socket.send_to(
ipv4.payload(),
SocketAddr::from(SocketAddrV4::new(dest_ip, 0)),
)?;
}
_ => {
log::warn!(
"不支持的ip代理Icmp协议:{}->{}->{}",
source,
destination,
dest_ip
);
}
}
Ok(true)
}
fn send_handle(&self, _ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()> {
Ok(())
}
}
+74 -106
View File
@@ -1,123 +1,91 @@
use std::io;
use std::net::Ipv4Addr;
use std::sync::Arc;
use crossbeam_utils::atomic::AtomicCell;
use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet;
use crate::channel::context::Context;
use crate::cipher::Cipher;
use crate::handle::CurrentDeviceInfo;
use crate::ip_proxy::icmp_proxy::IcmpProxy;
use crate::ip_proxy::tcp_proxy::TcpProxy;
use crate::ip_proxy::udp_proxy::UdpProxy;
use dashmap::DashMap;
#[cfg(not(target_os = "android"))]
use socket2::{SockAddr, Socket};
#[cfg(not(target_os = "android"))]
use std::net::Ipv4Addr;
use std::net::SocketAddrV4;
use std::sync::Arc;
use std::{io, thread};
use tokio::net::{TcpListener, UdpSocket};
use crate::util::{Scheduler, StopManager};
#[cfg(not(target_os = "android"))]
pub mod icmp_proxy;
pub mod tcp_proxy;
pub mod udp_proxy;
pub trait DashMapNew {
fn new0() -> Self;
fn new_cap(capacity: usize) -> Self;
}
impl<'a, K: 'a + Eq + std::hash::Hash, V: 'a> DashMapNew for DashMap<K, V> {
fn new0() -> Self {
Self::new_cap(0)
}
fn new_cap(capacity: usize) -> Self {
let shard_amount = (thread::available_parallelism().map_or(4, |v| {
// https://github.com/rust-lang/rust/issues/115868
let n: usize = v.get() * 4;
if n == 0 {
log::warn!("available_parallelism=0");
println!("warn available_parallelism=0");
}
if n < 4 {
return 4;
}
n
}))
.next_power_of_two();
DashMap::with_capacity_and_shard_amount(capacity, shard_amount)
}
}
#[derive(Eq, PartialEq, Ord, PartialOrd, Copy, Clone, Debug)]
pub enum Protocol {
Icmp,
Tcp,
Udp,
pub trait ProxyHandler {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
source: Ipv4Addr,
destination: Ipv4Addr,
) -> io::Result<bool>;
fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()>;
}
#[derive(Clone)]
pub struct IpProxyMap {
pub(crate) tcp_proxy_port: u16,
pub(crate) udp_proxy_port: u16,
//真实源地址 -> 目的地址
pub(crate) tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
pub(crate) udp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
// icmp用Identifier来区分,没有Identifier的一律不转发
#[cfg(not(target_os = "android"))]
pub(crate) icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>>,
#[cfg(not(target_os = "android"))]
icmp_socket: Arc<Socket>,
icmp_proxy: IcmpProxy,
tcp_proxy: TcpProxy,
udp_proxy: UdpProxy,
}
impl IpProxyMap {
#[cfg(not(target_os = "android"))]
pub fn send_icmp(&self, buf: &[u8], dest: &Ipv4Addr) -> io::Result<usize> {
self.icmp_socket
.send_to(buf, &SockAddr::from(SocketAddrV4::new(*dest, 0)))
}
}
pub fn init_proxy(
context: Context,
scheduler: Scheduler,
stop_manager: StopManager,
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
client_cipher: Cipher,
) -> io::Result<IpProxyMap> {
let icmp_proxy = IcmpProxy::new(context, stop_manager.clone(), current_device, client_cipher)?;
let tcp_proxy = TcpProxy::new(stop_manager.clone())?;
let udp_proxy = UdpProxy::new(scheduler, stop_manager)?;
pub async fn init_proxy(
#[cfg(not(target_os = "android"))] sender: crate::channel::sender::ChannelSender,
#[cfg(not(target_os = "android"))] current_device: Arc<
crossbeam_utils::atomic::AtomicCell<crate::handle::CurrentDeviceInfo>,
>,
#[cfg(not(target_os = "android"))] client_cipher: crate::cipher::Cipher,
) -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> {
let tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new0());
let udp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>> = Arc::new(DashMap::new0());
#[cfg(not(target_os = "android"))]
let icmp_proxy_map: Arc<DashMap<(Ipv4Addr, u16, u16), Ipv4Addr>> = Arc::new(DashMap::new0());
let tcp_listener = TcpListener::bind("0.0.0.0:0").await?;
let udp_socket = UdpSocket::bind("0.0.0.0:0").await?;
let tcp_proxy_port = tcp_listener.local_addr()?.port();
let udp_proxy_port = udp_socket.local_addr()?.port();
let tcp_proxy = TcpProxy::new(tcp_listener, tcp_proxy_map.clone());
let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone());
#[cfg(not(target_os = "android"))]
let icmp_socket = {
let addr = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
let icmp_proxy = icmp_proxy::IcmpProxy::new(
addr,
icmp_proxy_map.clone(),
sender.clone(),
current_device.clone(),
client_cipher,
)?;
let icmp_socket = icmp_proxy.icmp_socket();
thread::spawn(move || {
icmp_proxy.start();
});
icmp_socket
};
Ok((
Ok(IpProxyMap {
icmp_proxy,
tcp_proxy,
udp_proxy,
IpProxyMap {
tcp_proxy_port,
udp_proxy_port,
tcp_proxy_map,
udp_proxy_map,
#[cfg(not(target_os = "android"))]
icmp_proxy_map,
#[cfg(not(target_os = "android"))]
icmp_socket,
},
))
})
}
impl ProxyHandler for IpProxyMap {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
source: Ipv4Addr,
destination: Ipv4Addr,
) -> io::Result<bool> {
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),
ipv4::protocol::Protocol::Icmp => {
self.icmp_proxy.recv_handle(ipv4, source, destination)
}
_ => {
log::warn!(
"不支持的ip代理ipv4协议{:?}:{}->{}->{}",
ipv4.protocol(),
source,
destination,
ipv4.destination_ip()
);
Ok(false)
}
}
}
fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()> {
match ipv4.protocol() {
ipv4::protocol::Protocol::Tcp => self.tcp_proxy.send_handle(ipv4),
ipv4::protocol::Protocol::Udp => self.udp_proxy.send_handle(ipv4),
ipv4::protocol::Protocol::Icmp => self.icmp_proxy.send_handle(ipv4),
_ => Ok(()),
}
}
}
+433 -69
View File
@@ -1,93 +1,457 @@
use dashmap::DashMap;
use std::io;
use std::net::{SocketAddr, SocketAddrV4};
use std::io::{Read, Write};
use std::net::{Ipv4Addr, Shutdown, SocketAddrV4};
#[cfg(unix)]
use std::os::fd::AsRawFd;
#[cfg(windows)]
use std::os::windows::io::AsRawSocket;
use std::sync::Arc;
use std::time::Duration;
use std::{collections::HashMap, io, net::SocketAddr, thread};
use tokio::net::{TcpListener, TcpStream};
use bytes::{BufMut, BytesMut};
use mio::net::TcpStream;
use mio::{net::TcpListener, Events, Interest, Poll, Registry, Token, Waker};
use parking_lot::Mutex;
use packet::ip::ipv4::packet::IpV4Packet;
use packet::tcp::tcp::TcpPacket;
use crate::ip_proxy::ProxyHandler;
use crate::util::StopManager;
const SERVER_VAL: usize = 0;
const SERVER: Token = Token(SERVER_VAL);
const NOTIFY_VAL: usize = 1;
const NOTIFY: Token = Token(NOTIFY_VAL);
#[derive(Clone)]
pub struct TcpProxy {
tcp_listener: TcpListener,
tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
port: u16,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
}
impl TcpProxy {
pub fn new(
tcp_listener: TcpListener,
tcp_proxy_map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
) -> Self {
Self {
tcp_listener,
tcp_proxy_map,
pub fn new(stop_manager: StopManager) -> io::Result<Self> {
let nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>> =
Arc::new(Mutex::new(HashMap::with_capacity(16)));
let tcp_listener = TcpListener::bind(format!("0.0.0.0:{}", 0).parse().unwrap())?;
let port = tcp_listener.local_addr()?.port();
{
let nat_map = nat_map.clone();
thread::Builder::new()
.name("tcpProxy".into())
.spawn(move || {
if let Err(e) = tcp_proxy(tcp_listener, nat_map, stop_manager) {
log::warn!("tcp_proxy:{:?}", e);
}
})
.expect("tcpProxy");
}
Ok(Self { port, nat_map })
}
pub async fn start(self) {
let tcp_listener = self.tcp_listener;
let tcp_proxy_map = self.tcp_proxy_map;
loop {
match tcp_listener.accept().await {
Ok((tcp_stream, sender_addr)) => match sender_addr {
SocketAddr::V4(sender_addr) => {
if let Some(entry) = tcp_proxy_map.get(&sender_addr) {
let dest_addr = *entry.value();
drop(entry);
}
tokio::spawn(async move {
let peer_tcp_stream = match tokio::time::timeout(
Duration::from_secs(5),
TcpStream::connect(dest_addr),
)
.await
{
Ok(peer_tcp_stream) => match peer_tcp_stream {
Ok(peer_tcp_stream) => peer_tcp_stream,
Err(e) => {
log::warn!(
"tcp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
return;
}
},
Err(e) => {
log::warn!(
"tcp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
return;
}
};
if let Err(e) = proxy(tcp_stream, peer_tcp_stream).await {
log::warn!("{}->{},{}", sender_addr, dest_addr, e);
}
});
impl ProxyHandler for TcpProxy {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
source: Ipv4Addr,
destination: Ipv4Addr,
) -> io::Result<bool> {
let dest_ip = ipv4.destination_ip();
//转发到代理目标地址
let mut tcp_packet = TcpPacket::new(source, destination, ipv4.payload_mut())?;
let source_port = tcp_packet.source_port();
let dest_port = tcp_packet.destination_port();
tcp_packet.set_destination_port(self.port);
tcp_packet.update_checksum();
ipv4.set_destination_ip(destination);
ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port);
self.nat_map
.lock()
.insert(key, SocketAddrV4::new(dest_ip, dest_port));
Ok(false)
}
fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()> {
let src_ip = ipv4.source_ip();
let dest_ip = ipv4.destination_ip();
let dest_addr = {
let tcp_packet = TcpPacket::new(src_ip, dest_ip, ipv4.payload_mut())?;
SocketAddrV4::new(dest_ip, tcp_packet.destination_port())
};
if let Some(source_addr) = self.nat_map.lock().get(&dest_addr) {
let source_ip = *source_addr.ip();
let mut tcp_packet = TcpPacket::new(source_ip, dest_ip, ipv4.payload_mut())?;
tcp_packet.set_source_port(source_addr.port());
tcp_packet.update_checksum();
ipv4.set_source_ip(source_ip);
ipv4.update_checksum();
}
Ok(())
}
}
fn tcp_proxy(
mut tcp_listener: TcpListener,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
stop_manager: StopManager,
) -> io::Result<()> {
let mut poll = Poll::new()?;
poll.registry()
.register(&mut tcp_listener, SERVER, Interest::READABLE)?;
let mut events = Events::with_capacity(32);
let mut tcp_map: HashMap<usize, ProxyValue> = HashMap::with_capacity(16);
let mut mapping: HashMap<usize, usize> = HashMap::with_capacity(16);
let stop = Arc::new(Waker::new(poll.registry(), NOTIFY)?);
let _stop = stop.clone();
let _worker = stop_manager.add_listener("tcp_proxy".into(), move || {
if let Err(e) = stop.wake() {
log::warn!("stop tcp_proxy:{:?}", e);
}
})?;
loop {
poll.poll(&mut events, None)?;
if stop_manager.is_stop() {
return Ok(());
}
for event in events.iter() {
match event.token() {
SERVER => {
accept_handle(
poll.registry(),
&tcp_listener,
&nat_map,
&mut tcp_map,
&mut mapping,
);
}
NOTIFY => {
return Ok(());
}
Token(index) => {
let (val, src_index) = if let Some(v) = tcp_map.get_mut(&index) {
(v, index)
} else {
if let Some(dest_index) = mapping.get(&index) {
if let Some(v) = tcp_map.get_mut(dest_index) {
(v, *dest_index)
} else {
continue;
}
} else {
log::warn!("tcp代理异常: 来源:{},未找到目标", sender_addr);
continue;
}
};
let (stream1, stream2, buf1, buf2, state1, state2) = val.as_mut(index);
if event.is_readable() {
if let Err(_) = readable_handle(stream1, stream2, buf1, state2) {
*state1 |= READ_CLOSED;
}
}
SocketAddr::V6(_) => {}
},
Err(e) => {
log::warn!("tcp代理监听:{:?}", e);
if event.is_writable() {
let read = buf2.len() >= BUF_LEN;
if let Err(_) = writable_handle(stream1, buf2) {
*state1 |= WRITE_CLOSED;
} else if read {
if readable_handle(stream2, stream1, buf2, state1).is_err() {
*state2 |= READ_CLOSED;
}
}
}
if event.is_read_closed() || event.is_error() {
*state1 |= READ_CLOSED;
}
if event.is_write_closed() || event.is_error() {
*state1 |= WRITE_CLOSED;
}
if is_write_closed(*state1) {
let _ = stream1.shutdown(Shutdown::Write);
let _ = stream2.shutdown(Shutdown::Read);
}
if is_read_closed(*state1) {
let _ = stream1.shutdown(Shutdown::Read);
if buf1.is_empty() {
let _ = stream2.shutdown(Shutdown::Write);
}
}
if (is_both_closed(*state1) && buf1.is_empty())
|| (is_both_closed(*state2) && buf2.is_empty())
|| (is_write_closed(*state1) && is_write_closed(*state2)
|| (is_read_closed(*state1)
&& is_read_closed(*state2)
&& buf1.is_empty()
&& buf2.is_empty()))
{
close(src_index, &mut tcp_map, &mut mapping);
}
}
}
}
}
}
async fn proxy(mut client: TcpStream, mut server: TcpStream) -> io::Result<()> {
let (mut client_reader, mut client_writer) = client.split();
let (mut server_reader, mut server_writer) = server.split();
fn accept_handle(
registry: &Registry,
tcp_listener: &TcpListener,
nat_map: &Mutex<HashMap<SocketAddrV4, SocketAddrV4>>,
tcp_map: &mut HashMap<usize, ProxyValue>,
mapping: &mut HashMap<usize, usize>,
) {
loop {
match tcp_listener.accept() {
Ok((mut src_stream, addr)) => {
#[cfg(windows)]
let src_fd = src_stream.as_raw_socket() as usize;
#[cfg(unix)]
let src_fd = src_stream.as_raw_fd() as usize;
if src_fd == SERVER_VAL || src_fd == NOTIFY_VAL {
log::error!("fd错误:{:?}", src_fd);
continue;
}
let addr = match addr {
SocketAddr::V4(addr) => addr,
SocketAddr::V6(_) => {
// 忽略ipv6
continue;
}
};
let _ = src_stream.set_nodelay(false);
if let Some(dest_addr) = nat_map.lock().get(&addr).cloned() {
match tcp_connect(addr.port(), dest_addr.into()) {
Ok(mut dest_stream) => {
#[cfg(windows)]
let dest_fd = dest_stream.as_raw_socket() as usize;
#[cfg(unix)]
let dest_fd = dest_stream.as_raw_fd() as usize;
if dest_fd == SERVER_VAL || dest_fd == NOTIFY_VAL {
log::error!("fd错误:{:?}", dest_fd);
continue;
}
if let Err(e) = registry.register(
&mut src_stream,
Token(src_fd),
Interest::READABLE.add(Interest::WRITABLE),
) {
log::error!("register src_stream:{:?}", e);
continue;
}
if let Err(e) = registry.register(
&mut dest_stream,
Token(dest_fd),
Interest::READABLE.add(Interest::WRITABLE),
) {
log::error!("register dest_stream:{:?}", e);
continue;
}
tcp_map.insert(
src_fd,
ProxyValue::new(src_stream, dest_stream, src_fd, dest_fd),
);
mapping.insert(dest_fd, src_fd);
}
Err(e) => {
log::error!("connect:{:?} {}->{}", e, addr, dest_addr);
}
}
}
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::error!("accept:{:?}", e);
}
}
}
}
let client_to_server = tokio::io::copy(&mut client_reader, &mut server_writer);
let server_to_client = tokio::io::copy(&mut server_reader, &mut client_writer);
tokio::select! {
_ = tokio::time::timeout(Duration::from_secs(10), client_to_server) =>{},
_ = tokio::time::timeout(Duration::from_secs(10), server_to_client) =>{},
fn tcp_connect(src_port: u16, addr: SocketAddr) -> io::Result<TcpStream> {
let socket = socket2::Socket::new(
socket2::Domain::IPV4,
socket2::Type::STREAM,
Some(socket2::Protocol::TCP),
)?;
if socket
.bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into())
.is_err()
{
socket.bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?;
}
if let Err(e) = socket.set_tcp_keepalive(
&socket2::TcpKeepalive::new()
.with_time(Duration::from_secs(120))
.with_interval(Duration::from_secs(10)),
) {
log::warn!("set_tcp_keepalive err {:?}", e);
}
let _ = socket.set_nodelay(false);
socket.connect_timeout(&addr.into(), Duration::from_secs(3))?;
socket.set_nonblocking(true)?;
Ok(TcpStream::from_std(socket.into()))
}
#[derive(Debug)]
struct ProxyValue {
src_stream: TcpStream,
dest_stream: TcpStream,
src_fd: usize,
dest_fd: usize,
src_buf: BytesMut,
dest_buf: BytesMut,
src_state: u8,
dest_state: u8,
}
const BUF_LEN: usize = 65536;
impl ProxyValue {
fn new(src_stream: TcpStream, dest_stream: TcpStream, src_fd: usize, dest_fd: usize) -> Self {
Self {
src_stream,
dest_stream,
src_fd,
dest_fd,
src_buf: BytesMut::with_capacity(BUF_LEN),
dest_buf: BytesMut::with_capacity(BUF_LEN),
src_state: NORMAL,
dest_state: NORMAL,
}
}
fn as_mut(
&mut self,
index: usize,
) -> (
&mut TcpStream,
&mut TcpStream,
&mut BytesMut,
&mut BytesMut,
&mut u8,
&mut u8,
) {
if index == self.src_fd {
(
&mut self.src_stream,
&mut self.dest_stream,
&mut self.src_buf,
&mut self.dest_buf,
&mut self.src_state,
&mut self.dest_state,
)
} else {
(
&mut self.dest_stream,
&mut self.src_stream,
&mut self.dest_buf,
&mut self.src_buf,
&mut self.dest_state,
&mut self.src_state,
)
}
}
}
fn readable_handle(
stream1: &mut TcpStream,
stream2: &mut TcpStream,
mid_buf: &mut BytesMut,
state2: &mut u8,
) -> io::Result<()> {
let mut buf = [0; BUF_LEN];
loop {
if mid_buf.len() >= BUF_LEN {
// 达到上限不再继续读取
return Ok(());
}
match stream1.read(&mut buf) {
Ok(len) => {
if len == 0 {
return Err(io::Error::from(io::ErrorKind::UnexpectedEof));
}
let mut buf = &buf[..len];
if mid_buf.is_empty() {
// 直接写入,避免在buf中过渡
while !buf.is_empty() {
match stream2.write(buf) {
Ok(end) => {
if end == 0 {
*state2 |= WRITE_CLOSED;
return Err(io::Error::from(io::ErrorKind::WriteZero));
}
buf = &buf[end..];
}
Err(e) => {
if e.kind() != io::ErrorKind::WouldBlock {
*state2 |= WRITE_CLOSED;
return Err(e);
}
break;
}
}
}
if buf.is_empty() {
continue;
}
}
mid_buf.reserve(buf.len());
mid_buf.put_slice(buf);
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
}
}
Ok(())
}
fn writable_handle(stream: &mut TcpStream, mid_buf: &mut BytesMut) -> io::Result<()> {
while !mid_buf.is_empty() {
match stream.write(&mid_buf) {
Ok(len) => {
let _ = mid_buf.split_to(len);
}
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
}
}
Ok(())
}
fn close(
index: usize,
tcp_map: &mut HashMap<usize, ProxyValue>,
mapping: &mut HashMap<usize, usize>,
) {
if let Some(val) = tcp_map.remove(&index) {
let _ = val.src_stream.shutdown(Shutdown::Both);
let _ = val.dest_stream.shutdown(Shutdown::Both);
mapping.remove(&val.src_fd);
mapping.remove(&val.dest_fd);
}
}
const NORMAL: u8 = 0b00;
const READ_CLOSED: u8 = 0b01;
const WRITE_CLOSED: u8 = 0b10;
const BOTH_CLOSED: u8 = 0b11;
fn is_read_closed(state: u8) -> bool {
(state & READ_CLOSED == READ_CLOSED) || is_both_closed(state)
}
fn is_write_closed(state: u8) -> bool {
(state & WRITE_CLOSED == WRITE_CLOSED) || is_both_closed(state)
}
fn is_both_closed(state: u8) -> bool {
state & BOTH_CLOSED == BOTH_CLOSED
}
+286 -86
View File
@@ -1,112 +1,312 @@
use crate::ip_proxy::DashMapNew;
use dashmap::DashMap;
use std::io;
use std::net::{SocketAddr, SocketAddrV4};
use std::net::{Ipv4Addr, SocketAddrV4};
#[cfg(unix)]
use std::os::fd::AsRawFd;
#[cfg(windows)]
use std::os::windows::io::AsRawSocket;
use std::sync::Arc;
use std::time::Duration;
use tokio::net::UdpSocket;
use std::time::{Duration, Instant};
use std::{collections::HashMap, io, net::SocketAddr, rc::Rc, thread};
/// 一个udp代理,作用是利用系统协议栈,将udp数据报解析出来再转发到目的地址
use mio::{net::UdpSocket, Events, Interest, Poll, Token};
use mio::{Registry, Waker};
use parking_lot::Mutex;
use packet::ip::ipv4::packet::IpV4Packet;
use packet::udp::udp::UdpPacket;
use crate::ip_proxy::ProxyHandler;
use crate::util::{Scheduler, StopManager};
const SERVER_VAL: usize = 0;
const SERVER: Token = Token(SERVER_VAL);
const NOTIFY_VAL: usize = 1;
const NOTIFY: Token = Token(NOTIFY_VAL);
// 开了ip代理后使用mstsc,mstsc会误以为在真实局域网,从而不维护udp心跳,导致断连,所以这里尽量长一点过期时间
const NAT_TIMEOUT: Duration = Duration::from_secs(20 * 60);
const NAT_FAST_TIMEOUT: Duration = Duration::from_secs(5 * 60);
const NAT_MAX: usize = 5_000;
#[derive(Clone)]
pub struct UdpProxy {
udp_socket: Arc<UdpSocket>,
map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
port: u16,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
}
impl UdpProxy {
pub fn new(udp_socket: UdpSocket, map: Arc<DashMap<SocketAddrV4, SocketAddrV4>>) -> Self {
let udp_socket = Arc::new(udp_socket);
Self { udp_socket, map }
pub fn new(scheduler: Scheduler, stop_manager: StopManager) -> io::Result<Self> {
let nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>> =
Arc::new(Mutex::new(HashMap::with_capacity(16)));
let udp = UdpSocket::bind(format!("0.0.0.0:{}", 0).parse().unwrap())?;
let port = udp.local_addr()?.port();
{
let nat_map = nat_map.clone();
thread::Builder::new()
.name("udpProxy".into())
.spawn(move || {
if let Err(e) = udp_proxy(udp, nat_map, scheduler, stop_manager) {
log::warn!("udp_proxy:{:?}", e);
}
})
.expect("udpProxy");
}
Ok(Self { port, nat_map })
}
pub async fn start(self) {
let map = self.map;
let udp_socket = self.udp_socket;
let mut buf = [0u8; 65536];
}
let inner_map: Arc<DashMap<SocketAddrV4, Arc<UdpSocket>>> = Arc::new(DashMap::new0());
impl ProxyHandler for UdpProxy {
fn recv_handle(
&self,
ipv4: &mut IpV4Packet<&mut [u8]>,
source: Ipv4Addr,
destination: Ipv4Addr,
) -> io::Result<bool> {
let dest_ip = ipv4.destination_ip();
//转发到代理目标地址
let mut udp_packet = UdpPacket::new(source, destination, ipv4.payload_mut())?;
let source_port = udp_packet.source_port();
let dest_port = udp_packet.destination_port();
udp_packet.set_destination_port(self.port);
udp_packet.update_checksum();
ipv4.set_destination_ip(destination);
ipv4.update_checksum();
let key = SocketAddrV4::new(source, source_port);
self.nat_map
.lock()
.insert(key.into(), SocketAddrV4::new(dest_ip, dest_port).into());
Ok(false)
}
loop {
match udp_socket.recv_from(&mut buf).await {
Ok((len, sender_addr)) => match sender_addr {
SocketAddr::V4(sender_addr) => {
match start0(&buf[..len], sender_addr, &inner_map, &map, &udp_socket).await
{
Ok(_) => {}
Err(e) => {
log::warn!("udp代理异常:{:?},来源:{}", e, sender_addr);
}
fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()> {
let src_ip = ipv4.source_ip();
let dest_ip = ipv4.destination_ip();
let dest_addr = {
let udp_packet = UdpPacket::new(src_ip, dest_ip, ipv4.payload_mut())?;
SocketAddrV4::new(dest_ip, udp_packet.destination_port())
};
if let Some(source_addr) = self.nat_map.lock().get(&dest_addr) {
let source_ip = *source_addr.ip();
let mut udp_packet = UdpPacket::new(source_ip, dest_ip, ipv4.payload_mut())?;
udp_packet.set_source_port(source_addr.port());
udp_packet.update_checksum();
ipv4.set_source_ip(source_ip);
ipv4.update_checksum();
}
Ok(())
}
}
fn udp_proxy(
mut udp: UdpSocket,
nat_map: Arc<Mutex<HashMap<SocketAddrV4, SocketAddrV4>>>,
scheduler: Scheduler,
stop_manager: StopManager,
) -> io::Result<()> {
let mut poll = Poll::new()?;
poll.registry()
.register(&mut udp, SERVER, Interest::READABLE)?;
let mut events = Events::with_capacity(32);
let mut buf = [0; 65536];
let mut token_map: HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)> =
HashMap::with_capacity(64);
let mut udp_map: HashMap<SocketAddrV4, (Rc<UdpSocket>, Instant)> = HashMap::with_capacity(64);
let mut timeout = false;
let waker = Arc::new(Waker::new(poll.registry(), NOTIFY)?);
let stop = waker.clone();
let _worker = stop_manager.add_listener("udp_proxy".into(), move || {
if let Err(e) = stop.wake() {
log::warn!("stop udp_proxy:{:?}", e);
}
})?;
loop {
let mut check = false;
if token_map.is_empty() {
poll.poll(&mut events, None)?;
} else {
//所有事件 50分钟超时
if let Err(e) = poll.poll(&mut events, Some(Duration::from_secs(50 * 60))) {
if e.kind() == io::ErrorKind::TimedOut || e.kind() == io::ErrorKind::WouldBlock {
token_map.clear();
udp_map.clear();
continue;
}
return Err(e);
}
}
if stop_manager.is_stop() {
return Ok(());
}
for event in events.iter() {
match event.token() {
SERVER => server_handle(
poll.registry(),
&udp,
&nat_map,
&mut token_map,
&mut udp_map,
&mut buf,
),
NOTIFY => {
check = true;
}
token => {
if let Err(e) = readable_handle(&udp, &mut token_map, &token, &mut buf) {
log::error!("发送目标失败:{:?}", e);
if let Some((_, src_addr, _)) = token_map.remove(&token) {
udp_map.remove(&src_addr);
}
}
SocketAddr::V6(_) => {}
},
}
}
}
if check {
//超时校验
if token_map.len() > NAT_MAX / 2 {
check_handle(&mut token_map, &mut udp_map, NAT_FAST_TIMEOUT)
} else {
check_handle(&mut token_map, &mut udp_map, NAT_TIMEOUT)
}
timeout = false;
}
if !token_map.is_empty() && !timeout {
//注册超时监听
timeout = true;
let waker = waker.clone();
scheduler.timeout(NAT_FAST_TIMEOUT, move |_| {
let _ = waker.wake();
});
}
}
}
fn check_handle(
token_map: &mut HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)>,
udp_map: &mut HashMap<SocketAddrV4, (Rc<UdpSocket>, Instant)>,
timeout: Duration,
) {
let mut remove_list = Vec::new();
for (token, (_, addr, time)) in token_map.iter() {
if time.elapsed() > timeout {
if let Some((_, time)) = udp_map.get(addr) {
if time.elapsed() > timeout {
//映射超时,需要移除
remove_list.push(*token);
}
}
}
}
for token in remove_list {
if let Some((_, src_addr, _)) = token_map.remove(&token) {
udp_map.remove(&src_addr);
}
}
}
fn server_handle(
registry: &Registry,
udp: &UdpSocket,
nat_map: &Mutex<HashMap<SocketAddrV4, SocketAddrV4>>,
token_map: &mut HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)>,
udp_map: &mut HashMap<SocketAddrV4, (Rc<UdpSocket>, Instant)>,
buf: &mut [u8],
) {
loop {
let (len, src_addr) = match udp.recv_from(buf) {
Ok((len, src_addr)) => match src_addr {
SocketAddr::V4(addr) => (len, addr),
SocketAddr::V6(_) => {
continue;
}
},
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
log::error!("接收数据失败:{:?}", e);
break;
}
};
if let Some((dest_udp, time)) = udp_map.get_mut(&src_addr) {
//发送失败就当丢包了
let _ = dest_udp.send(&buf[..len]);
*time = Instant::now();
} else if let Some(dest_addr) = nat_map.lock().get(&src_addr).cloned() {
if token_map.len() >= NAT_MAX {
log::error!(
"UDP NAT_MAX:src_addr={:?},dest_addr={:?}",
src_addr,
dest_addr
);
continue;
}
match udp_connect(src_addr.port(), dest_addr.into()) {
Ok((token_val, mut dest_udp)) => {
let token = Token(token_val);
if let Err(e) = registry.register(&mut dest_udp, token, Interest::READABLE) {
log::error!("register失败:{:?},addr={:?}", e, dest_addr);
continue;
}
if dest_udp.send(&buf[..len]).is_ok() {
let dest_udp = Rc::new(dest_udp);
token_map.insert(token, (dest_udp.clone(), src_addr, Instant::now()));
udp_map.insert(src_addr, (dest_udp, Instant::now()));
}
}
Err(e) => {
log::warn!("udp代理异常:{:?}", e);
log::error!("绑定目标地址失败:{:?}", e);
continue;
}
};
}
}
}
async fn start0(
buf: &[u8],
sender_addr: SocketAddrV4,
inner_map: &Arc<DashMap<SocketAddrV4, Arc<UdpSocket>>>,
map: &Arc<DashMap<SocketAddrV4, SocketAddrV4>>,
udp_socket: &Arc<UdpSocket>,
/// 得到一个 fd不为SERVER_VAL或者NOTYFY_VAL的socket
fn udp_connect(src_port: u16, addr: SocketAddr) -> io::Result<(usize, UdpSocket)> {
loop {
let udp = if let Ok(udp) =
UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into())
{
udp
} else {
UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?
};
#[cfg(windows)]
let fd = udp.as_raw_socket() as usize;
#[cfg(unix)]
let fd = udp.as_raw_fd() as usize;
if fd == SERVER_VAL || fd == NOTIFY_VAL {
continue;
}
// 只接收目标的数据
udp.connect(addr)?;
return Ok((fd, udp));
}
}
fn readable_handle(
udp: &UdpSocket,
token_map: &mut HashMap<Token, (Rc<UdpSocket>, SocketAddrV4, Instant)>,
token: &Token,
buf: &mut [u8],
) -> io::Result<()> {
if let Some(entry) = inner_map.get(&sender_addr) {
let udp = entry.value().clone();
drop(entry);
udp.send(buf).await?;
} else if let Some(entry) = map.get(&sender_addr) {
let dest_addr = *entry.value();
drop(entry);
let peer_udp_socket = UdpSocket::bind("0.0.0.0:0").await?;
peer_udp_socket.connect(dest_addr).await?;
peer_udp_socket.send(buf).await?;
let peer_udp_socket = Arc::new(peer_udp_socket);
let inner_map = inner_map.clone();
inner_map.insert(sender_addr, peer_udp_socket.clone());
let udp_socket = udp_socket.clone();
let map = map.clone();
tokio::spawn(async move {
let mut buf = [0u8; 65536];
loop {
match tokio::time::timeout(Duration::from_secs(300), peer_udp_socket.recv(&mut buf))
.await
{
Ok(rs) => match rs {
Ok(len) => match udp_socket.send_to(&buf[..len], sender_addr).await {
Ok(_) => {}
Err(e) => {
log::warn!(
"udp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
break;
}
},
Err(e) => {
log::warn!(
"udp代理异常:{:?},来源:{},目标:{}",
e,
sender_addr,
dest_addr
);
break;
}
},
Err(_) => {
//超时关闭
log::warn!("udp代理超时关闭,来源:{},目标:{}", sender_addr, dest_addr);
if let Some((dest_udp, src_addr, time)) = token_map.get_mut(&token) {
loop {
let len = match dest_udp.recv(buf) {
Ok(rs) => rs,
Err(e) => {
if e.kind() == io::ErrorKind::WouldBlock {
break;
}
return Err(e);
}
};
if len == 0 {
return Err(io::Error::from(io::ErrorKind::UnexpectedEof));
}
inner_map.remove(&sender_addr);
map.remove(&sender_addr);
});
let _ = udp.send_to(&buf[..len], (*src_addr).into());
}
*time = Instant::now();
}
Ok(())
}
+4 -5
View File
@@ -1,17 +1,16 @@
use crate::error::Error;
pub const VNT_VERSION: &'static str = "1.2.3";
pub type Result<T> = std::result::Result<T, Error>;
pub const VNT_VERSION: &'static str = env!("CARGO_PKG_VERSION");
pub mod channel;
pub mod cipher;
pub mod core;
pub mod error;
pub mod external_route;
pub mod handle;
pub mod igmp_server;
#[cfg(feature = "ip_proxy")]
pub mod ip_proxy;
pub mod nat;
pub mod proto;
pub mod protocol;
pub mod tun_tap_device;
pub mod util;
pub use handle::callback::{DeviceInfo, ErrorInfo, HandshakeInfo, RegisterInfo, VntCallback};
+64 -89
View File
@@ -1,16 +1,19 @@
use std::io;
use std::net::UdpSocket;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::ops::Sub;
use std::sync::Arc;
use std::time::{Duration, Instant};
use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use crate::channel::punch::{NatInfo, NatType};
use crate::proto::message::PunchNatType;
mod stun_test;
mod stun;
pub fn local_ipv4() -> io::Result<Ipv4Addr> {
pub fn local_ipv4_() -> io::Result<Ipv4Addr> {
let socket = UdpSocket::bind("0.0.0.0:0")?;
socket.connect("8.8.8.8:80")?;
let addr = socket.local_addr()?;
@@ -19,33 +22,31 @@ pub fn local_ipv4() -> io::Result<Ipv4Addr> {
IpAddr::V6(_) => Ok(Ipv4Addr::UNSPECIFIED),
}
}
pub fn local_ipv4() -> Option<Ipv4Addr> {
match local_ipv4_() {
Ok(ipv4) => Some(ipv4),
Err(e) => {
log::warn!("获取ipv4失败:{:?}", e);
None
}
}
}
pub fn local_ipv6() -> io::Result<Ipv6Addr> {
pub fn local_ipv6_() -> io::Result<Ipv6Addr> {
let socket = UdpSocket::bind("[::]:0")?;
socket.connect("[2001:4860:4860::8888]:80")?;
socket.connect("[2001:4860:4860:0000:0000:0000:0000:8888]:80")?;
let addr = socket.local_addr()?;
match addr.ip() {
IpAddr::V4(_) => Ok(Ipv6Addr::UNSPECIFIED),
IpAddr::V6(ip) => Ok(ip),
}
}
pub fn local_ipv4_addr(port: u16) -> SocketAddrV4 {
match local_ipv4() {
Ok(ipv4) => SocketAddrV4::new(ipv4, port),
pub fn local_ipv6() -> Option<Ipv6Addr> {
match local_ipv6_() {
Ok(ipv6) => Some(ipv6),
Err(e) => {
log::warn!("获取本地ipv4地址失败:{}", e);
SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0)
}
}
}
pub fn local_ipv6_addr(port: u16) -> SocketAddrV6 {
match local_ipv6() {
Ok(ipv6) => SocketAddrV6::new(ipv6, port, 0, 0),
Err(e) => {
log::warn!("获取本地ipv6地址失败:{}", e);
SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, 0, 0, 0)
log::warn!("获取ipv6失败:{:?}", e);
None
}
}
}
@@ -54,6 +55,7 @@ pub fn local_ipv6_addr(port: u16) -> SocketAddrV6 {
pub struct NatTest {
stun_server: Vec<String>,
info: Arc<Mutex<NatInfo>>,
time: Arc<AtomicCell<Instant>>,
}
impl From<NatType> for PunchNatType {
@@ -76,89 +78,62 @@ impl Into<NatType> for PunchNatType {
impl NatTest {
pub fn new(
channel_num: usize,
mut stun_server: Vec<String>,
public_ip: Ipv4Addr,
public_port: u16,
local_ipv4_addr: SocketAddrV4,
ipv6_addr: SocketAddrV6,
local_ipv4: Option<Ipv4Addr>,
ipv6: Option<Ipv6Addr>,
udp_ports: Vec<u16>,
tcp_port: u16,
) -> NatTest {
let server = stun_server[0].clone();
stun_server.resize(3, server);
let mut ports = udp_ports.clone();
ports.resize(channel_num, 0);
let nat_info = NatInfo::new(
vec![public_ip],
public_port,
Vec::new(),
ports,
0,
local_ipv4_addr,
ipv6_addr,
local_ipv4,
ipv6,
udp_ports,
tcp_port,
NatType::Cone,
);
let info = Arc::new(Mutex::new(nat_info));
NatTest { stun_server, info }
NatTest {
stun_server,
info,
time: Arc::new(AtomicCell::new(
Instant::now().sub(Duration::from_secs(100)),
)),
}
}
pub fn can_update(&self) -> bool {
let last = self.time.load();
last.elapsed() > Duration::from_secs(10)
&& self.time.compare_exchange(last, Instant::now()).is_ok()
}
pub fn nat_info(&self) -> NatInfo {
self.info.lock().clone()
}
pub fn update_addr(&self, ip: Ipv4Addr, port: u16) {
pub fn update_addr(&self, index: usize, ip: Ipv4Addr, port: u16) {
let mut guard = self.info.lock();
guard.public_port = port;
if !guard.public_ips.contains(&ip) {
guard.public_ips.push(ip);
}
guard.update_addr(index, ip, port)
}
pub async fn re_test(
pub fn re_test(
&self,
public_ip: Ipv4Addr,
public_port: u16,
local_ipv4_addr: SocketAddrV4,
ipv6_addr: SocketAddrV6,
) -> NatInfo {
let info = NatTest::re_test_(
&self.stun_server,
public_ip,
public_port,
local_ipv4_addr,
ipv6_addr,
)
.await;
*self.info.lock() = info.clone();
info
}
async fn re_test_(
stun_server: &Vec<String>,
public_ip: Ipv4Addr,
public_port: u16,
local_ipv4_addr: SocketAddrV4,
ipv6_addr: SocketAddrV6,
) -> NatInfo {
return match stun_test::stun_test_nat(stun_server.clone()).await {
Ok((nat_type, ips, port_range)) => {
let mut public_ips = Vec::new();
public_ips.push(Ipv4Addr::from(public_ip));
for ip in ips {
if ip != public_ip {
public_ips.push(ip);
}
}
NatInfo::new(
public_ips,
public_port,
port_range,
local_ipv4_addr,
ipv6_addr,
nat_type,
)
}
Err(e) => {
log::warn!("{:?}", e);
NatInfo::new(
vec![public_ip],
public_port,
0,
local_ipv4_addr,
ipv6_addr,
NatType::Cone,
)
}
};
local_ipv4: Option<Ipv4Addr>,
ipv6: Option<Ipv6Addr>,
) -> io::Result<NatInfo> {
let (nat_type, public_ips, port_range) = stun::stun_test_nat(self.stun_server.clone())?;
let mut guard = self.info.lock();
guard.nat_type = nat_type;
guard.public_ips = public_ips;
guard.public_port_range = port_range;
guard.local_ipv4 = local_ipv4;
guard.ipv6 = ipv6;
Ok(guard.clone())
}
}
@@ -1,23 +1,23 @@
use std::collections::HashSet;
use std::io;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6};
use std::time::Duration;
use std::{io, thread};
use crate::channel::punch::NatType;
use std::net::UdpSocket;
use stun_format::Attr;
use tokio::net::UdpSocket;
pub async fn stun_test_nat(stun_servers: Vec<String>) -> io::Result<(NatType, Vec<Ipv4Addr>, u16)> {
pub fn stun_test_nat(stun_servers: Vec<String>) -> io::Result<(NatType, Vec<Ipv4Addr>, u16)> {
let mut h = Vec::new();
for x in stun_servers {
let handle = tokio::spawn(test_nat(x));
let handle = thread::spawn(move || test_nat(x));
h.push(handle);
}
let mut nat_type = NatType::Cone;
let mut port_range = 0;
let mut hash_set = HashSet::new();
for x in h {
if let Ok(rs) = x.await {
if let Ok(rs) = x.join() {
if let Ok((nat_type_t, ip_list_t, port_range_t)) = rs {
if nat_type_t == NatType::Symmetric {
nat_type = NatType::Symmetric;
@@ -34,13 +34,14 @@ pub async fn stun_test_nat(stun_servers: Vec<String>) -> io::Result<(NatType, Ve
Ok((nat_type, hash_set.into_iter().collect(), port_range))
}
async fn test_nat(stun_server: String) -> io::Result<(NatType, Vec<Ipv4Addr>, u16)> {
let udp = UdpSocket::bind("0.0.0.0:0").await?;
udp.connect(stun_server).await?;
let mut nat_type = NatType::Cone;
fn test_nat(stun_server: String) -> io::Result<(NatType, Vec<Ipv4Addr>, u16)> {
let udp = UdpSocket::bind("0.0.0.0:0")?;
udp.set_read_timeout(Some(Duration::from_millis(300)))?;
udp.connect(stun_server)?;
let mut port_range = 0;
let mut hash_set = HashSet::new();
match test_nat_(&udp, true, true).await {
let mut nat_type = NatType::Cone;
match test_nat_(&udp, true, true) {
Ok((mapped_addr1, changed_addr1)) => {
match mapped_addr1.ip() {
IpAddr::V4(ip) => {
@@ -48,18 +49,18 @@ async fn test_nat(stun_server: String) -> io::Result<(NatType, Vec<Ipv4Addr>, u1
}
IpAddr::V6(_) => {}
}
if udp.connect(changed_addr1).await.is_ok() {
if let Ok((mapped_addr2, _)) = test_nat_(&udp, false, false).await {
if udp.connect(changed_addr1).is_ok() {
if let Ok((mapped_addr2, _)) = test_nat_(&udp, false, false) {
match mapped_addr2.ip() {
IpAddr::V4(ip) => {
hash_set.insert(ip);
if mapped_addr1 != mapped_addr2 {
nat_type = NatType::Symmetric;
}
}
IpAddr::V6(_) => {}
}
port_range = mapped_addr2.port().abs_diff(mapped_addr1.port());
if mapped_addr1 != mapped_addr2 {
nat_type = NatType::Symmetric;
}
}
}
}
@@ -68,7 +69,7 @@ async fn test_nat(stun_server: String) -> io::Result<(NatType, Vec<Ipv4Addr>, u1
Ok((nat_type, hash_set.into_iter().collect(), port_range))
}
async fn test_nat_(
fn test_nat_(
udp: &UdpSocket,
change_ip: bool,
change_port: bool,
@@ -83,15 +84,14 @@ async fn test_nat_(
change_port,
})
.unwrap();
udp.send(msg.as_bytes()).await?;
udp.send(msg.as_bytes())?;
let mut buf = [0; 10240];
let (len, addr) =
match tokio::time::timeout(Duration::from_millis(300), udp.recv_from(&mut buf)).await {
Ok(rs) => rs?,
Err(_) => {
continue;
}
};
let (len, _addr) = match udp.recv_from(&mut buf) {
Ok(rs) => rs,
Err(_) => {
continue;
}
};
let msg = stun_format::Msg::from(&buf[..len]);
let mut mapped_addr = None;
let mut changed_addr = None;
@@ -118,8 +118,8 @@ async fn test_nat_(
return Ok((mapped_addr.unwrap(), changed_addr.unwrap()));
}
}
if mapped_addr.is_some() {
return Ok((mapped_addr.unwrap(), changed_addr.unwrap_or(addr)));
if let Some(addr) = mapped_addr {
return Ok((addr, changed_addr.unwrap_or(addr)));
}
}
Err(io::Error::new(io::ErrorKind::Other, "stun response err"))
+392 -5
View File
@@ -1317,6 +1317,12 @@ pub struct PunchInfo {
pub ipv6: ::std::vec::Vec<u8>,
// @@protoc_insertion_point(field:PunchInfo.ipv6_port)
pub ipv6_port: u32,
// @@protoc_insertion_point(field:PunchInfo.tcp_port)
pub tcp_port: u32,
// @@protoc_insertion_point(field:PunchInfo.udp_ports)
pub udp_ports: ::std::vec::Vec<u32>,
// @@protoc_insertion_point(field:PunchInfo.public_ports)
pub public_ports: ::std::vec::Vec<u32>,
// special fields
// @@protoc_insertion_point(special_field:PunchInfo.special_fields)
pub special_fields: ::protobuf::SpecialFields,
@@ -1334,7 +1340,7 @@ impl PunchInfo {
}
fn generated_message_descriptor_data() -> ::protobuf::reflect::GeneratedMessageDescriptorData {
let mut fields = ::std::vec::Vec::with_capacity(9);
let mut fields = ::std::vec::Vec::with_capacity(12);
let mut oneofs = ::std::vec::Vec::with_capacity(0);
fields.push(::protobuf::reflect::rt::v2::make_vec_simpler_accessor::<_, _>(
"public_ip_list",
@@ -1381,6 +1387,21 @@ impl PunchInfo {
|m: &PunchInfo| { &m.ipv6_port },
|m: &mut PunchInfo| { &mut m.ipv6_port },
));
fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"tcp_port",
|m: &PunchInfo| { &m.tcp_port },
|m: &mut PunchInfo| { &mut m.tcp_port },
));
fields.push(::protobuf::reflect::rt::v2::make_vec_simpler_accessor::<_, _>(
"udp_ports",
|m: &PunchInfo| { &m.udp_ports },
|m: &mut PunchInfo| { &mut m.udp_ports },
));
fields.push(::protobuf::reflect::rt::v2::make_vec_simpler_accessor::<_, _>(
"public_ports",
|m: &PunchInfo| { &m.public_ports },
|m: &mut PunchInfo| { &mut m.public_ports },
));
::protobuf::reflect::GeneratedMessageDescriptorData::new_2::<PunchInfo>(
"PunchInfo",
fields,
@@ -1429,6 +1450,21 @@ impl ::protobuf::Message for PunchInfo {
80 => {
self.ipv6_port = is.read_uint32()?;
},
88 => {
self.tcp_port = is.read_uint32()?;
},
98 => {
is.read_repeated_packed_uint32_into(&mut self.udp_ports)?;
},
96 => {
self.udp_ports.push(is.read_uint32()?);
},
106 => {
is.read_repeated_packed_uint32_into(&mut self.public_ports)?;
},
104 => {
self.public_ports.push(is.read_uint32()?);
},
tag => {
::protobuf::rt::read_unknown_or_skip_group(tag, is, self.special_fields.mut_unknown_fields())?;
},
@@ -1466,6 +1502,15 @@ impl ::protobuf::Message for PunchInfo {
if self.ipv6_port != 0 {
my_size += ::protobuf::rt::uint32_size(10, self.ipv6_port);
}
if self.tcp_port != 0 {
my_size += ::protobuf::rt::uint32_size(11, self.tcp_port);
}
for value in &self.udp_ports {
my_size += ::protobuf::rt::uint32_size(12, *value);
};
for value in &self.public_ports {
my_size += ::protobuf::rt::uint32_size(13, *value);
};
my_size += ::protobuf::rt::unknown_fields_size(self.special_fields.unknown_fields());
self.special_fields.cached_size().set(my_size as u32);
my_size
@@ -1499,6 +1544,15 @@ impl ::protobuf::Message for PunchInfo {
if self.ipv6_port != 0 {
os.write_uint32(10, self.ipv6_port)?;
}
if self.tcp_port != 0 {
os.write_uint32(11, self.tcp_port)?;
}
for v in &self.udp_ports {
os.write_uint32(12, *v)?;
};
for v in &self.public_ports {
os.write_uint32(13, *v)?;
};
os.write_unknown_fields(self.special_fields.unknown_fields())?;
::std::result::Result::Ok(())
}
@@ -1525,6 +1579,9 @@ impl ::protobuf::Message for PunchInfo {
self.local_port = 0;
self.ipv6.clear();
self.ipv6_port = 0;
self.tcp_port = 0;
self.udp_ports.clear();
self.public_ports.clear();
self.special_fields.clear();
}
@@ -1539,6 +1596,9 @@ impl ::protobuf::Message for PunchInfo {
local_port: 0,
ipv6: ::std::vec::Vec::new(),
ipv6_port: 0,
tcp_port: 0,
udp_ports: ::std::vec::Vec::new(),
public_ports: ::std::vec::Vec::new(),
special_fields: ::protobuf::SpecialFields::new(),
};
&instance
@@ -1562,6 +1622,323 @@ impl ::protobuf::reflect::ProtobufValue for PunchInfo {
type RuntimeType = ::protobuf::reflect::rt::RuntimeTypeMessage<Self>;
}
#[derive(PartialEq,Clone,Default,Debug)]
// @@protoc_insertion_point(message:ClientStatusInfo)
pub struct ClientStatusInfo {
// message fields
// @@protoc_insertion_point(field:ClientStatusInfo.source)
pub source: u32,
// @@protoc_insertion_point(field:ClientStatusInfo.p2p_list)
pub p2p_list: ::std::vec::Vec<RouteItem>,
// @@protoc_insertion_point(field:ClientStatusInfo.up_stream)
pub up_stream: u64,
// @@protoc_insertion_point(field:ClientStatusInfo.down_stream)
pub down_stream: u64,
// @@protoc_insertion_point(field:ClientStatusInfo.nat_type)
pub nat_type: ::protobuf::EnumOrUnknown<PunchNatType>,
// special fields
// @@protoc_insertion_point(special_field:ClientStatusInfo.special_fields)
pub special_fields: ::protobuf::SpecialFields,
}
impl<'a> ::std::default::Default for &'a ClientStatusInfo {
fn default() -> &'a ClientStatusInfo {
<ClientStatusInfo as ::protobuf::Message>::default_instance()
}
}
impl ClientStatusInfo {
pub fn new() -> ClientStatusInfo {
::std::default::Default::default()
}
fn generated_message_descriptor_data() -> ::protobuf::reflect::GeneratedMessageDescriptorData {
let mut fields = ::std::vec::Vec::with_capacity(5);
let mut oneofs = ::std::vec::Vec::with_capacity(0);
fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"source",
|m: &ClientStatusInfo| { &m.source },
|m: &mut ClientStatusInfo| { &mut m.source },
));
fields.push(::protobuf::reflect::rt::v2::make_vec_simpler_accessor::<_, _>(
"p2p_list",
|m: &ClientStatusInfo| { &m.p2p_list },
|m: &mut ClientStatusInfo| { &mut m.p2p_list },
));
fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"up_stream",
|m: &ClientStatusInfo| { &m.up_stream },
|m: &mut ClientStatusInfo| { &mut m.up_stream },
));
fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"down_stream",
|m: &ClientStatusInfo| { &m.down_stream },
|m: &mut ClientStatusInfo| { &mut m.down_stream },
));
fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"nat_type",
|m: &ClientStatusInfo| { &m.nat_type },
|m: &mut ClientStatusInfo| { &mut m.nat_type },
));
::protobuf::reflect::GeneratedMessageDescriptorData::new_2::<ClientStatusInfo>(
"ClientStatusInfo",
fields,
oneofs,
)
}
}
impl ::protobuf::Message for ClientStatusInfo {
const NAME: &'static str = "ClientStatusInfo";
fn is_initialized(&self) -> bool {
true
}
fn merge_from(&mut self, is: &mut ::protobuf::CodedInputStream<'_>) -> ::protobuf::Result<()> {
while let Some(tag) = is.read_raw_tag_or_eof()? {
match tag {
13 => {
self.source = is.read_fixed32()?;
},
18 => {
self.p2p_list.push(is.read_message()?);
},
24 => {
self.up_stream = is.read_uint64()?;
},
32 => {
self.down_stream = is.read_uint64()?;
},
40 => {
self.nat_type = is.read_enum_or_unknown()?;
},
tag => {
::protobuf::rt::read_unknown_or_skip_group(tag, is, self.special_fields.mut_unknown_fields())?;
},
};
}
::std::result::Result::Ok(())
}
// Compute sizes of nested messages
#[allow(unused_variables)]
fn compute_size(&self) -> u64 {
let mut my_size = 0;
if self.source != 0 {
my_size += 1 + 4;
}
for value in &self.p2p_list {
let len = value.compute_size();
my_size += 1 + ::protobuf::rt::compute_raw_varint64_size(len) + len;
};
if self.up_stream != 0 {
my_size += ::protobuf::rt::uint64_size(3, self.up_stream);
}
if self.down_stream != 0 {
my_size += ::protobuf::rt::uint64_size(4, self.down_stream);
}
if self.nat_type != ::protobuf::EnumOrUnknown::new(PunchNatType::Symmetric) {
my_size += ::protobuf::rt::int32_size(5, self.nat_type.value());
}
my_size += ::protobuf::rt::unknown_fields_size(self.special_fields.unknown_fields());
self.special_fields.cached_size().set(my_size as u32);
my_size
}
fn write_to_with_cached_sizes(&self, os: &mut ::protobuf::CodedOutputStream<'_>) -> ::protobuf::Result<()> {
if self.source != 0 {
os.write_fixed32(1, self.source)?;
}
for v in &self.p2p_list {
::protobuf::rt::write_message_field_with_cached_size(2, v, os)?;
};
if self.up_stream != 0 {
os.write_uint64(3, self.up_stream)?;
}
if self.down_stream != 0 {
os.write_uint64(4, self.down_stream)?;
}
if self.nat_type != ::protobuf::EnumOrUnknown::new(PunchNatType::Symmetric) {
os.write_enum(5, ::protobuf::EnumOrUnknown::value(&self.nat_type))?;
}
os.write_unknown_fields(self.special_fields.unknown_fields())?;
::std::result::Result::Ok(())
}
fn special_fields(&self) -> &::protobuf::SpecialFields {
&self.special_fields
}
fn mut_special_fields(&mut self) -> &mut ::protobuf::SpecialFields {
&mut self.special_fields
}
fn new() -> ClientStatusInfo {
ClientStatusInfo::new()
}
fn clear(&mut self) {
self.source = 0;
self.p2p_list.clear();
self.up_stream = 0;
self.down_stream = 0;
self.nat_type = ::protobuf::EnumOrUnknown::new(PunchNatType::Symmetric);
self.special_fields.clear();
}
fn default_instance() -> &'static ClientStatusInfo {
static instance: ClientStatusInfo = ClientStatusInfo {
source: 0,
p2p_list: ::std::vec::Vec::new(),
up_stream: 0,
down_stream: 0,
nat_type: ::protobuf::EnumOrUnknown::from_i32(0),
special_fields: ::protobuf::SpecialFields::new(),
};
&instance
}
}
impl ::protobuf::MessageFull for ClientStatusInfo {
fn descriptor() -> ::protobuf::reflect::MessageDescriptor {
static descriptor: ::protobuf::rt::Lazy<::protobuf::reflect::MessageDescriptor> = ::protobuf::rt::Lazy::new();
descriptor.get(|| file_descriptor().message_by_package_relative_name("ClientStatusInfo").unwrap()).clone()
}
}
impl ::std::fmt::Display for ClientStatusInfo {
fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result {
::protobuf::text_format::fmt(self, f)
}
}
impl ::protobuf::reflect::ProtobufValue for ClientStatusInfo {
type RuntimeType = ::protobuf::reflect::rt::RuntimeTypeMessage<Self>;
}
#[derive(PartialEq,Clone,Default,Debug)]
// @@protoc_insertion_point(message:RouteItem)
pub struct RouteItem {
// message fields
// @@protoc_insertion_point(field:RouteItem.next_ip)
pub next_ip: u32,
// special fields
// @@protoc_insertion_point(special_field:RouteItem.special_fields)
pub special_fields: ::protobuf::SpecialFields,
}
impl<'a> ::std::default::Default for &'a RouteItem {
fn default() -> &'a RouteItem {
<RouteItem as ::protobuf::Message>::default_instance()
}
}
impl RouteItem {
pub fn new() -> RouteItem {
::std::default::Default::default()
}
fn generated_message_descriptor_data() -> ::protobuf::reflect::GeneratedMessageDescriptorData {
let mut fields = ::std::vec::Vec::with_capacity(1);
let mut oneofs = ::std::vec::Vec::with_capacity(0);
fields.push(::protobuf::reflect::rt::v2::make_simpler_field_accessor::<_, _>(
"next_ip",
|m: &RouteItem| { &m.next_ip },
|m: &mut RouteItem| { &mut m.next_ip },
));
::protobuf::reflect::GeneratedMessageDescriptorData::new_2::<RouteItem>(
"RouteItem",
fields,
oneofs,
)
}
}
impl ::protobuf::Message for RouteItem {
const NAME: &'static str = "RouteItem";
fn is_initialized(&self) -> bool {
true
}
fn merge_from(&mut self, is: &mut ::protobuf::CodedInputStream<'_>) -> ::protobuf::Result<()> {
while let Some(tag) = is.read_raw_tag_or_eof()? {
match tag {
13 => {
self.next_ip = is.read_fixed32()?;
},
tag => {
::protobuf::rt::read_unknown_or_skip_group(tag, is, self.special_fields.mut_unknown_fields())?;
},
};
}
::std::result::Result::Ok(())
}
// Compute sizes of nested messages
#[allow(unused_variables)]
fn compute_size(&self) -> u64 {
let mut my_size = 0;
if self.next_ip != 0 {
my_size += 1 + 4;
}
my_size += ::protobuf::rt::unknown_fields_size(self.special_fields.unknown_fields());
self.special_fields.cached_size().set(my_size as u32);
my_size
}
fn write_to_with_cached_sizes(&self, os: &mut ::protobuf::CodedOutputStream<'_>) -> ::protobuf::Result<()> {
if self.next_ip != 0 {
os.write_fixed32(1, self.next_ip)?;
}
os.write_unknown_fields(self.special_fields.unknown_fields())?;
::std::result::Result::Ok(())
}
fn special_fields(&self) -> &::protobuf::SpecialFields {
&self.special_fields
}
fn mut_special_fields(&mut self) -> &mut ::protobuf::SpecialFields {
&mut self.special_fields
}
fn new() -> RouteItem {
RouteItem::new()
}
fn clear(&mut self) {
self.next_ip = 0;
self.special_fields.clear();
}
fn default_instance() -> &'static RouteItem {
static instance: RouteItem = RouteItem {
next_ip: 0,
special_fields: ::protobuf::SpecialFields::new(),
};
&instance
}
}
impl ::protobuf::MessageFull for RouteItem {
fn descriptor() -> ::protobuf::reflect::MessageDescriptor {
static descriptor: ::protobuf::rt::Lazy<::protobuf::reflect::MessageDescriptor> = ::protobuf::rt::Lazy::new();
descriptor.get(|| file_descriptor().message_by_package_relative_name("RouteItem").unwrap()).clone()
}
}
impl ::std::fmt::Display for RouteItem {
fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result {
::protobuf::text_format::fmt(self, f)
}
}
impl ::protobuf::reflect::ProtobufValue for RouteItem {
type RuntimeType = ::protobuf::reflect::rt::RuntimeTypeMessage<Self>;
}
#[derive(Clone,Copy,PartialEq,Eq,Debug,Hash)]
// @@protoc_insertion_point(enum:PunchNatType)
pub enum PunchNatType {
@@ -1643,7 +2020,7 @@ static file_descriptor_proto_data: &'static [u8] = b"\
\n\rdevice_status\x18\x03\x20\x01(\rR\x0cdeviceStatus\x12#\n\rclient_sec\
ret\x18\x04\x20\x01(\x08R\x0cclientSecret\"Y\n\nDeviceList\x12\x14\n\x05\
epoch\x18\x01\x20\x01(\rR\x05epoch\x125\n\x10device_info_list\x18\x02\
\x20\x03(\x0b2\x0b.DeviceInfoR\x0edeviceInfoList\"\xa9\x02\n\tPunchInfo\
\x20\x03(\x0b2\x0b.DeviceInfoR\x0edeviceInfoList\"\x84\x03\n\tPunchInfo\
\x12$\n\x0epublic_ip_list\x18\x02\x20\x03(\x07R\x0cpublicIpList\x12\x1f\
\n\x0bpublic_port\x18\x03\x20\x01(\rR\npublicPort\x12*\n\x11public_port_\
range\x18\x04\x20\x01(\rR\x0fpublicPortRange\x12(\n\x08nat_type\x18\x05\
@@ -1651,8 +2028,16 @@ static file_descriptor_proto_data: &'static [u8] = b"\
\x01(\x08R\x05reply\x12\x19\n\x08local_ip\x18\x07\x20\x01(\x07R\x07local\
Ip\x12\x1d\n\nlocal_port\x18\x08\x20\x01(\rR\tlocalPort\x12\x12\n\x04ipv\
6\x18\t\x20\x01(\x0cR\x04ipv6\x12\x1b\n\tipv6_port\x18\n\x20\x01(\rR\x08\
ipv6Port*'\n\x0cPunchNatType\x12\r\n\tSymmetric\x10\0\x12\x08\n\x04Cone\
\x10\x01b\x06proto3\
ipv6Port\x12\x19\n\x08tcp_port\x18\x0b\x20\x01(\rR\x07tcpPort\x12\x1b\n\
\tudp_ports\x18\x0c\x20\x03(\rR\x08udpPorts\x12!\n\x0cpublic_ports\x18\r\
\x20\x03(\rR\x0bpublicPorts\"\xb9\x01\n\x10ClientStatusInfo\x12\x16\n\
\x06source\x18\x01\x20\x01(\x07R\x06source\x12%\n\x08p2p_list\x18\x02\
\x20\x03(\x0b2\n.RouteItemR\x07p2pList\x12\x1b\n\tup_stream\x18\x03\x20\
\x01(\x04R\x08upStream\x12\x1f\n\x0bdown_stream\x18\x04\x20\x01(\x04R\nd\
ownStream\x12(\n\x08nat_type\x18\x05\x20\x01(\x0e2\r.PunchNatTypeR\x07na\
tType\"$\n\tRouteItem\x12\x17\n\x07next_ip\x18\x01\x20\x01(\x07R\x06next\
Ip*'\n\x0cPunchNatType\x12\r\n\tSymmetric\x10\0\x12\x08\n\x04Cone\x10\
\x01b\x06proto3\
";
/// `FileDescriptorProto` object which was a source for this generated file
@@ -1670,7 +2055,7 @@ pub fn file_descriptor() -> &'static ::protobuf::reflect::FileDescriptor {
file_descriptor.get(|| {
let generated_file_descriptor = generated_file_descriptor_lazy.get(|| {
let mut deps = ::std::vec::Vec::with_capacity(0);
let mut messages = ::std::vec::Vec::with_capacity(8);
let mut messages = ::std::vec::Vec::with_capacity(10);
messages.push(HandshakeRequest::generated_message_descriptor_data());
messages.push(HandshakeResponse::generated_message_descriptor_data());
messages.push(SecretHandshakeRequest::generated_message_descriptor_data());
@@ -1679,6 +2064,8 @@ pub fn file_descriptor() -> &'static ::protobuf::reflect::FileDescriptor {
messages.push(DeviceInfo::generated_message_descriptor_data());
messages.push(DeviceList::generated_message_descriptor_data());
messages.push(PunchInfo::generated_message_descriptor_data());
messages.push(ClientStatusInfo::generated_message_descriptor_data());
messages.push(RouteItem::generated_message_descriptor_data());
let mut enums = ::std::vec::Vec::with_capacity(1);
enums.push(PunchNatType::generated_enum_descriptor_data());
::protobuf::reflect::GeneratedFileDescriptor::new_generated(
+1 -1
View File
@@ -1,6 +1,6 @@
use std::{fmt, io};
pub const ENCRYPTION_RESERVED: usize = 32 + 12;
pub const ENCRYPTION_RESERVED: usize = 16 + 32 + 12;
pub const AES_GCM_ENCRYPTION_RESERVED: usize = 32;
pub const RSA_ENCRYPTION_RESERVED: usize = 32;
+5 -5
View File
@@ -1,4 +1,4 @@
use crate::error::*;
use std::io;
#[derive(Eq, PartialEq, Copy, Clone, Debug)]
pub enum Protocol {
@@ -50,7 +50,7 @@ pub enum InErrorPacket<B> {
}
impl<B: AsRef<[u8]>> InErrorPacket<B> {
pub fn new(protocol: u8, buffer: B) -> Result<InErrorPacket<B>> {
pub fn new(protocol: u8, buffer: B) -> io::Result<InErrorPacket<B>> {
match Protocol::from(protocol) {
Protocol::TokenError => Ok(InErrorPacket::TokenError),
Protocol::Disconnect => Ok(InErrorPacket::Disconnect),
@@ -68,16 +68,16 @@ pub struct ErrorPacket<B> {
}
impl<B: AsRef<[u8]>> ErrorPacket<B> {
pub fn new(buffer: B) -> Result<ErrorPacket<B>> {
pub fn new(buffer: B) -> io::Result<ErrorPacket<B>> {
Ok(Self { buffer })
}
}
impl<B: AsRef<[u8]>> ErrorPacket<B> {
pub fn message(&self) -> Result<String> {
pub fn message(&self) -> io::Result<String> {
match String::from_utf8(self.buffer.as_ref().to_vec()) {
Ok(str) => Ok(str),
Err(_) => Err(Error::InvalidPacket),
Err(_) => Err(io::Error::new(io::ErrorKind::Other, "Utf8Error")),
}
}
}
+20 -8
View File
@@ -28,14 +28,14 @@ pub mod service_packet;
#[derive(Eq, PartialEq, Copy, Clone, Debug)]
pub enum Version {
V1,
UnKnow(u8),
Unknown(u8),
}
impl From<u8> for Version {
fn from(value: u8) -> Self {
match value {
1 => Version::V1,
val => Version::UnKnow(val),
val => Version::Unknown(val),
}
}
}
@@ -44,7 +44,7 @@ impl Into<u8> for Version {
fn into(self) -> u8 {
match self {
Version::V1 => 1,
Version::UnKnow(val) => val,
Version::Unknown(val) => val,
}
}
}
@@ -61,7 +61,7 @@ pub enum Protocol {
IpTurn,
/// 转发其他数据
OtherTurn,
UnKnow(u8),
Unknown(u8),
}
impl From<u8> for Protocol {
@@ -72,7 +72,7 @@ impl From<u8> for Protocol {
3 => Protocol::Control,
4 => Protocol::IpTurn,
5 => Protocol::OtherTurn,
val => Protocol::UnKnow(val),
val => Protocol::Unknown(val),
}
}
}
@@ -85,7 +85,7 @@ impl Into<u8> for Protocol {
Protocol::Control => 3,
Protocol::IpTurn => 4,
Protocol::OtherTurn => 5,
Protocol::UnKnow(val) => val,
Protocol::Unknown(val) => val,
}
}
}
@@ -123,7 +123,7 @@ impl<B: AsRef<[u8]>> NetPacket<B> {
));
}
// 不能大于udp最大载荷长度
if data_len < 12 || buffer.as_ref().len() > 65535 - 20 - 8 {
if data_len < 12 || data_len > 65535 - 20 - 8 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"length overflow",
@@ -155,7 +155,7 @@ impl<B: AsRef<[u8]>> NetPacket<B> {
}
/// 网关通信的标识
pub fn is_gateway(&self) -> bool {
self.buffer.as_ref()[0] & 0x50 == 0x50
self.buffer.as_ref()[0] & 0x40 == 0x40
}
pub fn version(&self) -> Version {
Version::from(self.buffer.as_ref()[0] & 0x0F)
@@ -183,6 +183,9 @@ impl<B: AsRef<[u8]>> NetPacket<B> {
pub fn payload(&self) -> &[u8] {
&self.buffer.as_ref()[12..self.data_len]
}
pub fn head(&self) -> &[u8] {
&self.buffer.as_ref()[..12]
}
}
impl<B: AsRef<[u8]> + AsMut<[u8]>> NetPacket<B> {
@@ -198,6 +201,7 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> NetPacket<B> {
}
pub fn set_gateway_flag(&mut self, is_gateway: bool) {
if is_gateway {
// 后面的版本再改为0x40,改了之后不兼容1.2.5之前的版本
self.buffer.as_mut()[0] = self.buffer.as_ref()[0] | 0x50
} else {
self.buffer.as_mut()[0] = self.buffer.as_ref()[0] & 0xBF
@@ -213,12 +217,20 @@ impl<B: AsRef<[u8]> + AsMut<[u8]>> NetPacket<B> {
pub fn set_transport_protocol(&mut self, transport_protocol: u8) {
self.buffer.as_mut()[2] = transport_protocol;
}
pub fn set_transport_protocol_into<P: Into<u8>>(&mut self, transport_protocol: P) {
self.buffer.as_mut()[2] = transport_protocol.into();
}
pub fn first_set_ttl(&mut self, ttl: u8) {
self.buffer.as_mut()[3] = ttl << 4 | ttl;
}
pub fn set_ttl(&mut self, ttl: u8) {
self.buffer.as_mut()[3] = (self.buffer.as_mut()[3] & MAX_SOURCE) | (MAX_TTL & ttl);
}
pub fn incr_ttl(&mut self) -> u8 {
let ttl = self.ttl() - 1;
self.set_ttl(ttl);
ttl
}
pub fn set_source_ttl(&mut self, source_ttl: u8) {
self.buffer.as_mut()[3] = (source_ttl << 4) | (MAX_TTL & self.buffer.as_ref()[3]);
}
+4
View File
@@ -13,6 +13,8 @@ pub enum Protocol {
HandshakeResponse,
SecretHandshakeRequest,
SecretHandshakeResponse,
/// 客户端上报状态
ClientStatusInfo,
Unknown(u8),
}
@@ -27,6 +29,7 @@ impl From<u8> for Protocol {
6 => Self::HandshakeResponse,
7 => Self::SecretHandshakeRequest,
8 => Self::SecretHandshakeResponse,
9 => Self::ClientStatusInfo,
val => Self::Unknown(val),
}
}
@@ -43,6 +46,7 @@ impl Into<u8> for Protocol {
Self::HandshakeResponse => 6,
Self::SecretHandshakeRequest => 7,
Self::SecretHandshakeResponse => 8,
Self::ClientStatusInfo => 9,
Self::Unknown(val) => val,
}
}
-48
View File
@@ -1,48 +0,0 @@
use std::io;
use std::os::unix::io::RawFd;
#[derive(Clone)]
pub struct DeviceWriter(RawFd);
pub struct DeviceReader(RawFd);
impl DeviceWriter {
pub fn write_ipv4_tun(&self, buf: &[u8]) -> io::Result<()> {
unsafe {
let amount = libc::write(self.0, buf.as_ptr() as *const _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
///写入ipv4数据,为了兼容其他代码,头部空了14个字节
pub fn write_ipv4(&self, buf: &[u8]) -> io::Result<()> {
let buf = &buf[14..];
self.write_ipv4_tun(buf)
}
pub fn close(&self) -> io::Result<()> {
// unsafe {
// libc::close(self.0);
// }
Ok(())
}
}
impl DeviceReader {
pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
unsafe {
let amount = libc::read(self.0, buf.as_mut_ptr() as *mut _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(amount as usize)
}
}
}
pub fn create(fd: i32) -> (DeviceWriter, DeviceReader) {
(DeviceWriter(fd as _), DeviceReader(fd as _))
}
-175
View File
@@ -1,175 +0,0 @@
use crate::tun_tap_device::linux_mac::DeviceW;
use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo};
use parking_lot::Mutex;
use std::io;
use std::net::Ipv4Addr;
use std::process::Command;
use std::sync::Arc;
use tun::Device;
pub const TUN_INTERFACE_NAME: &str = "vnt-tun";
pub const TAP_INTERFACE_NAME: &str = "vnt-tap";
impl DeviceWriter {
pub fn change_ip(
&self,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
_old_netmask: Ipv4Addr,
_old_gateway: Ipv4Addr,
) -> io::Result<()> {
let mut config = tun::Configuration::default();
let broadcast_address =
(!u32::from_be_bytes(netmask.octets())) | u32::from_be_bytes(gateway.octets());
let broadcast_address = Ipv4Addr::from(broadcast_address);
config
.destination(gateway)
.address(address)
.netmask(netmask)
.broadcast(broadcast_address)
// .queues(2)
.up();
let mut dev = self.lock.lock();
if let Err(e) = dev.configure(&config) {
return Err(io::Error::new(io::ErrorKind::Other, format!("{:?}", e)));
}
let name = dev.name();
for (address, netmask) in &self.in_ips {
add_route(name, *address, *netmask)?;
}
// 当前网段路由
// add_route(name, address, netmask)?;
// 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
return Ok(());
}
}
pub fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
let route_add_str: String = format!("ip route add {:?}/{:?} dev {}", address, netmask, name);
let route_add_out = Command::new("sh")
.arg("-c")
.arg(&route_add_str)
.output()
.expect("sh exec error!");
if !route_add_out.status.success() {
return Err(io::Error::new(
io::ErrorKind::Other,
format!(
"添加路由失败: cmd:{},out:{:?}",
route_add_str, route_add_out
),
));
}
Ok(())
}
pub fn create_device(
device_type: DeviceType,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
let mut config = tun::Configuration::default();
let broadcast_address =
(!u32::from_be_bytes(netmask.octets())) | u32::from_be_bytes(gateway.octets());
let broadcast_address = Ipv4Addr::from(broadcast_address);
config
.destination(gateway)
.address(address)
.netmask(netmask)
.mtu(mtu.into())
.broadcast(broadcast_address)
// .queues(2) 用多个队列有兼容性问题
.up();
match device_type {
DeviceType::Tun => {
config.name(TUN_INTERFACE_NAME);
}
DeviceType::Tap => {
config.name(TAP_INTERFACE_NAME);
config.layer(tun::Layer::L2);
}
}
let dev = tun::create(&config).expect("tun/tap failed to create");
let packet_information = dev.has_packet_information();
let queue = dev.queue(0).unwrap();
let reader = queue.reader();
let writer = queue.writer();
let name = dev.name();
for (address, netmask) in &in_ips {
add_route(name, *address, *netmask)?;
}
// 当前网段路由
// add_route(name, address, netmask)?;
// 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
let device_w = match device_type {
DeviceType::Tun => DeviceW::Tun(writer),
DeviceType::Tap => {
let get_mac_cmd = format!("cat /sys/class/net/{}/address", name);
let mac_out = Command::new("sh")
.arg("-c")
.arg(get_mac_cmd)
.output()
.expect("sh exec error!");
if !mac_out.status.success() {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("获取mac地址错误: {:?}", mac_out),
));
}
let mac_str = String::from_utf8(mac_out.stdout).unwrap();
let mut mac = [0; 6];
let mut split = mac_str.split(":");
for i in 0..6 {
mac[i] = u8::from_str_radix(&split.next().unwrap()[..2], 16).unwrap();
}
DeviceW::Tap((writer, mac))
}
};
let driver_info = DriverInfo {
device_type,
name: name.to_string(),
version: String::new(),
mac: None,
};
Ok((
DeviceWriter::new(
device_w,
Arc::new(Mutex::new(dev)),
in_ips,
address,
packet_information,
),
DeviceReader::new(reader),
driver_info,
))
}
pub fn delete_device(_device_type: DeviceType) {
for name in [TUN_INTERFACE_NAME, TAP_INTERFACE_NAME] {
let cmd = format!("ip link delete {}", name);
let delete_tun = Command::new("sh")
.arg("-c")
.arg(&cmd)
.output()
.expect("sh exec error!");
if !delete_tun.status.success() {
log::warn!("删除网卡失败:{:?}", delete_tun);
}
}
}
-144
View File
@@ -1,144 +0,0 @@
use std::io;
use std::sync::Arc;
use bytes::BufMut;
use packet::ethernet;
use parking_lot::Mutex;
use std::net::Ipv4Addr;
use std::os::unix::io::AsRawFd;
#[cfg(any(target_os = "linux"))]
use tun::platform::linux::Device;
#[cfg(any(target_os = "macos"))]
use tun::platform::macos::Device;
use tun::platform::posix::{Reader, Writer};
use packet::ethernet::packet::EthernetPacket;
#[derive(Clone)]
pub enum DeviceW {
Tun(Writer),
Tap((Writer, [u8; 6])),
}
impl DeviceW {
pub fn is_tun(&self) -> bool {
match self {
DeviceW::Tun(_) => true,
DeviceW::Tap(_) => false,
}
}
}
#[derive(Clone)]
pub struct DeviceWriter {
writer: DeviceW,
pub lock: Arc<Mutex<Device>>,
pub in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
packet_information: bool,
}
impl DeviceWriter {
pub fn new(
writer: DeviceW,
lock: Arc<Mutex<Device>>,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
_ip: Ipv4Addr,
packet_information: bool,
) -> Self {
Self {
writer,
lock,
in_ips,
packet_information,
}
}
}
impl DeviceWriter {
pub fn write(packet_information: bool, writer: &Writer, packet: &[u8]) -> io::Result<()> {
if packet_information {
let mut buf = Vec::<u8>::with_capacity(4 + packet.len());
buf.put_u16(0);
#[cfg(any(target_os = "macos", target_os = "ios"))]
buf.put_u16(libc::PF_INET as u16);
#[cfg(any(target_os = "linux", target_os = "android"))]
buf.put_u16(libc::ETH_P_IP as u16);
buf.extend_from_slice(packet);
let len = writer.write(&buf)?;
if len != buf.len() {
log::error!("tun write error");
}
} else {
let len = writer.write(packet)?;
if len != packet.len() {
log::error!("tun write error");
}
}
Ok(())
}
///tun网卡写入ipv4数据
pub fn write_ipv4_tun(&self, buf: &[u8]) -> io::Result<()> {
match &self.writer {
DeviceW::Tun(writer) => Self::write(self.packet_information, writer, buf),
DeviceW::Tap(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
}
}
/// tap网卡写入以太网帧
pub fn write_ethernet_tap(&self, buf: &[u8]) -> io::Result<()> {
match &self.writer {
DeviceW::Tun(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
DeviceW::Tap((writer, _)) => Self::write(self.packet_information, writer, buf),
}
}
///写入ipv4数据,头部必须留14字节,给tap写入以太网帧头
pub fn write_ipv4(&self, buf: &mut [u8]) -> io::Result<()> {
match &self.writer {
DeviceW::Tun(writer) => Self::write(self.packet_information, writer, &buf[14..]),
DeviceW::Tap((writer, mac)) => {
let source_mac = [
buf[14 + 12],
buf[14 + 13],
buf[14 + 14],
buf[14 + 15],
!mac[5],
234,
];
let mut ethernet_packet = EthernetPacket::unchecked(buf);
ethernet_packet.set_source(&source_mac);
ethernet_packet.set_destination(mac);
ethernet_packet.set_protocol(ethernet::protocol::Protocol::Ipv4);
Self::write(self.packet_information, writer, &ethernet_packet.buffer)
}
}
}
pub fn close(&self) -> io::Result<()> {
unsafe {
match &self.writer {
DeviceW::Tun(writer) => {
libc::close(writer.as_raw_fd());
}
DeviceW::Tap((writer, _)) => {
libc::close(writer.as_raw_fd());
}
}
}
Ok(())
}
pub fn is_tun(&self) -> bool {
self.writer.is_tun()
}
}
pub struct DeviceReader(Reader);
impl DeviceReader {
pub fn new(device: Reader) -> Self {
DeviceReader(device)
}
}
impl DeviceReader {
pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
self.0.read(buf)
}
}
-153
View File
@@ -1,153 +0,0 @@
use crate::tun_tap_device::linux_mac::DeviceW;
use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo};
use parking_lot::Mutex;
use std::io;
use std::net::Ipv4Addr;
use std::process::Command;
use std::sync::Arc;
use tun::Device;
impl DeviceWriter {
pub fn change_ip(
&self,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
_old_netmask: Ipv4Addr,
_old_gateway: Ipv4Addr,
) -> io::Result<()> {
let mut config = tun::Configuration::default();
config
.destination(gateway)
.address(address)
.netmask(netmask)
.up();
let mut dev = self.lock.lock();
if let Err(e) = dev.configure(&config) {
return Err(io::Error::new(io::ErrorKind::Other, format!("{:?}", e)));
}
if let Err(e) = config_ip(dev.name(), address, netmask, gateway) {
log::error!("{}", e);
}
let name = dev.name();
for (address, netmask) in &self.in_ips {
add_route(name, *address, *netmask)?;
}
// 当前网段路由
add_route(name, address, netmask)?;
// 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
return Ok(());
}
}
pub fn create_device(
device_type: DeviceType,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
match device_type {
DeviceType::Tun => {}
DeviceType::Tap => {
unimplemented!()
}
}
let mut config = tun::Configuration::default();
config
.destination(gateway)
.address(address)
.netmask(netmask)
.mtu(mtu.into())
.up();
let dev = tun::create(&config).unwrap();
let name = dev.name();
config_ip(name, address, netmask, gateway)?;
for (address, netmask) in &in_ips {
add_route(name, *address, *netmask)?;
}
// 当前网段路由
add_route(name, address, netmask)?;
// 广播和组播路由
add_route(name, Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST)?;
add_route(
name,
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
)?;
let packet_information = dev.has_packet_information();
let queue = dev.queue(0).unwrap();
let reader = queue.reader();
let writer = queue.writer();
let driver_info = DriverInfo {
device_type,
name: name.to_string(),
version: String::new(),
mac: None,
};
Ok((
DeviceWriter::new(
DeviceW::Tun(writer),
Arc::new(Mutex::new(dev)),
in_ips,
address,
packet_information,
),
DeviceReader::new(reader),
driver_info,
))
}
fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
let route_add_str: String = format!(
"route -n add {} -netmask {} -interface {}",
address, netmask, name
);
let route_add_out = Command::new("sh")
.arg("-c")
.arg(&route_add_str)
.output()
.expect("sh exec error!");
if !route_add_out.status.success() {
return Err(io::Error::new(
io::ErrorKind::Other,
format!(
"添加路由失败: cmd:{},out:{:?}",
route_add_str, route_add_out
),
));
}
Ok(())
}
fn config_ip(
name: &str,
address: Ipv4Addr,
_netmask: Ipv4Addr,
gateway: Ipv4Addr,
) -> io::Result<()> {
let up_eth_str: String = format!("ifconfig {} {:?} {:?} up ", name, address, gateway);
let up_eth_out = Command::new("sh")
.arg("-c")
.arg(&up_eth_str)
.output()
.expect("sh exec error!");
if !up_eth_out.status.success() {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置网络地址失败: cmd:{},out:{:?}", up_eth_str, up_eth_out),
));
}
Ok(())
}
pub fn delete_device(_device_type: DeviceType) {}
+63 -45
View File
@@ -1,52 +1,70 @@
#[cfg(target_os = "android")]
mod android;
#[cfg(any(target_os = "linux"))]
mod linux;
#[cfg(any(target_os = "linux", target_os = "macos"))]
mod linux_mac;
#[cfg(target_os = "macos")]
mod mac;
#[cfg(target_os = "windows")]
mod windows;
use std::io;
use std::sync::Arc;
#[cfg(target_os = "android")]
pub use android::create;
#[cfg(target_os = "android")]
pub use android::{DeviceReader, DeviceWriter};
#[cfg(any(target_os = "linux"))]
pub use linux::create_device;
#[cfg(any(target_os = "linux"))]
pub use linux::delete_device;
#[cfg(any(target_os = "linux", target_os = "macos"))]
pub use linux_mac::{DeviceReader, DeviceWriter};
#[cfg(target_os = "macos")]
pub use mac::create_device;
#[cfg(target_os = "macos")]
pub use mac::delete_device;
use tun::device::IFace;
use tun::Device;
#[cfg(target_os = "windows")]
pub use windows::create_device;
#[cfg(target_os = "windows")]
pub use windows::delete_device;
#[cfg(target_os = "windows")]
pub use windows::{DeviceReader, DeviceWriter};
use crate::core::Config;
#[cfg(any(target_os = "windows", target_os = "linux"))]
const DEFAULT_TUN_NAME: &str = "vnt-tun";
#[cfg(any(target_os = "windows", target_os = "linux"))]
const DEFAULT_TAP_NAME: &str = "vnt-tap";
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum DeviceType {
Tun,
Tap,
pub fn create_device(config: &Config) -> io::Result<Arc<Device>> {
#[cfg(any(target_os = "windows", target_os = "linux"))]
let default_name: &str = if config.tap {
DEFAULT_TAP_NAME
} else {
DEFAULT_TUN_NAME
};
#[cfg(target_os = "linux")]
let device = {
let device_name = config
.device_name
.clone()
.unwrap_or(default_name.to_string());
if &device_name == default_name {
delete_device(default_name);
}
Arc::new(Device::new(Some(device_name), config.tap)?)
};
#[cfg(target_os = "macos")]
let device = Arc::new(Device::new(config.device_name.clone())?);
#[cfg(target_os = "windows")]
let device = Arc::new(Device::new(
config
.device_name
.clone()
.unwrap_or(default_name.to_string()),
config.tap,
)?);
#[cfg(target_os = "android")]
let device = Arc::new(Device::new(config.device_fd as _)?);
#[cfg(not(target_os = "android"))]
{
let mtu = config.mtu.unwrap_or_else(|| {
if config.password.is_none() {
1450
} else {
1410
}
});
device.set_mtu(mtu)?;
}
Ok(device)
}
impl DeviceType {
pub fn is_tun(&self) -> bool {
*self == DeviceType::Tun
#[cfg(target_os = "linux")]
fn delete_device(name: &str) {
// 删除默认网卡,此操作有风险,后续可能去除
use std::process::Command;
let cmd = format!("ip link delete {}", name);
let delete_tun = Command::new("sh")
.arg("-c")
.arg(&cmd)
.output()
.expect("sh exec error!");
if !delete_tun.status.success() {
log::warn!("删除网卡失败:{:?}", delete_tun);
}
}
#[derive(Clone)]
pub struct DriverInfo {
pub device_type: DeviceType,
pub name: String,
pub version: String,
pub mac: Option<String>,
}
-363
View File
@@ -1,363 +0,0 @@
use crate::tun_tap_device::{DeviceType, DriverInfo};
use libloading::Library;
use packet::ethernet;
use packet::ethernet::packet::EthernetPacket;
use parking_lot::Mutex;
use std::net::Ipv4Addr;
use std::os::windows::process::CommandExt;
use std::sync::Arc;
use std::time::Duration;
use std::{io, thread};
use win_tun_tap::{IFace, TapDevice, TunDevice};
pub const TUN_INTERFACE_NAME: &str = "Vnt-Tun-V1";
pub const TUN_POOL_NAME: &str = "Vnt-Tun-V1";
pub const TAP_INTERFACE_NAME: &str = "Vnt-Tap-V1";
pub enum Device {
Tun(TunDevice),
Tap((TapDevice, [u8; 6])),
}
impl Device {
pub fn is_tun(&self) -> bool {
match self {
Device::Tun(_) => true,
Device::Tap(_) => false,
}
}
}
#[derive(Clone)]
pub struct DeviceWriter {
device: Arc<Device>,
lock: Arc<Mutex<()>>,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
}
impl DeviceWriter {
pub fn new(device: Arc<Device>, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, _ip: Ipv4Addr) -> Self {
Self {
device,
lock: Arc::new(Default::default()),
in_ips,
}
}
}
impl DeviceWriter {
///tun网卡写入ipv4数据
pub fn write_ipv4_tun(&self, buf: &[u8]) -> io::Result<()> {
match self.device.as_ref() {
Device::Tun(dev) => {
let mut packet = dev.allocate_send_packet(buf.len() as u16)?;
packet.bytes_mut().copy_from_slice(buf);
dev.send_packet(packet);
Ok(())
}
Device::Tap(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
}
}
/// tap网卡写入以太网帧
pub fn write_ethernet_tap(&self, buf: &[u8]) -> io::Result<()> {
match self.device.as_ref() {
Device::Tun(_) => Err(io::Error::from(io::ErrorKind::Unsupported)),
Device::Tap((dev, _)) => {
dev.write(buf)?;
Ok(())
}
}
}
///写入ipv4数据,头部必须留14字节,给tap写入以太网帧头
pub fn write_ipv4(&self, buf: &mut [u8]) -> io::Result<()> {
match self.device.as_ref() {
Device::Tun(dev) => {
let mut packet = dev.allocate_send_packet((buf.len() - 14) as u16)?;
packet.bytes_mut().copy_from_slice(&buf[14..]);
dev.send_packet(packet);
}
Device::Tap((dev, mac)) => {
let source_mac = [
buf[14 + 12],
buf[14 + 13],
buf[14 + 14],
buf[14 + 15],
!mac[5],
234,
];
let mut ethernet_packet = EthernetPacket::unchecked(buf);
ethernet_packet.set_source(&source_mac);
ethernet_packet.set_destination(mac);
ethernet_packet.set_protocol(ethernet::protocol::Protocol::Ipv4);
dev.write(&ethernet_packet.buffer)?;
}
}
Ok(())
}
pub fn change_ip(
&self,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
old_netmask: Ipv4Addr,
old_gateway: Ipv4Addr,
) -> io::Result<()> {
let _guard = self.lock.lock();
let dev: &dyn IFace = match self.device.as_ref() {
Device::Tun(dev) => dev as &dyn IFace,
Device::Tap((dev, _)) => dev as &dyn IFace,
};
if let Err(e) = dev.delete_route(dest(old_gateway, old_gateway), old_netmask, old_gateway) {
log::warn!("{:?}", e);
}
dev.set_ip(address, netmask)?;
for (address, netmask) in &self.in_ips {
dev.add_route(*address, *netmask, gateway, 1)?;
}
// 当前网段路由
dev.add_route(address, netmask, gateway, 1)?;
// 广播和组播路由
dev.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?;
dev.add_route(
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
gateway,
1,
)?;
delete_cache();
Ok(())
}
pub fn close(&self) -> io::Result<()> {
match self.device.as_ref() {
Device::Tun(dev) => dev.shutdown(),
Device::Tap((dev, _)) => dev.shutdown(),
}
}
pub fn is_tun(&self) -> bool {
self.device.is_tun()
}
}
fn dest(ip: Ipv4Addr, mask: Ipv4Addr) -> Ipv4Addr {
let ip = ip.octets();
let mask = mask.octets();
Ipv4Addr::from([
ip[0] & mask[0],
ip[1] & mask[1],
ip[2] & mask[2],
ip[3] & mask[3],
])
}
pub struct DeviceReader {
device: Arc<Device>,
}
impl DeviceReader {
pub fn new(device: Arc<Device>) -> Self {
Self { device }
}
}
impl DeviceReader {
pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
match self.device.as_ref() {
Device::Tun(dev) => {
let packet = dev.receive_blocking()?;
let packet = packet.bytes();
let len = packet.len();
if len > buf.len() {
return Err(io::Error::new(io::ErrorKind::InvalidData, "data too long"));
}
buf[..len].copy_from_slice(packet);
Ok(len)
}
Device::Tap((dev, _)) => dev.read(buf),
}
}
}
fn create_tun(
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
unsafe {
match Library::new("wintun.dll") {
Ok(lib) => match TunDevice::delete_for_name(lib, TUN_INTERFACE_NAME) {
Ok(_) => {
thread::sleep(Duration::from_millis(5));
}
Err(_) => {}
},
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("wintun.dll not found {:?}", e),
));
}
}
let tun_device = match TunDevice::create(
Library::new("wintun.dll").unwrap(),
TUN_POOL_NAME,
TUN_INTERFACE_NAME,
) {
Ok(tun_device) => tun_device,
Err(_) => {
thread::sleep(Duration::from_millis(200));
match TunDevice::create(
Library::new("wintun.dll").unwrap(),
TUN_POOL_NAME,
TUN_INTERFACE_NAME,
) {
Ok(tun_device) => tun_device,
Err(e) => {
return Err(io::Error::new(io::ErrorKind::Other, format!("{:?}", e)));
}
}
}
};
let name = tun_device.get_name()?;
let version = format!("{:?}", tun_device.version()?);
tun_device.set_ip(address, netmask)?;
tun_device.set_metric(1)?;
tun_device.set_mtu(mtu)?;
// ip代理路由
for (address, netmask) in &in_ips {
tun_device.add_route(*address, *netmask, gateway, 1)?;
}
// 当前网段路由
tun_device.add_route(address, netmask, gateway, 1)?;
// 广播和组播路由
tun_device.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?;
tun_device.add_route(
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
gateway,
1,
)?;
delete_cache();
let device = Arc::new(Device::Tun(tun_device));
let driver_info = DriverInfo {
device_type: DeviceType::Tun,
name,
version,
mac: None,
};
Ok((
DeviceWriter::new(device.clone(), in_ips, address),
DeviceReader::new(device),
driver_info,
))
}
}
fn delete_cache() {
//清除路由缓存
let delete_cache = "netsh interface ip delete destinationcache";
let out = std::process::Command::new("cmd")
.creation_flags(0x08000000)
.arg("/C")
.arg(delete_cache)
.output()
.unwrap();
if !out.status.success() {
log::warn!("删除缓存失败:{:?}", out);
}
}
fn delete_tun() {
unsafe {
match Library::new("wintun.dll") {
Ok(lib) => match TunDevice::delete_for_name(lib, TUN_INTERFACE_NAME) {
Ok(_) => {}
Err(_) => {}
},
Err(_) => {}
}
}
}
fn create_tap(
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
let tap_device = match TapDevice::open(TAP_INTERFACE_NAME) {
Ok(tap_device) => tap_device,
Err(e) => {
log::warn!("{:?}", e);
let tap_device = TapDevice::create()?;
tap_device.set_name(TAP_INTERFACE_NAME)?;
tap_device
}
};
let mac = tap_device.get_mac()?;
let name = tap_device.get_name()?;
let version = format!("{:?}", tap_device.get_version()?);
let mac_str = format!("mac:{:x?}", mac);
tap_device.set_ip(address, netmask)?;
tap_device.set_metric(1)?;
tap_device.set_mtu(mtu)?;
tap_device.set_status(true)?;
tap_device.add_route(address, netmask, gateway, 1)?;
for (address, netmask) in &in_ips {
tap_device.add_route(*address, *netmask, gateway, 1)?;
}
// 广播和组播路由
tap_device.add_route(Ipv4Addr::BROADCAST, Ipv4Addr::BROADCAST, gateway, 1)?;
tap_device.add_route(
Ipv4Addr::from([224, 0, 0, 0]),
Ipv4Addr::from([240, 0, 0, 0]),
gateway,
1,
)?;
delete_cache();
let tap = Arc::new(Device::Tap((tap_device, mac)));
let driver_info = DriverInfo {
device_type: DeviceType::Tap,
name,
version,
mac: Some(mac_str),
};
Ok((
DeviceWriter::new(tap.clone(), in_ips, address),
DeviceReader::new(tap),
driver_info,
))
}
fn delete_tap() {
let tap_device = match TapDevice::open(TAP_INTERFACE_NAME) {
Ok(tap_device) => tap_device,
Err(_) => {
return;
}
};
let _ = tap_device.delete();
}
pub fn create_device(
device_type: DeviceType,
address: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
in_ips: Vec<(Ipv4Addr, Ipv4Addr)>,
mtu: u16,
) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> {
match device_type {
DeviceType::Tun => create_tun(address, netmask, gateway, in_ips, mtu),
DeviceType::Tap => create_tap(address, netmask, gateway, in_ips, mtu),
}
}
pub fn delete_device(device_type: DeviceType) {
match device_type {
DeviceType::Tun => delete_tun(),
DeviceType::Tap => delete_tap(),
}
}
+138
View File
@@ -0,0 +1,138 @@
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
/// 不安全的并发计数器,谨慎使用
pub struct U64Adder {
global_index: Arc<AtomicUsize>,
inner: Arc<U64AdderInner>,
index: usize,
}
pub struct SingleU64Adder {
inner: Arc<SingleU64AdderInner>,
}
impl SingleU64Adder {
pub fn new() -> Self {
Self {
inner: Arc::new(SingleU64AdderInner::new()),
}
}
pub fn add(&mut self, num: u64) {
self.inner.add(num);
}
pub fn get(&self) -> u64 {
self.inner.get()
}
pub fn watch(&self) -> WatchSingleU64Adder {
WatchSingleU64Adder {
inner: self.inner.clone(),
}
}
}
struct SingleU64AdderInner {
ptr: *mut u64,
}
impl SingleU64AdderInner {
fn new() -> Self {
Self {
ptr: Box::into_raw(Box::new(0)),
}
}
#[inline(always)]
fn add(&self, num: u64) {
unsafe { *self.ptr += num }
}
fn get(&self) -> u64 {
unsafe { *self.ptr }
}
}
impl Drop for SingleU64AdderInner {
fn drop(&mut self) {
unsafe {
let _ = Box::from_raw(self.ptr);
}
}
}
unsafe impl Send for SingleU64AdderInner {}
unsafe impl Sync for SingleU64AdderInner {}
struct U64AdderInner {
base: Vec<SingleU64AdderInner>,
}
impl U64AdderInner {
pub fn get(&self) -> u64 {
let mut count = 0;
for counter in self.base.iter() {
count += counter.get()
}
count
}
}
impl U64Adder {
/// 计数槽容量
pub fn with_capacity(capacity: usize) -> Self {
let mut base = Vec::with_capacity(capacity);
for _ in 0..capacity {
base.push(SingleU64AdderInner::new())
}
let inner = Arc::new(U64AdderInner { base });
U64Adder {
global_index: Arc::new(AtomicUsize::new(1)),
inner,
index: 0,
}
}
pub fn add(&mut self, num: u64) {
self.inner.base[self.index].add(num);
}
pub fn get(&self) -> u64 {
self.inner.get()
}
pub fn watch(&self) -> WatchU64Adder {
WatchU64Adder {
inner: self.inner.clone(),
}
}
}
impl Clone for U64Adder {
fn clone(&self) -> Self {
let index = self.global_index.fetch_add(1, Ordering::AcqRel);
if index > self.inner.base.len() {
panic!()
}
Self {
global_index: self.global_index.clone(),
inner: self.inner.clone(),
index,
}
}
}
#[derive(Clone)]
pub struct WatchU64Adder {
inner: Arc<U64AdderInner>,
}
impl WatchU64Adder {
pub fn get(&self) -> u64 {
self.inner.get()
}
}
#[derive(Clone)]
pub struct WatchSingleU64Adder {
inner: Arc<SingleU64AdderInner>,
}
impl WatchSingleU64Adder {
pub fn get(&self) -> u64 {
self.inner.get()
}
}
+2
View File
@@ -0,0 +1,2 @@
mod adder;
pub use adder::*;
+9 -1
View File
@@ -1 +1,9 @@
pub mod wait;
mod notify;
mod result_convert;
pub use result_convert::io_convert;
mod scheduler;
pub use notify::StopManager;
pub use scheduler::Scheduler;
mod counter;
pub use counter::*;
+143
View File
@@ -0,0 +1,143 @@
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::Arc;
use std::thread::Thread;
use std::{io, thread};
use parking_lot::Mutex;
#[derive(Clone)]
pub struct StopManager {
inner: Arc<StopManagerInner>,
}
impl StopManager {
pub fn new<F>(f: F) -> Self
where
F: FnOnce() + Send + 'static,
{
Self {
inner: Arc::new(StopManagerInner::new(f)),
}
}
pub fn add_listener<F>(&self, name: String, f: F) -> io::Result<Worker>
where
F: FnOnce() + Send + 'static,
{
self.inner.add_listener(name, f)
}
pub fn stop(&self) {
self.inner.stop("");
}
pub fn wait(&self) {
self.inner.wait();
}
pub fn is_stop(&self) -> bool {
self.inner.state.load(Ordering::Acquire)
}
}
struct StopManagerInner {
listeners: Mutex<(bool, Vec<(String, Box<dyn FnOnce() + Send>)>)>,
park_threads: Mutex<Vec<Thread>>,
worker_num: AtomicUsize,
state: AtomicBool,
stop_call: Mutex<Option<Box<dyn FnOnce() + Send>>>,
}
impl StopManagerInner {
fn new<F>(f: F) -> Self
where
F: FnOnce() + Send + 'static,
{
Self {
listeners: Mutex::new((false, Vec::with_capacity(32))),
park_threads: Mutex::new(Vec::with_capacity(4)),
worker_num: AtomicUsize::new(0),
state: AtomicBool::new(false),
stop_call: Mutex::new(Some(Box::new(f))),
}
}
fn add_listener<F>(self: &Arc<Self>, name: String, f: F) -> io::Result<Worker>
where
F: FnOnce() + Send + 'static,
{
if name.is_empty() {
return Err(io::Error::new(io::ErrorKind::Other, "name cannot be empty"));
}
let mut guard = self.listeners.lock();
if guard.0 {
return Err(io::Error::new(io::ErrorKind::Other, "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),
));
}
}
guard.1.push((name.clone(), Box::new(f)));
Ok(Worker::new(name, self.clone()))
}
fn stop(&self, skip_name: &str) {
self.state.store(true, Ordering::Release);
let mut guard = self.listeners.lock();
guard.0 = true;
for (name, listener) in guard.1.drain(..) {
if &name == skip_name {
continue;
}
listener();
}
}
fn wait(&self) {
{
let mut guard = self.park_threads.lock();
guard.push(thread::current());
drop(guard);
}
loop {
if self.worker_num.load(Ordering::Acquire) == 0 {
return;
}
thread::park()
}
}
fn stop_call(&self) {
if let Some(call) = self.stop_call.lock().take() {
call();
}
}
}
pub struct Worker {
name: String,
inner: Arc<StopManagerInner>,
}
impl Worker {
fn new(name: String, inner: Arc<StopManagerInner>) -> Self {
let _ = inner.worker_num.fetch_add(1, Ordering::AcqRel);
Self { name, inner }
}
fn release0(&self) {
let inner = &self.inner;
let count = inner.worker_num.fetch_sub(1, Ordering::AcqRel);
if count == 1 {
for x in inner.park_threads.lock().drain(..) {
x.unpark();
}
self.inner.stop_call();
}
}
pub fn stop_all(self) {
self.inner.stop(&self.name)
}
}
impl Drop for Worker {
fn drop(&mut self) {
self.release0();
log::info!("stop {}", self.name);
}
}
+10
View File
@@ -0,0 +1,10 @@
use std::fmt::Display;
use std::io;
#[inline]
pub fn io_convert<T, R: Display, F: FnOnce(&io::Error) -> R>(
rs: io::Result<T>,
f: F,
) -> io::Result<T> {
rs.map_err(|e| io::Error::new(e.kind(), format!("{},internal error:{:?}", f(&e), e)))
}
+132
View File
@@ -0,0 +1,132 @@
use crate::util::StopManager;
use std::collections::BinaryHeap;
use std::{
cmp::Ordering,
io,
sync::mpsc::{sync_channel, Receiver, SyncSender},
time::{Duration, Instant},
};
struct DelayedTask {
f: Box<dyn FnOnce(&Scheduler) + Send>,
next: Instant,
}
impl Eq for DelayedTask {}
impl PartialEq for DelayedTask {
fn eq(&self, other: &Self) -> bool {
self.next.eq(&other.next)
}
}
impl PartialOrd for DelayedTask {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
self.next.partial_cmp(&other.next).map(|ord| ord.reverse())
}
}
impl Ord for DelayedTask {
fn cmp(&self, other: &Self) -> Ordering {
self.next.cmp(&other.next).reverse()
}
}
enum Op {
Task(DelayedTask),
Stop,
}
#[derive(Clone)]
pub struct Scheduler {
sender: SyncSender<Op>,
}
impl Scheduler {
pub fn new(stop_manager: StopManager) -> io::Result<Self> {
let (sender, receiver) = sync_channel::<Op>(32);
let s = Self { sender };
let s_inner = s.clone();
let worker = {
let scheduler = s.clone();
stop_manager.add_listener("Scheduler".into(), move || {
scheduler.shutdown();
})?
};
std::thread::Builder::new()
.name("Scheduler".into())
.spawn(move || {
run(receiver, s_inner);
worker.stop_all();
})
.expect("Scheduler");
Ok(s)
}
pub fn timeout<F>(&self, time: Duration, f: F) -> bool
where
F: FnOnce(&Scheduler) + Send + 'static,
{
let task = DelayedTask {
f: Box::new(f),
next: Instant::now().checked_add(time).unwrap(),
};
self.sender.send(Op::Task(task)).is_ok()
}
pub fn shutdown(self) {
let _ = self.sender.send(Op::Stop);
}
}
fn run(receiver: Receiver<Op>, s_inner: Scheduler) {
let mut binary_heap = BinaryHeap::<DelayedTask>::with_capacity(32);
loop {
while let Some(task) = binary_heap.peek() {
let now = Instant::now();
if now < task.next {
//需要等待对应时间
match receiver.recv_timeout(task.next - now) {
Ok(op) => {
if add_task(op, &mut binary_heap) {
continue;
}
return;
}
Err(e) => match e {
std::sync::mpsc::RecvTimeoutError::Timeout => continue,
std::sync::mpsc::RecvTimeoutError::Disconnected => return,
},
}
} else {
if let Some(task) = binary_heap.pop() {
(task.f)(&s_inner);
}
}
}
//取出所有任务
loop {
match receiver.try_recv() {
Ok(op) => {
if add_task(op, &mut binary_heap) {
continue;
}
return;
}
Err(e) => match e {
std::sync::mpsc::TryRecvError::Empty => break,
std::sync::mpsc::TryRecvError::Disconnected => return,
},
}
}
if binary_heap.is_empty() {
//任务队列为空时陷入等待
if let Ok(op) = receiver.recv() {
if add_task(op, &mut binary_heap) {
continue;
}
}
return;
}
}
}
fn add_task(op: Op, binary_heap: &mut BinaryHeap<DelayedTask>) -> bool {
return match op {
Op::Task(task) => {
binary_heap.push(task);
true
}
Op::Stop => false,
};
}
-44
View File
@@ -1,44 +0,0 @@
use std::sync::atomic::{AtomicIsize, Ordering};
use std::sync::Arc;
use tokio::sync::watch::{channel, Receiver, Sender};
#[derive(Clone)]
pub struct WaitGroup {
count: Arc<AtomicIsize>,
receiver: Receiver<usize>,
sender: Arc<Sender<usize>>,
}
impl WaitGroup {
pub fn new() -> Self {
let (sender, receiver) = channel(1);
Self {
count: Arc::new(Default::default()),
receiver,
sender: Arc::new(sender),
}
}
pub fn add(&self) {
let _ = self.count.fetch_add(1, Ordering::Relaxed);
}
pub fn done(&self) {
let i = self.count.fetch_sub(1, Ordering::Relaxed);
if i == 1 {
let _ = self.sender.send(0);
}
}
pub async fn wait(&mut self) {
loop {
if 0 == *self.receiver.borrow() {
return;
}
if self.receiver.changed().await.is_ok() {
if 0 == *self.receiver.borrow() {
return;
}
} else {
return;
}
}
}
}
+32
View File
@@ -0,0 +1,32 @@
[package]
name = "tun"
version = "0.1.0"
edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
[dependencies]
libc = "0.2.153"
log = { version = "0.4.20", features = [] }
rand = "0.8.5"
[target.'cfg(any(target_os = "linux", target_os = "macos"))'.dependencies]
ioctl = { version = "0.8", package = "ioctl-sys" }
[target.'cfg(target_os = "windows")'.dependencies]
libloading = "0.8.0"
widestring = "1.0.2"
winapi = {version = "0.3",features = [
"errhandlingapi",
"combaseapi",
"ioapiset",
"winioctl",
"setupapi",
"synchapi",
"netioapi",
"fileapi","handleapi","winerror","minwindef","ifdef","basetsd","winnt","winreg","winbase","minwinbase",
"impl-default"
]}
+56
View File
@@ -0,0 +1,56 @@
use crate::device::IFace;
use crate::Fd;
use std::io;
use std::net::Ipv4Addr;
use std::os::fd::RawFd;
pub struct Device {
fd: Fd,
}
impl Device {
pub fn new(fd: RawFd) -> io::Result<Self> {
Ok(Self { fd: Fd::new(fd)? })
}
}
impl IFace for Device {
fn version(&self) -> io::Result<String> {
Ok(String::new())
}
fn name(&self) -> io::Result<String> {
Ok(String::new())
}
fn shutdown(&self) -> io::Result<()> {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
fn mtu(&self) -> io::Result<u32> {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
fn set_mtu(&self, value: u32) -> io::Result<()> {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, metric: u16) -> io::Result<()> {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
self.fd.read(buf)
}
fn write(&self, buf: &[u8]) -> io::Result<usize> {
self.fd.write(buf)
}
}
+24
View File
@@ -0,0 +1,24 @@
use io::Result;
use std::io;
use std::net::Ipv4Addr;
pub trait IFace {
fn version(&self) -> Result<String>;
/// Get the device name.
fn name(&self) -> Result<String>;
fn shutdown(&self) -> Result<()>;
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> Result<()>;
/// Get the MTU.
fn mtu(&self) -> Result<u32>;
/// Set the MTU.
fn set_mtu(&self, value: u32) -> Result<()>;
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, metric: u16) -> Result<()>;
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr) -> Result<()>;
fn read(&self, buf: &mut [u8]) -> Result<usize>;
fn write(&self, buf: &[u8]) -> Result<usize>;
}
+31
View File
@@ -0,0 +1,31 @@
/// 参考
/// https://github.com/meh/rust-tun
/// https://github.com/Tazdevil971/tap-windows
/// https://github.com/nulldotblack/wintun
pub mod device;
mod packet;
#[cfg(target_os = "linux")]
mod linux;
#[cfg(target_os = "linux")]
pub use linux::Device;
#[cfg(target_os = "android")]
mod android;
#[cfg(target_os = "android")]
pub use android::Device;
#[cfg(target_os = "macos")]
mod macos;
#[cfg(target_os = "macos")]
pub use macos::Device;
#[cfg(unix)]
mod unix;
#[cfg(unix)]
pub use unix::Fd;
#[cfg(windows)]
mod windows;
#[cfg(windows)]
pub use windows::Device;
+307
View File
@@ -0,0 +1,307 @@
use std::ffi::{CStr, CString};
use std::net::Ipv4Addr;
use std::os::fd::AsRawFd;
use std::process::Command;
use std::{io, mem, ptr};
use libc::{
c_char, c_short, ifreq, AF_INET, IFF_MULTI_QUEUE, IFF_NO_PI, IFF_RUNNING, IFF_TAP, IFF_TUN,
IFF_UP, IFNAMSIZ, O_RDWR, SOCK_DGRAM,
};
use crate::device::IFace;
use crate::linux::route;
use crate::linux::sys::*;
use crate::packet;
use crate::unix::{exe_cmd, Fd, SockAddr};
pub struct Device {
name: String,
ctl: Fd,
tun: Fd,
mac: Option<[u8; 6]>,
}
impl Device {
pub fn new(name: Option<String>, tap: bool) -> io::Result<Self> {
let device = unsafe {
let dev = match name {
Some(name) => {
let name =
CString::new(name).map_err(|e| io::Error::new(io::ErrorKind::Other, e))?;
if name.as_bytes_with_nul().len() > IFNAMSIZ {
return Err(io::Error::new(io::ErrorKind::InvalidInput, "name too long"));
}
Some(name)
}
None => None,
};
let mut req: ifreq = mem::zeroed();
if let Some(dev) = dev.as_ref() {
ptr::copy_nonoverlapping(
dev.as_ptr() as *const c_char,
req.ifr_name.as_mut_ptr(),
dev.as_bytes().len(),
);
}
let device_type: c_short = if tap { IFF_TAP } else { IFF_TUN } as c_short;
let queues_num = 1;
let iff_no_pi = IFF_NO_PI as c_short;
let iff_multi_queue = IFF_MULTI_QUEUE as c_short;
let packet_information = false;
req.ifr_ifru.ifru_flags = device_type
| if packet_information { 0 } else { iff_no_pi }
| if queues_num > 1 { iff_multi_queue } else { 0 };
let tun = Fd::new(libc::open(b"/dev/net/tun\0".as_ptr() as *const _, O_RDWR))
.map_err(|_| io::Error::last_os_error())?;
if tunsetiff(tun.0, &mut req as *mut _ as *mut _) < 0 {
return Err(io::Error::last_os_error());
}
let ctl = Fd::new(libc::socket(AF_INET, SOCK_DGRAM, 0))?;
let name = CStr::from_ptr(req.ifr_name.as_ptr())
.to_string_lossy()
.to_string();
let mac = if tap {
let get_mac_cmd = format!("cat /sys/class/net/{}/address", name);
let mac_out = exe_cmd(&get_mac_cmd)?;
let mac_str = String::from_utf8(mac_out.stdout).unwrap();
let mut mac = [0; 6];
let mut split = mac_str.split(":");
for i in 0..6 {
mac[i] = u8::from_str_radix(&split.next().unwrap()[..2], 16).unwrap();
}
Some(mac)
} else {
None
};
let set_txqueuelen = format!("ifconfig {} txqueuelen 1000", name);
if let Err(e) = exe_cmd(&set_txqueuelen){
log::warn!("{:?}",e);
}
Device {
name,
tun,
ctl,
mac,
}
};
device.enabled(true)?;
Ok(device)
}
}
impl Device {
fn enabled(&self, value: bool) -> io::Result<()> {
unsafe {
let mut req = self.request();
if siocgifflags(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
if value {
req.ifr_ifru.ifru_flags |= (IFF_UP | IFF_RUNNING) as c_short;
} else {
req.ifr_ifru.ifru_flags &= !(IFF_UP as c_short);
}
if siocsifflags(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
unsafe fn request(&self) -> ifreq {
let mut req: ifreq = mem::zeroed();
ptr::copy_nonoverlapping(
self.name.as_ptr() as *const c_char,
req.ifr_name.as_mut_ptr(),
self.name.len(),
);
req
}
fn address(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error().into());
}
SockAddr::new(&req.ifr_ifru.ifru_addr).map(Into::into)
}
}
fn set_address(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifr_ifru.ifru_addr = SockAddr::from(value).into();
if siocsifaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn destination(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifdstaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
SockAddr::new(&req.ifr_ifru.ifru_dstaddr).map(Into::into)
}
}
fn set_destination(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifr_ifru.ifru_dstaddr = SockAddr::from(value).into();
if siocsifdstaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn broadcast(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifbrdaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
SockAddr::new(&req.ifr_ifru.ifru_broadaddr).map(Into::into)
}
}
fn set_broadcast(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifr_ifru.ifru_broadaddr = SockAddr::from(value).into();
if siocsifbrdaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn netmask(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifnetmask(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
SockAddr::new(&req.ifr_ifru.ifru_netmask).map(Into::into)
}
}
fn set_netmask(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifr_ifru.ifru_netmask = SockAddr::from(value).into();
if siocsifnetmask(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
}
impl IFace for Device {
fn version(&self) -> io::Result<String> {
Ok(String::new())
}
fn name(&self) -> io::Result<String> {
Ok(self.name.clone())
}
fn shutdown(&self) -> io::Result<()> {
exe_cmd(&format!("ip link delete {}", self.name))?;
Ok(())
}
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
self.set_address(address)?;
self.set_netmask(mask)
}
fn mtu(&self) -> io::Result<u32> {
unsafe {
let mut req = self.request();
if siocgifmtu(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(req.ifr_ifru.ifru_mtu as u32)
}
}
fn set_mtu(&self, value: u32) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifr_ifru.ifru_mtu = value as _;
if siocsifmtu(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, _metric: u16) -> io::Result<()> {
route::add_route(&self.name, dest, netmask)
}
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
route::del_route(&self.name, dest, netmask)
}
fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
if self.mac.is_some() {
packet::read_tap(
buf,
|eth_buf| self.tun.read(eth_buf),
|eth_buf| self.tun.write(eth_buf),
)
} else {
self.tun.read(buf)
}
}
fn write(&self, buf: &[u8]) -> io::Result<usize> {
if let Some(mac) = &self.mac {
packet::write_tap(buf, |eth_buf| self.tun.write(eth_buf), mac)
} else {
self.tun.write(buf)
}
}
}
+4
View File
@@ -0,0 +1,4 @@
mod device;
pub use device::Device;
mod route;
mod sys;
+16
View File
@@ -0,0 +1,16 @@
use std::io;
use std::net::Ipv4Addr;
use crate::unix::exe_cmd;
pub fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
let cmd = format!("ip route add {:?}/{:?} dev {}", address, netmask, name);
exe_cmd(&cmd)?;
Ok(())
}
pub fn del_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
let cmd = format!("ip route del {:?}/{:?} dev {}", address, netmask, name);
exe_cmd(&cmd)?;
Ok(())
}
+21
View File
@@ -0,0 +1,21 @@
use ioctl::*;
use libc::{c_int, ifreq};
ioctl!(bad read siocgifflags with 0x8913; ifreq);
ioctl!(bad write siocsifflags with 0x8914; ifreq);
ioctl!(bad read siocgifaddr with 0x8915; ifreq);
ioctl!(bad write siocsifaddr with 0x8916; ifreq);
ioctl!(bad read siocgifdstaddr with 0x8917; ifreq);
ioctl!(bad write siocsifdstaddr with 0x8918; ifreq);
ioctl!(bad read siocgifbrdaddr with 0x8919; ifreq);
ioctl!(bad write siocsifbrdaddr with 0x891a; ifreq);
ioctl!(bad read siocgifnetmask with 0x891b; ifreq);
ioctl!(bad write siocsifnetmask with 0x891c; ifreq);
ioctl!(bad read siocgifmtu with 0x8921; ifreq);
ioctl!(bad write siocsifmtu with 0x8922; ifreq);
ioctl!(bad write siocsifname with 0x8923; ifreq);
ioctl!(write tunsetiff with b'T', 202; c_int);
ioctl!(write tunsetpersist with b'T', 203; c_int);
ioctl!(write tunsetowner with b'T', 204; c_int);
ioctl!(write tunsetgroup with b'T', 206; c_int);
+291
View File
@@ -0,0 +1,291 @@
use std::ffi::{c_void, CStr};
use std::net::Ipv4Addr;
use std::os::fd::AsRawFd;
use std::{io, mem, ptr};
use libc::{
c_char, c_short, c_uint, sockaddr, socklen_t, AF_INET, AF_SYSTEM, AF_SYS_CONTROL, IFF_RUNNING,
IFF_UP, IFNAMSIZ, PF_SYSTEM, SOCK_DGRAM, SYSPROTO_CONTROL, UTUN_OPT_IFNAME,
};
use crate::device::IFace;
use crate::macos::route;
use crate::macos::sys::*;
use crate::unix::{Fd, SockAddr};
pub struct Device {
name: String,
ctl: Fd,
tun: Fd,
}
impl Device {
pub fn new(name: Option<String>) -> io::Result<Self> {
let id = if let Some(name) = name {
if name.len() > IFNAMSIZ {
return Err(io::Error::new(io::ErrorKind::InvalidInput, "name too long"));
}
if !name.starts_with("utun") {
return Err(io::Error::new(io::ErrorKind::InvalidInput, "invalid name"));
}
name[4..]
.parse::<u32>()
.map_err(|e| io::Error::new(io::ErrorKind::Other, e))?
+ 1u32
} else {
0u32
};
let device = unsafe {
let tun = Fd::new(libc::socket(PF_SYSTEM, SOCK_DGRAM, SYSPROTO_CONTROL))?;
let mut info = ctl_info {
ctl_id: 0,
ctl_name: {
let mut buffer = [0; 96];
for (i, o) in UTUN_CONTROL_NAME.as_bytes().iter().zip(buffer.iter_mut()) {
*o = *i as _;
}
buffer
},
};
if ctliocginfo(tun.0, &mut info as *mut _ as *mut _) < 0 {
return Err(io::Error::last_os_error());
}
let addr = sockaddr_ctl {
sc_id: info.ctl_id,
sc_len: mem::size_of::<sockaddr_ctl>() as _,
sc_family: AF_SYSTEM as _,
ss_sysaddr: AF_SYS_CONTROL as _,
sc_unit: id as c_uint,
sc_reserved: [0; 5],
};
let address = &addr as *const sockaddr_ctl as *const sockaddr;
if libc::connect(tun.0, address, mem::size_of_val(&addr) as socklen_t) < 0 {
return Err(io::Error::last_os_error());
}
let mut name = [0u8; 64];
let mut name_len: socklen_t = 64;
let optval = &mut name as *mut _ as *mut c_void;
let optlen = &mut name_len as *mut socklen_t;
if libc::getsockopt(tun.0, SYSPROTO_CONTROL, UTUN_OPT_IFNAME, optval, optlen) < 0 {
return Err(io::Error::last_os_error());
}
let ctl = Fd::new(libc::socket(AF_INET, SOCK_DGRAM, 0))?;
Device {
name: CStr::from_ptr(name.as_ptr() as *const c_char)
.to_string_lossy()
.into(),
tun,
ctl,
}
};
device.enabled(true)?;
Ok(device)
}
}
impl Device {
fn enabled(&self, value: bool) -> io::Result<()> {
unsafe {
let mut req = self.request();
if siocgifflags(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
if value {
req.ifru.flags |= (IFF_UP | IFF_RUNNING) as c_short;
} else {
req.ifru.flags &= !(IFF_UP as c_short);
}
if siocsifflags(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
unsafe fn request(&self) -> ifreq {
let mut req: ifreq = mem::zeroed();
ptr::copy_nonoverlapping(
self.name.as_ptr() as *const c_char,
req.ifrn.name.as_mut_ptr(),
self.name.len(),
);
req
}
fn address(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
SockAddr::new(&req.ifru.addr).map(Into::into)
}
}
fn set_address(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifru.addr = SockAddr::from(value).into();
if siocsifaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn destination(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifdstaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
SockAddr::new(&req.ifru.dstaddr).map(Into::into)
}
}
fn set_destination(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifru.dstaddr = SockAddr::from(value).into();
if siocsifdstaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn broadcast(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifbrdaddr(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
SockAddr::new(&req.ifru.broadaddr).map(Into::into)
}
}
fn set_broadcast(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifru.broadaddr = SockAddr::from(value).into();
if siocsifbrdaddr(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn netmask(&self) -> io::Result<Ipv4Addr> {
unsafe {
let mut req = self.request();
if siocgifnetmask(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
SockAddr::unchecked(&req.ifru.addr).map(Into::into)
}
}
fn set_netmask(&self, value: Ipv4Addr) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifru.addr = SockAddr::from(value).into();
if siocsifnetmask(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
}
impl IFace for Device {
fn version(&self) -> io::Result<String> {
Ok(String::new())
}
fn name(&self) -> io::Result<String> {
Ok(self.name.clone())
}
fn shutdown(&self) -> io::Result<()> {
Ok(())
}
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
self.set_address(address)?;
self.set_netmask(mask)
}
fn mtu(&self) -> io::Result<u32> {
unsafe {
let mut req = self.request();
if siocgifmtu(self.ctl.as_raw_fd(), &mut req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(req.ifru.mtu as _)
}
}
fn set_mtu(&self, value: u32) -> io::Result<()> {
unsafe {
let mut req = self.request();
req.ifru.mtu = value as _;
if siocsifmtu(self.ctl.as_raw_fd(), &req) < 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, _metric: u16) -> io::Result<()> {
route::add_route(&self.name, dest, netmask)
}
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
route::del_route(&self.name, dest, netmask)
}
fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
self.tun.read(buf)
}
fn write(&self, buf: &[u8]) -> io::Result<usize> {
let mut packet = Vec::<u8>::with_capacity(4 + buf.len());
packet.push(0);
packet.push(0);
packet.extend_from_slice(&(libc::PF_INET as u16).to_be_bytes());
packet.extend_from_slice(buf);
self.tun.write(&packet)
}
}
+5
View File
@@ -0,0 +1,5 @@
mod device;
pub use device::Device;
mod sys;
mod route;
+20
View File
@@ -0,0 +1,20 @@
use crate::unix::exe_cmd;
use std::io;
use std::net::Ipv4Addr;
pub fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
let cmd = format!(
"route -n add {} -netmask {} -interface {}",
address, netmask, name
);
exe_cmd(&cmd)?;
Ok(())
}
pub fn del_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
let cmd = format!(
"route -n delete {} -netmask {} -interface {}",
address, netmask, name
);
exe_cmd(&cmd)?;
Ok(())
}
@@ -1,35 +1,11 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
//! Bindings to internal macOS stuff.
use ioctl::*;
use libc::sockaddr;
use libc::{c_char, c_int, c_short, c_uint, c_ushort, c_void};
use libc::{c_char, c_int, c_short, c_uint, c_ushort, c_void, sockaddr, IFNAMSIZ};
pub const IFNAMSIZ: usize = 16;
pub const IFF_UP: c_short = 0x1;
pub const IFF_RUNNING: c_short = 0x40;
pub const AF_SYS_CONTROL: c_ushort = 2;
pub const AF_SYSTEM: c_char = 32;
pub const PF_SYSTEM: c_int = AF_SYSTEM as c_int;
pub const SYSPROTO_CONTROL: c_int = 2;
pub const UTUN_OPT_IFNAME: c_int = 2;
pub const UTUN_CONTROL_NAME: &str = "com.apple.net.utun_control";
#[allow(non_camel_case_types)]
#[repr(C)]
#[derive(Copy, Clone)]
pub struct ctl_info {
@@ -37,6 +13,7 @@ pub struct ctl_info {
pub ctl_name: [c_char; 96],
}
#[allow(non_camel_case_types)]
#[repr(C)]
#[derive(Copy, Clone)]
pub struct sockaddr_ctl {
@@ -54,6 +31,7 @@ pub union ifrn {
pub name: [c_char; IFNAMSIZ],
}
#[allow(non_camel_case_types)]
#[repr(C)]
#[derive(Copy, Clone)]
pub struct ifdevmtu {
@@ -69,6 +47,7 @@ pub union ifku {
pub value: c_int,
}
#[allow(non_camel_case_types)]
#[repr(C)]
#[derive(Copy, Clone)]
pub struct ifkpi {
@@ -98,6 +77,7 @@ pub union ifru {
pub functional_type: c_uint,
}
#[allow(non_camel_case_types)]
#[repr(C)]
#[derive(Copy, Clone)]
pub struct ifreq {
@@ -105,6 +85,7 @@ pub struct ifreq {
pub ifru: ifru,
}
#[allow(non_camel_case_types)]
#[repr(C)]
#[derive(Copy, Clone)]
pub struct ifaliasreq {
+1
View File
@@ -0,0 +1 @@
pub mod packet;
+122
View File
@@ -0,0 +1,122 @@
use std::{fmt, io};
/// 地址解析协议,由IP地址找到MAC地址
/// https://www.ietf.org/rfc/rfc6747.txt
/*
0 2 4 5 6 8 10 ()
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| |||||
| MAC地址 | ip地址 |
| MAC地址 | ip地址 |
*/
pub struct ArpPacket<B> {
buffer: B,
}
impl<B: AsRef<[u8]>> ArpPacket<B> {
pub fn unchecked(buffer: B) -> Self {
Self { buffer }
}
pub fn new(buffer: B) -> io::Result<Self> {
if buffer.as_ref().len() != 28 {
Err(io::Error::from(io::ErrorKind::InvalidData))?;
}
let packet = Self::unchecked(buffer);
Ok(packet)
}
}
impl<B: AsRef<[u8]>> ArpPacket<B> {
/// 硬件类型 以太网类型为1
pub fn hardware_type(&self) -> u16 {
u16::from_be_bytes(self.buffer.as_ref()[0..2].try_into().unwrap())
}
/// 上层协议类型,ipv4是0x0800
pub fn protocol_type(&self) -> u16 {
u16::from_be_bytes(self.buffer.as_ref()[2..4].try_into().unwrap())
}
/// 如果是MAC地址 则长度为6
pub fn hardware_size(&self) -> u8 {
self.buffer.as_ref()[4]
}
/// 如果是IPv4 则长度为4
pub fn protocol_size(&self) -> u8 {
self.buffer.as_ref()[5]
}
/// 操作类型,请求和响应 1:ARP请求,2:ARP响应,3RARP请求,4RARP响应
pub fn op_code(&self) -> u16 {
u16::from_be_bytes(self.buffer.as_ref()[6..8].try_into().unwrap())
}
/// 发送端硬件地址,仅支持以太网
pub fn sender_hardware_addr(&self) -> &[u8] {
&self.buffer.as_ref()[8..14]
}
/// 发送端协议地址,仅支持IPv4
pub fn sender_protocol_addr(&self) -> &[u8] {
&self.buffer.as_ref()[14..18]
}
/// 接收端硬件地址,仅支持以太网
pub fn target_hardware_addr(&self) -> &[u8] {
&self.buffer.as_ref()[18..24]
}
/// 接收端协议地址,仅支持IPv4
pub fn target_protocol_addr(&self) -> &[u8] {
&self.buffer.as_ref()[24..28]
}
}
impl<B: AsRef<[u8]> + AsMut<[u8]>> ArpPacket<B> {
/// 硬件类型 以太网类型为1
pub fn set_hardware_type(&mut self, value: u16) {
self.buffer.as_mut()[0..2].copy_from_slice(&value.to_be_bytes())
}
/// 上层协议类型,ipv4是0x0800
pub fn set_protocol_type(&mut self, value: u16) {
self.buffer.as_mut()[2..4].copy_from_slice(&value.to_be_bytes())
}
/// 如果是MAC地址 则长度为6
pub fn set_hardware_size(&mut self, value: u8) {
self.buffer.as_mut()[4] = value
}
/// 如果是IPv4 则长度为4
pub fn set_protocol_size(&mut self, value: u8) {
self.buffer.as_mut()[5] = value
}
/// 操作类型,请求和响应 1:ARP请求,2:ARP响应,3RARP请求,4RARP响应
pub fn set_op_code(&mut self, value: u16) {
self.buffer.as_mut()[6..8].copy_from_slice(&value.to_be_bytes())
}
/// 发送端硬件地址,仅支持以太网
pub fn set_sender_hardware_addr(&mut self, buf: &[u8]) {
self.buffer.as_mut()[8..14].copy_from_slice(buf)
}
/// 发送端协议地址,仅支持IPv4
pub fn set_sender_protocol_addr(&mut self, buf: &[u8]) {
self.buffer.as_mut()[14..18].copy_from_slice(buf)
}
/// 接收端硬件地址,仅支持以太网
pub fn set_target_hardware_addr(&mut self, buf: &[u8]) {
self.buffer.as_mut()[18..24].copy_from_slice(buf)
}
/// 接收端协议地址,仅支持IPv4
pub fn set_target_protocol_addr(&mut self, buf: &[u8]) {
self.buffer.as_mut()[24..28].copy_from_slice(buf)
}
}
impl<B: AsRef<[u8]>> fmt::Debug for ArpPacket<B> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ArpPacket")
.field("hardware_type", &self.hardware_type())
.field("protocol_type", &self.protocol_type())
.field("hardware_size", &self.hardware_size())
.field("protocol_size", &self.protocol_size())
.field("op_code", &self.op_code())
.field("sender_hardware_addr", &self.sender_hardware_addr())
.field("sender_protocol_addr", &self.sender_protocol_addr())
.field("target_hardware_addr", &self.target_hardware_addr())
.field("target_protocol_addr", &self.target_protocol_addr())
.finish()
}
}
+2
View File
@@ -0,0 +1,2 @@
pub mod packet;
pub mod protocol;
+77
View File
@@ -0,0 +1,77 @@
use crate::packet::ethernet::protocol::Protocol;
use std::{fmt, io};
/// 以太网帧协议
/// https://www.ietf.org/rfc/rfc894.txt
/*
0 6 12 14 ()
+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+
| | | |
*/
pub struct EthernetPacket<B> {
pub buffer: B,
}
impl<B: AsRef<[u8]>> EthernetPacket<B> {
pub fn unchecked(buffer: B) -> EthernetPacket<B> {
EthernetPacket { buffer }
}
pub fn new(buffer: B) -> io::Result<EthernetPacket<B>> {
let packet = EthernetPacket::unchecked(buffer);
//头部固定14位
if packet.buffer.as_ref().len() < 14 {
Err(io::Error::new(io::ErrorKind::InvalidData,format!("len={}", packet.buffer.as_ref().len())))?;
}
Ok(packet)
}
}
impl<B: AsRef<[u8]>> EthernetPacket<B> {
/// 目的MAC地址
pub fn destination(&self) -> &[u8] {
&self.buffer.as_ref()[0..6]
}
/// 源MAC地址
pub fn source(&self) -> &[u8] {
&self.buffer.as_ref()[6..12]
}
/// 3层协议
pub fn protocol(&self) -> Protocol {
u16::from_be_bytes(self.buffer.as_ref()[12..14].try_into().unwrap()).into()
}
/// 载荷
pub fn payload(&self) -> &[u8] {
&self.buffer.as_ref()[14..]
}
}
impl<B: AsRef<[u8]> + AsMut<[u8]>> EthernetPacket<B> {
pub fn set_destination(&mut self, value: &[u8]) {
self.buffer.as_mut()[0..6].copy_from_slice(value);
}
pub fn set_source(&mut self, value: &[u8]) {
self.buffer.as_mut()[6..12].copy_from_slice(value);
}
pub fn set_protocol(&mut self, value: Protocol) {
let p: u16 = value.into();
self.buffer.as_mut()[12..14].copy_from_slice(&p.to_be_bytes())
}
pub fn payload_mut(&mut self) -> &mut [u8] {
&mut self.buffer.as_mut()[14..]
}
}
impl<B: AsRef<[u8]>> fmt::Debug for EthernetPacket<B> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EthernetPacket")
.field("destination", &self.destination())
.field("source", &self.source())
.field("protocol", &self.protocol())
.field("payload", &self.payload())
.finish()
}
}
+141
View File
@@ -0,0 +1,141 @@
/// 以太网帧协议
#[derive(Eq, PartialEq, Copy, Clone, Debug)]
pub enum Protocol {
///
Ipv4,
///
Arp,
///
WakeOnLan,
///
Trill,
///
DecNet,
///
Rarp,
///
AppleTalk,
///
Aarp,
///
Ipx,
///
Qnx,
///
Ipv6,
///
FlowControl,
///
CobraNet,
///
Mpls,
///
MplsMulticast,
///
PppoeDiscovery,
///
PppoeSession,
///
Vlan,
///
PBridge,
///
Lldp,
///
Ptp,
///
Cfm,
///
QinQ,
///
Unknown(u16),
}
impl From<u16> for Protocol {
fn from(value: u16) -> Protocol {
use self::Protocol::*;
match value {
0x0800 => Ipv4,
0x0806 => Arp,
0x0842 => WakeOnLan,
0x22f3 => Trill,
0x6003 => DecNet,
0x8035 => Rarp,
0x809b => AppleTalk,
0x80f3 => Aarp,
0x8137 => Ipx,
0x8204 => Qnx,
0x86dd => Ipv6,
0x8808 => FlowControl,
0x8819 => CobraNet,
0x8847 => Mpls,
0x8848 => MplsMulticast,
0x8863 => PppoeDiscovery,
0x8864 => PppoeSession,
0x8100 => Vlan,
0x88a8 => PBridge,
0x88cc => Lldp,
0x88f7 => Ptp,
0x8902 => Cfm,
0x9100 => QinQ,
n => Unknown(n),
}
}
}
impl Into<u16> for Protocol {
fn into(self) -> u16 {
use self::Protocol::*;
match self {
Ipv4 => 0x0800,
Arp => 0x0806,
WakeOnLan => 0x0842,
Trill => 0x22f3,
DecNet => 0x6003,
Rarp => 0x8035,
AppleTalk => 0x809b,
Aarp => 0x80f3,
Ipx => 0x8137,
Qnx => 0x8204,
Ipv6 => 0x86dd,
FlowControl => 0x8808,
CobraNet => 0x8819,
Mpls => 0x8847,
MplsMulticast => 0x8848,
PppoeDiscovery => 0x8863,
PppoeSession => 0x8864,
Vlan => 0x8100,
PBridge => 0x88a8,
Lldp => 0x88cc,
Ptp => 0x88f7,
Cfm => 0x8902,
QinQ => 0x9100,
Unknown(n) => n,
}
}
}
+67
View File
@@ -0,0 +1,67 @@
use crate::packet::ethernet::protocol::Protocol;
use std::io;
pub mod arp;
pub mod ethernet;
const MAC: [u8; 6] = [0xf, 0xf, 0xf, 0xf, 0xe, 0x9];
pub fn read_tap<W, R>(buf: &mut [u8], read_fn: R, write_fn: W) -> io::Result<usize>
where
W: Fn(&[u8]) -> io::Result<usize>,
R: Fn(&mut [u8]) -> io::Result<usize>,
{
let mut eth_buf = [0; 65536];
loop {
let len = read_fn(&mut eth_buf)?;
if len == 0{
return Ok(len);
}
//处理arp包
let mut ether = ethernet::packet::EthernetPacket::new(&mut eth_buf[..len])?;
match ether.protocol() {
Protocol::Ipv4 => {
let len = ether.payload().len();
if len > buf.len() {
return Err(io::Error::new(io::ErrorKind::Other, "short"));
}
buf[..len].copy_from_slice(ether.payload());
return Ok(len);
}
Protocol::Arp => {
let mut arp_packet = arp::packet::ArpPacket::unchecked(ether.payload_mut());
let sender_h: [u8; 6] = arp_packet.sender_hardware_addr().try_into().unwrap();
let sender_p: [u8; 4] = arp_packet.sender_protocol_addr().try_into().unwrap();
let target_p: [u8; 4] = arp_packet.target_protocol_addr().try_into().unwrap();
if target_p == [0, 0, 0, 0] || sender_p == [0, 0, 0, 0] || target_p == sender_p {
continue;
}
if arp_packet.op_code() == 1 {
//回复一个默认的MAC
arp_packet.set_op_code(2);
arp_packet.set_target_hardware_addr(&sender_h);
arp_packet.set_target_protocol_addr(&sender_p);
arp_packet.set_sender_protocol_addr(&target_p);
arp_packet.set_sender_hardware_addr(&MAC);
ether.set_destination(&sender_h);
ether.set_source(&MAC);
write_fn(ether.buffer)?;
}
}
_ => {
//忽略这些数据
}
}
}
}
pub fn write_tap<W>(buf: &[u8], write_fn: W, mac: &[u8; 6]) -> io::Result<usize>
where
W: Fn(&[u8]) -> io::Result<usize>,
{
// 封装二层数据
let mut ether = ethernet::packet::EthernetPacket::unchecked(vec![0; 14 + buf.len()]);
ether.set_source(&MAC);
ether.set_destination(mac);
ether.set_protocol(Protocol::Ipv4);
ether.payload_mut().copy_from_slice(buf);
write_fn(&ether.buffer)
}
+62
View File
@@ -0,0 +1,62 @@
use std::io;
use std::os::fd::{AsRawFd, IntoRawFd, RawFd};
pub struct Fd(pub RawFd);
impl Fd {
pub fn new(value: RawFd) -> io::Result<Self> {
if value < 0 {
return Err(io::Error::from(io::ErrorKind::InvalidInput));
}
Ok(Fd(value))
}
}
impl Fd {
pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
unsafe {
let amount = libc::read(self.0, buf.as_mut_ptr() as *mut _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(amount as usize)
}
}
pub fn write(&self, buf: &[u8]) -> io::Result<usize> {
unsafe {
let amount = libc::write(self.0, buf.as_ptr() as *const _, buf.len());
if amount < 0 {
return Err(io::Error::last_os_error());
}
Ok(amount as usize)
}
}
}
impl AsRawFd for Fd {
fn as_raw_fd(&self) -> RawFd {
self.0
}
}
impl IntoRawFd for Fd {
fn into_raw_fd(mut self) -> RawFd {
let fd = self.0;
self.0 = -1;
fd
}
}
impl Drop for Fd {
fn drop(&mut self) {
unsafe {
if self.0 >= 0 {
libc::close(self.0);
}
}
}
}
+27
View File
@@ -0,0 +1,27 @@
mod fd;
pub use fd::Fd;
use std::process::Output;
#[cfg(any(target_os = "macos", target_os = "linux"))]
mod sockaddr;
#[cfg(any(target_os = "macos", target_os = "linux"))]
pub use sockaddr::SockAddr;
#[cfg(any(target_os = "macos", target_os = "linux"))]
pub fn exe_cmd(cmd: &str) -> std::io::Result<Output> {
use std::io;
use std::process::Command;
println!("exe cmd: {}", cmd);
let out = Command::new("sh")
.arg("-c")
.arg(cmd)
.output()
.expect("sh exec error!");
if !out.status.success() {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("cmd={},out={:?}", cmd, out),
));
}
Ok(out)
}
@@ -1,46 +1,17 @@
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// Version 2, December 2004
//
// Copyleft (ↄ) meh. <[email protected]> | http://meh.schizofreni.co
//
// Everyone is permitted to copy and distribute verbatim or modified
// copies of this license document, and changing it is allowed as long
// as the name is changed.
//
// DO WHAT THE FUCK YOU WANT TO PUBLIC LICENSE
// TERMS AND CONDITIONS FOR COPYING, DISTRIBUTION AND MODIFICATION
//
// 0. You just DO WHAT THE FUCK YOU WANT TO.
use std::mem;
use std::net::Ipv4Addr;
use std::ptr;
#[cfg(any(target_os = "macos", target_os = "ios"))]
use libc::c_uchar;
#[cfg(any(target_os = "linux", target_os = "android"))]
use libc::c_ushort;
use libc::AF_INET as _AF_INET;
use libc::{in_addr, sockaddr, sockaddr_in};
use std::{io, mem, net::Ipv4Addr, ptr};
use crate::error::*;
use io::Result;
/// A wrapper for `sockaddr_in`.
#[derive(Copy, Clone)]
pub struct SockAddr(sockaddr_in);
#[cfg(any(target_os = "linux", target_os = "android"))]
const AF_INET: c_ushort = _AF_INET as c_ushort;
#[cfg(any(target_os = "macos", target_os = "ios"))]
const AF_INET: c_uchar = _AF_INET as c_uchar;
impl SockAddr {
/// Create a new `SockAddr` from a generic `sockaddr`.
pub fn new(value: &sockaddr) -> Result<Self> {
if value.sa_family != AF_INET {
return Err(Error::InvalidAddress);
if value.sa_family != libc::AF_INET as libc::sa_family_t {
return Err(io::Error::new(io::ErrorKind::Other, "invalid address"));
}
unsafe { Self::unchecked(value) }
@@ -64,7 +35,7 @@ impl From<Ipv4Addr> for SockAddr {
let octets = ip.octets();
let mut addr = unsafe { mem::zeroed::<sockaddr_in>() };
addr.sin_family = AF_INET;
addr.sin_family = libc::AF_INET as libc::sa_family_t;
addr.sin_port = 0;
addr.sin_addr = in_addr {
s_addr: u32::from_ne_bytes(octets),
+91
View File
@@ -0,0 +1,91 @@
use crate::device::IFace;
use crate::windows::{tap, tun};
use std::io;
use std::net::Ipv4Addr;
pub enum Device {
Tap(tap::Device),
Tun(tun::Device),
}
impl Device {
pub fn new(name: String, tap: bool) -> io::Result<Self> {
if tap {
Ok(Device::Tap(tap::Device::new(name)?))
} else {
Ok(Device::Tun(tun::Device::new(name)?))
}
}
}
impl IFace for Device {
fn version(&self) -> io::Result<String> {
match self {
Device::Tap(dev) => dev.version(),
Device::Tun(dev) => dev.version(),
}
}
fn name(&self) -> io::Result<String> {
match self {
Device::Tap(dev) => dev.name(),
Device::Tun(dev) => dev.name(),
}
}
fn shutdown(&self) -> io::Result<()> {
match self {
Device::Tap(dev) => dev.shutdown(),
Device::Tun(dev) => dev.shutdown(),
}
}
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
match self {
Device::Tap(dev) => dev.set_ip(address, mask),
Device::Tun(dev) => dev.set_ip(address, mask),
}
}
fn mtu(&self) -> io::Result<u32> {
match self {
Device::Tap(dev) => dev.mtu(),
Device::Tun(dev) => dev.mtu(),
}
}
fn set_mtu(&self, value: u32) -> io::Result<()> {
match self {
Device::Tap(dev) => dev.set_mtu(value),
Device::Tun(dev) => dev.set_mtu(value),
}
}
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, metric: u16) -> io::Result<()> {
match self {
Device::Tap(dev) => dev.add_route(dest, netmask, metric),
Device::Tun(dev) => dev.add_route(dest, netmask, metric),
}
}
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
match self {
Device::Tap(dev) => dev.delete_route(dest, netmask),
Device::Tun(dev) => dev.delete_route(dest, netmask),
}
}
fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Device::Tap(dev) => dev.read(buf),
Device::Tun(dev) => dev.read(buf),
}
}
fn write(&self, buf: &[u8]) -> io::Result<usize> {
match self {
Device::Tap(dev) => dev.write(buf),
Device::Tun(dev) => dev.write(buf),
}
}
}
@@ -137,7 +137,7 @@ pub fn read_file(handle: HANDLE, buffer: &mut [u8]) -> io::Result<DWORD> {
&mut ip_overlapped,
) {
let e = io::Error::last_os_error();
if e.raw_os_error().unwrap_or(0) == 997 {
if e.raw_os_error().unwrap_or(0) == ERROR_IO_PENDING as _ {
if 0 == GetOverlappedResult(handle, &mut ip_overlapped, &mut ret, 1) {
return Err(e);
}
@@ -166,7 +166,7 @@ pub fn write_file(handle: HANDLE, buffer: &[u8]) -> io::Result<DWORD> {
&mut ip_overlapped,
) {
let e = io::Error::last_os_error();
if e.raw_os_error().unwrap_or(0) == 997 {
if e.raw_os_error().unwrap_or(0) == ERROR_IO_PENDING as _ {
if 0 == GetOverlappedResult(handle, &mut ip_overlapped, &mut ret, 1) {
return Err(e);
}
+43
View File
@@ -0,0 +1,43 @@
use std::io;
use std::os::windows::process::CommandExt;
use winapi::shared::minwindef::DWORD;
use winapi::um::winbase::CREATE_NO_WINDOW;
mod device;
mod ffi;
mod netsh;
mod route;
mod tap;
mod tun;
pub use device::Device;
/// Encode a string as a utf16 buffer
pub fn encode_utf16(string: &str) -> Vec<u16> {
use std::iter::once;
string.encode_utf16().chain(once(0)).collect()
}
pub fn decode_utf16(string: &[u16]) -> String {
let end = string.iter().position(|b| *b == 0).unwrap_or(string.len());
String::from_utf16_lossy(&string[..end])
}
pub const fn ctl_code(device_type: DWORD, function: DWORD, method: DWORD, access: DWORD) -> DWORD {
(device_type << 16) | (access << 14) | (function << 2) | method
}
pub fn exe_cmd(cmd: &str) -> io::Result<()> {
println!("exe cmd: {}", cmd);
let out = std::process::Command::new("cmd")
.creation_flags(CREATE_NO_WINDOW)
.arg("/C")
.arg(&cmd)
.output()?;
if !out.status.success() {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("cmd={},out={:?}", cmd, String::from_utf8(out.stderr)),
));
}
Ok(())
}
+47
View File
@@ -0,0 +1,47 @@
use crate::windows::exe_cmd;
use std::net::Ipv4Addr;
use std::{io, process};
/// 设置网卡名称
pub fn set_interface_name(old_name: &str, new_name: &str) -> io::Result<()> {
let cmd = format!(
" netsh interface set interface name={:?} newname={:?}",
old_name, new_name
);
exe_cmd(&cmd)
}
/// 删除缓存
pub fn delete_cache() -> io::Result<()> {
//清除缓存
let cmd = "netsh interface ip delete destinationcache";
exe_cmd(cmd)
}
/// 设置网卡ip
pub fn set_interface_ip(index: u32, address: &Ipv4Addr, netmask: &Ipv4Addr) -> io::Result<()> {
let cmd = format!(
"netsh interface ip set address {} static {:?} {:?} ",
index, address, netmask,
);
exe_cmd(&cmd)
}
pub fn set_interface_mtu(index: u32, mtu: u32) -> io::Result<()> {
let cmd = format!(
"netsh interface ipv4 set subinterface {} mtu={} store=persistent",
index, mtu
);
exe_cmd(&cmd)
}
pub fn set_interface_metric(index: u32, metric: u16) -> io::Result<()> {
let cmd = format!(
"netsh interface ip set interface {} metric={}",
index, metric
);
exe_cmd(&cmd)
}
/// 禁用ipv6
pub fn disabled_ipv6(index: u32) -> io::Result<()> {
let cmd = format!("netsh interface ipv6 set interface {} disabled", index);
exe_cmd(&cmd)
}
+33
View File
@@ -0,0 +1,33 @@
use std::io;
use std::net::Ipv4Addr;
use crate::windows::exe_cmd;
/// 添加路由
pub fn add_route(
index: u32,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()> {
let cmd = format!(
"route add {:?} mask {:?} {:?} metric {} if {}",
dest, netmask, gateway, metric, index
);
exe_cmd(&cmd)
}
/// 删除路由
pub fn delete_route(
index: u32,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
) -> io::Result<()> {
let cmd = format!(
"route delete {:?} mask {:?} {:?} if {}",
dest, netmask, gateway, index
);
exe_cmd(&cmd)
}
+188
View File
@@ -0,0 +1,188 @@
use std::io;
use std::net::Ipv4Addr;
use winapi::shared::ifdef::NET_LUID;
use winapi::shared::minwindef::DWORD;
use winapi::um::fileapi::OPEN_EXISTING;
use winapi::um::winbase::FILE_FLAG_OVERLAPPED;
use winapi::um::winioctl::{FILE_ANY_ACCESS, FILE_DEVICE_UNKNOWN, METHOD_BUFFERED};
use winapi::um::winnt::{
FILE_ATTRIBUTE_SYSTEM, FILE_SHARE_READ, FILE_SHARE_WRITE, GENERIC_READ, GENERIC_WRITE, HANDLE,
};
use crate::device::IFace;
use crate::packet;
use crate::packet::ethernet::protocol::Protocol;
use crate::packet::{arp, ethernet};
use crate::windows::{ctl_code, decode_utf16, encode_utf16, ffi, netsh, route};
/* Present in 8.1 */
const TAP_WIN_IOCTL_GET_MAC: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 1, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_GET_VERSION: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 2, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_GET_MTU: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 3, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_GET_INFO: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 4, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_CONFIG_POINT_TO_POINT: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 5, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_SET_MEDIA_STATUS: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 6, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_CONFIG_DHCP_MASQ: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 7, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_GET_LOG_LINE: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 8, METHOD_BUFFERED, FILE_ANY_ACCESS);
const TAP_WIN_IOCTL_CONFIG_DHCP_SET_OPT: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 9, METHOD_BUFFERED, FILE_ANY_ACCESS);
/* Added in 8.2 */
/* obsoletes TAP_WIN_IOCTL_CONFIG_POINT_TO_POINT */
const TAP_WIN_IOCTL_CONFIG_TUN: DWORD =
ctl_code(FILE_DEVICE_UNKNOWN, 10, METHOD_BUFFERED, FILE_ANY_ACCESS);
pub struct Device {
handle: HANDLE,
index: u32,
luid: NET_LUID,
mac: [u8; 6],
}
unsafe impl Send for Device {}
unsafe impl Sync for Device {}
impl Device {
/// 打开设备,设置为TUN模式,激活网卡
pub fn new(name: String) -> io::Result<Self> {
let luid = ffi::alias_to_luid(&encode_utf16(&name)).map_err(|e| {
io::Error::new(e.kind(), format!("alias_to_luid name={},err={:?}", name, e))
})?;
let guid = ffi::luid_to_guid(&luid)
.and_then(|guid| ffi::string_from_guid(&guid))
.map_err(|e| {
io::Error::new(e.kind(), format!("luid_to_guid name={},err={:?}", name, e))
})?;
let path = format!(r"\\.\Global\{}.tap", decode_utf16(&guid));
let handle = ffi::create_file(
&encode_utf16(&path),
GENERIC_READ | GENERIC_WRITE,
FILE_SHARE_READ | FILE_SHARE_WRITE,
OPEN_EXISTING,
FILE_ATTRIBUTE_SYSTEM | FILE_FLAG_OVERLAPPED,
)
.map_err(|e| io::Error::new(e.kind(), format!("tap name={},err={:?}", name, e)))?;
// ep保存tun网卡的IP地址和掩码
// let mut ep = [0;3];
// ep[0] = Ipv4Addr::new(10,26,0,11).into();
// ep[2] = Ipv4Addr::new(255,255,255,0).into();;
// ep[1] = ep[0] & ep[2];
// //tun模式收不到ipv4包,原因未知 https://github.com/OpenVPN/tap-windows6/issues/111
// ffi::device_io_control(handle, TAP_WIN_IOCTL_CONFIG_TUN, &ep, &mut ()).map_err(
// |e| {
// io::Error::new(
// e.kind(),
// format!("TAP_WIN_IOCTL_CONFIG_TUN name={},err={:?}", name_str, e),
// )
// },
// )?;
let mut mac = [0u8; 6];
ffi::device_io_control(handle, TAP_WIN_IOCTL_GET_MAC, &(), &mut mac)
.map_err(|e| {
io::Error::new(
e.kind(),
format!("TAP_WIN_IOCTL_CONFIG_TUN name={},err={:?}", name, e),
)
})
.map_err(|e| io::Error::new(e.kind(), format!("TAP_WIN_IOCTL_GET_MAC,err={:?}", e)))?;
let index = ffi::luid_to_index(&luid).map(|index| index as u32)?;
// 设置网卡跃点
if let Err(e) = netsh::set_interface_metric(index, 0) {
log::warn!("{:?}",e);
}
let device = Self {
handle,
index,
luid,
mac,
};
device.enabled(true)?;
Ok(device)
}
fn write_tap(&self, buf: &[u8]) -> io::Result<usize> {
ffi::write_file(self.handle, buf).map(|res| res as _)
}
fn enabled(&self, value: bool) -> io::Result<()> {
let status: u32 = if value { 1 } else { 0 };
ffi::device_io_control(
self.handle,
TAP_WIN_IOCTL_SET_MEDIA_STATUS,
&status,
&mut (),
)
}
}
const MAC: [u8; 6] = [0xf, 0xf, 0xf, 0xf, 0xe, 0x9];
impl IFace for Device {
fn version(&self) -> io::Result<String> {
let mut version = [0u32; 3];
ffi::device_io_control(self.handle, TAP_WIN_IOCTL_GET_VERSION, &(), &mut version)?;
Ok(format!("{}.{}.{}", version[0], version[1], version[2]))
}
fn name(&self) -> io::Result<String> {
ffi::luid_to_alias(&self.luid).map(|name| decode_utf16(&name))
}
fn shutdown(&self) -> io::Result<()> {
self.enabled(false)
}
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
netsh::set_interface_ip(self.index, &address, &mask)
}
fn mtu(&self) -> io::Result<u32> {
let mut mtu = 0;
ffi::device_io_control(self.handle, TAP_WIN_IOCTL_GET_MTU, &(), &mut mtu).map(|_| mtu)
}
fn set_mtu(&self, value: u32) -> io::Result<()> {
netsh::set_interface_mtu(self.index, value)
}
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, metric: u16) -> io::Result<()> {
route::add_route(self.index, dest, netmask, Ipv4Addr::UNSPECIFIED, metric)?;
netsh::delete_cache()
}
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
route::delete_route(self.index, dest, netmask, Ipv4Addr::UNSPECIFIED)?;
netsh::delete_cache()
}
fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
packet::read_tap(
buf,
|eth_buf| ffi::read_file(self.handle, eth_buf).map(|res| res as usize),
|eth_buf| ffi::write_file(self.handle, eth_buf).map(|res| res as _),
)
}
fn write(&self, buf: &[u8]) -> io::Result<usize> {
// 封装二层数据
packet::write_tap(
buf,
|eth_buf| ffi::write_file(self.handle, eth_buf).map(|res| res as _),
&self.mac,
)
}
}
impl Drop for Device {
fn drop(&mut self) {
if let Err(e) = ffi::close_handle(self.handle) {
log::warn!("close_handle={:?}", e)
}
}
}
@@ -1,12 +1,17 @@
use libloading::{Error, Library};
use std::io;
use std::net::Ipv4Addr;
use winapi::um::{handleapi, synchapi, winbase, winnt};
use crate::{decode_utf16, encode_utf16, ffi, netsh, route, IFace};
use rand::Rng;
mod log;
pub mod packet;
use winapi::um::winbase;
use winapi::um::{synchapi, winnt};
use crate::device::IFace;
use crate::windows::decode_utf16;
use crate::windows::{encode_utf16, ffi, netsh, route};
mod packet;
mod wintun_log;
mod wintun_raw;
/// The maximum size of wintun's internal ring buffer (in bytes)
@@ -18,7 +23,7 @@ pub const MIN_RING_CAPACITY: u32 = 0x2_0000;
/// Maximum pool name length including zero terminator
pub const MAX_POOL: usize = 256;
pub struct TunDevice {
pub struct Device {
pub(crate) luid: u64,
pub(crate) index: u32,
/// The session handle given to us by WintunStartSession
@@ -39,109 +44,103 @@ pub struct TunDevice {
pub(crate) adapter: wintun_raw::WINTUN_ADAPTER_HANDLE,
}
unsafe impl Send for TunDevice {}
unsafe impl Send for Device {}
unsafe impl Sync for TunDevice {}
unsafe impl Sync for Device {}
impl TunDevice {
pub unsafe fn create<L>(library: L, pool: &str, name: &str) -> io::Result<Self>
where
L: Into<libloading::Library>,
{
let win_tun = match wintun_raw::wintun::from_library(library) {
Ok(win_tun) => win_tun,
Err(e) => {
impl Device {
pub fn new(name: String) -> io::Result<Self> {
unsafe {
let library = match Library::new("wintun.dll") {
Ok(library) => library,
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("wintun.dll not found {:?}", e),
));
}
};
let win_tun = match wintun_raw::wintun::from_library(library) {
Ok(win_tun) => win_tun,
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("library error {:?} ", e),
));
}
};
let name_utf16 = encode_utf16(&name);
if name_utf16.len() > MAX_POOL {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("library error {:?} ", e),
format!("too long {}:{:?}", MAX_POOL, name),
));
}
};
let pool_utf16 = encode_utf16(pool);
if pool_utf16.len() > MAX_POOL {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("长度大于{}:{:?}", MAX_POOL, pool),
));
}
let name_utf16 = encode_utf16(name);
if name_utf16.len() > MAX_POOL {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("长度大于{}:{:?}", MAX_POOL, pool),
));
}
let mut guid_bytes: [u8; 16] = [0u8; 16];
rand::thread_rng().fill(&mut guid_bytes);
let guid = u128::from_ne_bytes(guid_bytes);
//SAFETY: guid is a unique integer so transmuting either all zeroes or the user's preferred
//guid to the winapi guid type is safe and will allow the windows kernel to see our GUID
wintun_log::set_default_logger_if_unset(&win_tun);
let _ = Self::delete_for_name(&win_tun, &name_utf16);
let mut guid_bytes: [u8; 16] = [0u8; 16];
rand::thread_rng().fill(&mut guid_bytes);
let guid = u128::from_ne_bytes(guid_bytes);
//SAFETY: guid is a unique integer so transmuting either all zeroes or the user's preferred
//guid to the winapi guid type is safe and will allow the windows kernel to see our GUID
let guid_struct: wintun_raw::GUID = unsafe { std::mem::transmute(guid) };
let guid_ptr = &guid_struct as *const wintun_raw::GUID;
let guid_struct: wintun_raw::GUID = unsafe { std::mem::transmute(guid) };
let guid_ptr = &guid_struct as *const wintun_raw::GUID;
log::set_default_logger_if_unset(&win_tun);
//SAFETY: the function is loaded from the wintun dll properly, we are providing valid
//pointers, and all the strings are correct null terminated UTF-16. This safety rationale
//applies for all Wintun* functions below
let adapter =
win_tun.WintunCreateAdapter(pool_utf16.as_ptr(), name_utf16.as_ptr(), guid_ptr);
if adapter.is_null() {
return Err(io::Error::new(
io::ErrorKind::Other,
"Failed to crate adapter",
));
}
Self::init(win_tun, adapter)
}
pub unsafe fn init(
win_tun: wintun_raw::wintun,
adapter: wintun_raw::WINTUN_ADAPTER_HANDLE,
) -> io::Result<Self> {
// 开启session
let session = win_tun.WintunStartSession(adapter, 128 * 1024);
if session.is_null() {
return Err(io::Error::new(
io::ErrorKind::Other,
"WintunStartSession failed",
));
}
//SAFETY: We follow the contract required by CreateEventA. See MSDN
//(the pointers are allowed to be null, and 0 is okay for the others)
let shutdown_event =
synchapi::CreateEventA(std::ptr::null_mut(), 0, 0, std::ptr::null_mut());
let read_event = win_tun.WintunGetReadWaitEvent(session) as winnt::HANDLE;
let mut luid: wintun_raw::NET_LUID = std::mem::zeroed();
win_tun.WintunGetAdapterLUID(adapter, &mut luid as *mut wintun_raw::NET_LUID);
let index = ffi::luid_to_index(&std::mem::transmute(luid)).map(|index| index as u32)?;
Ok(TunDevice {
luid: std::mem::transmute(luid),
index,
session,
win_tun,
read_event,
shutdown_event,
adapter,
})
}
pub unsafe fn delete_for_name<L>(library: L, name: &str) -> io::Result<()>
where
L: Into<libloading::Library>,
{
let win_tun = match wintun_raw::wintun::from_library(library) {
Ok(win_tun) => win_tun,
Err(e) => {
//SAFETY: the function is loaded from the wintun dll properly, we are providing valid
//pointers, and all the strings are correct null terminated UTF-16. This safety rationale
//applies for all Wintun* functions below
let adapter =
win_tun.WintunCreateAdapter(name_utf16.as_ptr(), name_utf16.as_ptr(), guid_ptr);
if adapter.is_null() {
log::error!("adapter.is_null {:?}", io::Error::last_os_error());
return Err(io::Error::new(
io::ErrorKind::Other,
format!("library error {:?} ", e),
"Failed to crate adapter",
));
}
};
log::set_default_logger_if_unset(&win_tun);
let name_utf16 = encode_utf16(name);
// 开启session
let session = win_tun.WintunStartSession(adapter, MAX_RING_CAPACITY);
if session.is_null() {
log::error!("session.is_null {:?}", io::Error::last_os_error());
return Err(io::Error::new(
io::ErrorKind::Other,
"WintunStartSession failed",
));
}
//SAFETY: We follow the contract required by CreateEventA. See MSDN
//(the pointers are allowed to be null, and 0 is okay for the others)
let shutdown_event =
synchapi::CreateEventA(std::ptr::null_mut(), 0, 0, std::ptr::null_mut());
let read_event = win_tun.WintunGetReadWaitEvent(session) as winnt::HANDLE;
let mut luid: wintun_raw::NET_LUID = std::mem::zeroed();
win_tun.WintunGetAdapterLUID(adapter, &mut luid as *mut wintun_raw::NET_LUID);
let index = ffi::luid_to_index(&std::mem::transmute(luid)).map(|index| index as u32)?;
// 设置网卡跃点
if let Err(e) = netsh::set_interface_metric(index, 0) {
log::warn!("{:?}",e);
}
Ok(Self {
luid: std::mem::transmute(luid),
index,
session,
win_tun,
read_event,
shutdown_event,
adapter,
})
}
}
pub unsafe fn delete_for_name(
win_tun: &wintun_raw::wintun,
name_utf16: &Vec<u16>,
) -> io::Result<()> {
let adapter = win_tun.WintunOpenAdapter(name_utf16.as_ptr());
if adapter.is_null() {
log::error!(
"delete_for_name adapter.is_null {:?}",
io::Error::last_os_error()
);
return Err(io::Error::new(
io::ErrorKind::Other,
"Failed to open adapter",
@@ -151,11 +150,10 @@ impl TunDevice {
win_tun.WintunDeleteDriver();
Ok(())
}
pub fn delete(self) -> io::Result<()> {
drop(self);
Ok(())
}
pub fn version(&self) -> io::Result<Version> {
}
impl IFace for Device {
fn version(&self) -> io::Result<String> {
let version = unsafe { self.win_tun.WintunGetRunningDriverVersion() };
if version == 0 {
return Err(io::Error::new(
@@ -163,78 +161,66 @@ impl TunDevice {
"WintunGetRunningDriverVersion",
));
} else {
Ok(Version {
major: ((version >> 16) & 0xFF) as u16,
minor: (version & 0xFF) as u16,
})
Ok(format!("{}.{}", (version >> 16) & 0xFFFF, version & 0xFFFF))
}
}
}
#[derive(Copy, Clone, PartialEq, Eq, Debug)]
pub struct Version {
pub major: u16,
pub minor: u16,
}
// impl TunDevice {
// fn get_adapter_luid(&self) -> u64 {
// let mut luid: wintun_raw::NET_LUID = unsafe { std::mem::zeroed() };
// unsafe { self.win_tun.WintunGetAdapterLUID(self.adapter, &mut luid as *mut wintun_raw::NET_LUID) };
// unsafe { std::mem::transmute(luid) }
// }
// }
impl IFace for TunDevice {
fn shutdown(&self) -> io::Result<()> {
let _ = unsafe { synchapi::SetEvent(self.shutdown_event) };
let _ = unsafe { handleapi::CloseHandle(self.shutdown_event) };
Ok(())
}
fn get_index(&self) -> io::Result<u32> {
Ok(self.index)
}
fn get_name(&self) -> io::Result<String> {
fn name(&self) -> io::Result<String> {
let luid = self.luid;
ffi::luid_to_alias(&unsafe { std::mem::transmute(luid) }).map(|name| decode_utf16(&name))
}
fn set_name(&self, new_name: &str) -> io::Result<()> {
let name = self.get_name()?;
netsh::set_interface_name(&name, new_name)
fn shutdown(&self) -> io::Result<()> {
unsafe {
if 0 == synchapi::SetEvent(self.shutdown_event) {
Ok(())
} else {
Err(io::Error::last_os_error())
}
}
}
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
netsh::set_interface_ip(self.get_index()?, &address, &mask)
netsh::set_interface_ip(self.index, &address, &mask)
}
fn add_route(
&self,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()> {
route::add_route(self.get_index()?, dest, netmask, gateway, metric)
fn mtu(&self) -> io::Result<u32> {
Err(io::Error::from(io::ErrorKind::Unsupported))
}
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> {
route::delete_route(self.get_index()?, dest, netmask, gateway)
fn set_mtu(&self, value: u32) -> io::Result<()> {
netsh::set_interface_mtu(self.index, value)
}
fn set_mtu(&self, mtu: u16) -> io::Result<()> {
netsh::set_interface_mtu(self.get_index()?, mtu)
fn add_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, metric: u16) -> io::Result<()> {
route::add_route(self.index, dest, netmask, Ipv4Addr::UNSPECIFIED, metric)?;
netsh::delete_cache()
}
fn set_metric(&self, metric: u16) -> io::Result<()> {
let index = self.get_index()?;
netsh::set_interface_metric(index, metric)
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> {
route::delete_route(self.index, dest, netmask, Ipv4Addr::UNSPECIFIED)?;
netsh::delete_cache()
}
fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
let packet = self.receive_blocking()?;
let packet = packet.bytes();
let len = packet.len();
if len > buf.len() {
return Err(io::Error::new(io::ErrorKind::InvalidData, "data too long"));
}
buf[..len].copy_from_slice(packet);
Ok(len)
}
fn write(&self, buf: &[u8]) -> io::Result<usize> {
let mut packet = self.allocate_send_packet(buf.len() as u16)?;
packet.bytes_mut().copy_from_slice(buf);
self.send_packet(packet);
Ok(buf.len())
}
}
impl TunDevice {
impl Device {
pub fn try_receive(&self) -> io::Result<Option<packet::TunPacket>> {
let mut size = 0u32;
@@ -265,9 +251,9 @@ impl TunDevice {
}
pub fn receive_blocking(&self) -> io::Result<packet::TunPacket> {
loop {
//Try 5 times to receive without blocking so we don't have to issue a syscall to wait
//Try 16 times to receive without blocking so we don't have to issue a syscall to wait
//for the event if packets are being received at a rapid rate
for _ in 0..5 {
for _i in 0..20 {
match self.try_receive()? {
None => {
continue;
@@ -291,7 +277,7 @@ impl TunDevice {
};
match result {
winbase::WAIT_FAILED => {
return Err(io::Error::new(io::ErrorKind::Other, "WAIT_FAILED"))
return Err(io::Error::new(io::ErrorKind::Other, "WAIT_FAILED"));
}
_ => {
if result == winbase::WAIT_OBJECT_0 {
@@ -308,9 +294,6 @@ impl TunDevice {
}
}
}
}
impl TunDevice {
pub fn allocate_send_packet(&self, size: u16) -> io::Result<packet::TunPacket> {
let bytes_ptr = unsafe {
self.win_tun
@@ -344,13 +327,17 @@ impl TunDevice {
}
}
impl Drop for TunDevice {
impl Drop for Device {
fn drop(&mut self) {
//Close adapter on drop
//This is why we need an Arc of wintun
unsafe {
if let Err(e) = ffi::close_handle(self.shutdown_event) {
log::warn!("close shutdown_event={:?}", e)
}
self.win_tun.WintunEndSession(self.session);
self.win_tun.WintunCloseAdapter(self.adapter);
self.win_tun.WintunDeleteDriver()
};
if 0 != self.win_tun.WintunDeleteDriver() {
log::warn!("WintunDeleteDriver failed")
}
}
}
}
@@ -1,4 +1,4 @@
use crate::TunDevice;
use crate::windows::tun::Device;
pub(crate) enum Kind {
SendPacketPending,
@@ -16,7 +16,7 @@ pub struct TunPacket<'a> {
//Share ownership of session to prevent the session from being dropped before packets that
//belong to it
pub(crate) tun_device: Option<&'a TunDevice>,
pub(crate) tun_device: Option<&'a Device>,
}
impl<'a> TunPacket<'a> {
@@ -1,6 +1,6 @@
use log::*;
use crate::tun::wintun_raw;
use crate::windows::tun::wintun_raw;
use std::sync::atomic::{AtomicBool, Ordering};
use widestring::U16CStr;
-33
View File
@@ -1,33 +0,0 @@
[package]
name = "win-tun-tap"
version = "0.1.0"
edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
[dependencies]
log = "0.4.17"
winreg = "0.7"
scopeguard = "1.1"
libloading = "0.7"
widestring = "0.4"
once_cell = "1.8"
itertools = "0.10.1"
rand = "0.8.5"
[dependencies.winapi]
version = "0.3"
features = [
"errhandlingapi",
"combaseapi",
"ioapiset",
"winioctl",
"setupapi",
"synchapi",
"netioapi",
"fileapi",
"winbase",
"winerror",
"ipexport",
"iphlpapi",
"handleapi"
]
-49
View File
@@ -1,49 +0,0 @@
#![cfg(windows)]
mod ffi;
mod netsh;
mod route;
mod tap;
mod tun;
use std::io;
use std::net::Ipv4Addr;
pub use tap::TapDevice;
pub use tun::*;
/// Encode a string as a utf16 buffer
fn encode_utf16(string: &str) -> Vec<u16> {
use std::iter::once;
string.encode_utf16().chain(once(0)).collect()
}
/// Decode a string from a utf16 buffer
fn decode_utf16(string: &[u16]) -> String {
let end = string.iter().position(|b| *b == 0).unwrap_or(string.len());
String::from_utf16_lossy(&string[..end])
}
pub trait IFace {
fn shutdown(&self) -> io::Result<()>;
/// 获取接口索引
fn get_index(&self) -> io::Result<u32>;
/// 获取名称
fn get_name(&self) -> io::Result<String>;
/// 设置名称
fn set_name(&self, new_name: &str) -> io::Result<()>;
/// 设置ip
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()>;
/// 设置路由
fn add_route(
&self,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()>;
/// 删除路由
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()>;
/// 设置最大传输单元
fn set_mtu(&self, mtu: u16) -> io::Result<()>;
/// 设置跃点
fn set_metric(&self, metric: u16) -> io::Result<()>;
}
-80
View File
@@ -1,80 +0,0 @@
use std::io;
use std::net::Ipv4Addr;
use std::os::windows::process::CommandExt;
/// 设置网卡名称
pub fn set_interface_name(old_name: &str, new_name: &str) -> io::Result<()> {
let cmd = format!(
" netsh interface set interface name={:?} newname={:?}",
old_name, new_name
);
let out = std::process::Command::new("cmd")
.creation_flags(0x08000000) //winapi-0.3.9/src/um/winbase.rs:283
.arg("/C")
.arg(&cmd)
.output()?;
if !out.status.success() {
log::warn!("修改网卡名称失败:cmd={:?},out={:?}", cmd, out);
return Err(io::Error::new(io::ErrorKind::Other, "修改网卡名称失败"));
}
Ok(())
}
/// 设置网卡ip
pub fn set_interface_ip(index: u32, address: &Ipv4Addr, netmask: &Ipv4Addr) -> io::Result<()> {
let set_address = format!(
"netsh interface ip set address {} static {:?} {:?} ",
index, address, netmask,
);
let out = std::process::Command::new("cmd")
.creation_flags(0x08000000)
.arg("/C")
.arg(&set_address)
.output()?;
if !out.status.success() {
log::error!("cmd={:?},out={:?}", set_address, out);
return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置网络地址失败: {:?}", out),
));
}
Ok(())
}
pub fn set_interface_mtu(index: u32, mtu: u16) -> io::Result<()> {
let set_mtu = format!(
"netsh interface ipv4 set subinterface {} mtu={} store=persistent",
index, mtu
);
let out = std::process::Command::new("cmd")
.creation_flags(0x08000000)
.arg("/C")
.arg(&set_mtu)
.output()?;
if !out.status.success() {
log::error!("cmd={:?},out={:?}", set_mtu, out);
return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置mtu失败: {:?}", out),
));
}
Ok(())
}
pub fn set_interface_metric(index: u32, metric: u16) -> io::Result<()> {
let set_metric = format!(
"netsh interface ip set interface {} metric={}",
index, metric
);
let out = std::process::Command::new("cmd")
.creation_flags(0x08000000)
.arg("/C")
.arg(&set_metric)
.output()?;
if !out.status.success() {
log::error!("cmd={:?},out={:?}", set_metric, out);
return Err(io::Error::new(
io::ErrorKind::Other,
format!("设置metric失败: {:?}", out),
));
}
Ok(())
}
-65
View File
@@ -1,65 +0,0 @@
use std::io;
use std::net::Ipv4Addr;
use std::os::windows::process::CommandExt;
/// 添加路由
pub fn add_route(
index: u32,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()> {
let set_route = format!(
"route add {:?} mask {:?} {:?} metric {} if {}",
dest, netmask, gateway, metric, index
);
// 执行添加路由命令
let out = std::process::Command::new("cmd")
.creation_flags(0x08000000)
.arg("/C")
.arg(&set_route)
.output()
.unwrap();
if !out.status.success() {
log::error!("cmd={:?},out={:?}", set_route, out);
return Err(io::Error::new(
io::ErrorKind::Other,
format!("添加路由失败: {:?}", out),
));
}
Ok(())
}
/// 删除路由
pub fn delete_route(
index: u32,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
) -> io::Result<()> {
if index == 0 {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("网络接口索引错误: {:?}", index),
));
}
let delete_route = format!(
"route delete {:?} mask {:?} {:?} if {}",
dest, netmask, gateway, index
);
// 删除路由
let out = std::process::Command::new("cmd")
.creation_flags(0x08000000)
.arg("/C")
.arg(delete_route)
.output()
.unwrap();
if !out.status.success() {
return Err(io::Error::new(
io::ErrorKind::Other,
format!("删除路由失败: {:?}", out),
));
}
Ok(())
}
-297
View File
@@ -1,297 +0,0 @@
use winapi::shared::ifdef::NET_LUID;
use winapi::shared::minwindef::*;
use winapi::um::fileapi::*;
use winapi::um::setupapi::*;
use winapi::um::winnt::*;
use scopeguard::{guard, ScopeGuard};
use winreg::RegKey;
use std::io;
use winapi::um::winbase::FILE_FLAG_OVERLAPPED;
use crate::{decode_utf16, encode_utf16, ffi};
/// tap-windows hardware ID
const HARDWARE_ID: &str = "tap0901";
winapi::DEFINE_GUID! {
GUID_NETWORK_ADAPTER,
0x4d36e972, 0xe325, 0x11ce,
0xbf, 0xc1, 0x08, 0x00, 0x2b, 0xe1, 0x03, 0x18
}
/// Create a new interface and returns its NET_LUID
pub fn create_interface() -> io::Result<NET_LUID> {
let devinfo = ffi::create_device_info_list(&GUID_NETWORK_ADAPTER)?;
let _guard = guard((), |_| {
let _ = ffi::destroy_device_info_list(devinfo);
});
let class_name = ffi::class_name_from_guid(&GUID_NETWORK_ADAPTER)?;
let devinfo_data = ffi::create_device_info(
devinfo,
&class_name,
&GUID_NETWORK_ADAPTER,
&encode_utf16(""),
DICD_GENERATE_ID,
)?;
ffi::set_selected_device(devinfo, &devinfo_data)?;
ffi::set_device_registry_property(
devinfo,
&devinfo_data,
SPDRP_HARDWAREID,
&encode_utf16(HARDWARE_ID),
)?;
ffi::build_driver_info_list(devinfo, &devinfo_data, SPDIT_COMPATDRIVER)?;
let _guard = guard((), |_| {
let _ = ffi::destroy_driver_info_list(devinfo, &devinfo_data, SPDIT_COMPATDRIVER);
});
let mut driver_version = 0;
let mut member_index = 0;
while let Some(drvinfo_data) =
ffi::enum_driver_info(devinfo, &devinfo_data, SPDIT_COMPATDRIVER, member_index)
{
member_index += 1;
let drvinfo_data = match drvinfo_data {
Ok(drvinfo_data) => drvinfo_data,
_ => continue,
};
if drvinfo_data.DriverVersion <= driver_version {
continue;
}
let drvinfo_detail =
match ffi::get_driver_info_detail(devinfo, &devinfo_data, &drvinfo_data) {
Ok(drvinfo_detail) => drvinfo_detail,
_ => continue,
};
let is_compatible = drvinfo_detail
.HardwareID
.split(|b| *b == 0)
.map(|id| decode_utf16(id))
.any(|id| id.eq_ignore_ascii_case(HARDWARE_ID));
if !is_compatible {
continue;
}
match ffi::set_selected_driver(devinfo, &devinfo_data, &drvinfo_data) {
Ok(_) => (),
_ => continue,
}
driver_version = drvinfo_data.DriverVersion;
}
if driver_version == 0 {
return Err(io::Error::new(io::ErrorKind::NotFound, "No driver found"));
}
let uninstaller = guard((), |_| {
let _ = ffi::call_class_installer(devinfo, &devinfo_data, DIF_REMOVE);
});
ffi::call_class_installer(devinfo, &devinfo_data, DIF_REGISTERDEVICE)?;
let _ = ffi::call_class_installer(devinfo, &devinfo_data, DIF_REGISTER_COINSTALLERS);
let _ = ffi::call_class_installer(devinfo, &devinfo_data, DIF_INSTALLINTERFACES);
ffi::call_class_installer(devinfo, &devinfo_data, DIF_INSTALLDEVICE)?;
let key = ffi::open_dev_reg_key(
devinfo,
&devinfo_data,
DICS_FLAG_GLOBAL,
0,
DIREG_DRV,
KEY_QUERY_VALUE | KEY_NOTIFY,
)?;
let key = RegKey::predef(key);
while let Err(_) = key.get_value::<DWORD, &str>("*IfType") {
ffi::notify_change_key_value(key.raw_handle(), TRUE, REG_NOTIFY_CHANGE_NAME, 2000)?;
}
while let Err(_) = key.get_value::<DWORD, &str>("NetLuidIndex") {
ffi::notify_change_key_value(key.raw_handle(), TRUE, REG_NOTIFY_CHANGE_NAME, 2000)?;
}
let if_type: DWORD = key.get_value("*IfType")?;
let luid_index: DWORD = key.get_value("NetLuidIndex")?;
// Defuse the uninstaller
ScopeGuard::into_inner(uninstaller);
let mut luid = NET_LUID { Value: 0 };
luid.set_IfType(if_type as _);
luid.set_NetLuidIndex(luid_index as _);
Ok(luid)
}
/// Check if the given interface exists and is a valid tap-windows device
pub fn check_interface(luid: &NET_LUID) -> io::Result<()> {
let devinfo = ffi::get_class_devs(&GUID_NETWORK_ADAPTER, DIGCF_PRESENT)?;
let _guard = guard((), |_| {
let _ = ffi::destroy_device_info_list(devinfo);
});
let mut member_index = 0;
while let Some(devinfo_data) = ffi::enum_device_info(devinfo, member_index) {
member_index += 1;
let devinfo_data = match devinfo_data {
Ok(devinfo_data) => devinfo_data,
Err(_) => continue,
};
let hardware_id =
match ffi::get_device_registry_property(devinfo, &devinfo_data, SPDRP_HARDWAREID) {
Ok(hardware_id) => hardware_id,
Err(_) => continue,
};
if !decode_utf16(&hardware_id).eq_ignore_ascii_case(HARDWARE_ID) {
continue;
}
let key = match ffi::open_dev_reg_key(
devinfo,
&devinfo_data,
DICS_FLAG_GLOBAL,
0,
DIREG_DRV,
KEY_QUERY_VALUE | KEY_NOTIFY,
) {
Ok(key) => RegKey::predef(key),
Err(_) => continue,
};
let if_type: DWORD = match key.get_value("*IfType") {
Ok(if_type) => if_type,
Err(_) => continue,
};
let luid_index: DWORD = match key.get_value("NetLuidIndex") {
Ok(luid_index) => luid_index,
Err(_) => continue,
};
let mut luid2 = NET_LUID { Value: 0 };
luid2.set_IfType(if_type as _);
luid2.set_NetLuidIndex(luid_index as _);
if luid.Value != luid2.Value {
continue;
}
// Found it!
return Ok(());
}
Err(io::Error::new(
io::ErrorKind::NotFound,
"TAP Device not found",
))
}
/// Deletes an existing interface
pub fn delete_interface(luid: &NET_LUID) -> io::Result<()> {
let devinfo = ffi::get_class_devs(&GUID_NETWORK_ADAPTER, DIGCF_PRESENT)?;
let _guard = guard((), |_| {
let _ = ffi::destroy_device_info_list(devinfo);
});
let mut member_index = 0;
while let Some(devinfo_data) = ffi::enum_device_info(devinfo, member_index) {
member_index += 1;
let devinfo_data = match devinfo_data {
Ok(devinfo_data) => devinfo_data,
Err(_) => continue,
};
let hardware_id =
match ffi::get_device_registry_property(devinfo, &devinfo_data, SPDRP_HARDWAREID) {
Ok(hardware_id) => hardware_id,
Err(_) => continue,
};
if !decode_utf16(&hardware_id).eq_ignore_ascii_case(HARDWARE_ID) {
continue;
}
let key = match ffi::open_dev_reg_key(
devinfo,
&devinfo_data,
DICS_FLAG_GLOBAL,
0,
DIREG_DRV,
KEY_QUERY_VALUE | KEY_NOTIFY,
) {
Ok(key) => RegKey::predef(key),
Err(_) => continue,
};
let if_type: DWORD = match key.get_value("*IfType") {
Ok(if_type) => if_type,
Err(_) => continue,
};
let luid_index: DWORD = match key.get_value("NetLuidIndex") {
Ok(luid_index) => luid_index,
Err(_) => continue,
};
let mut luid2 = NET_LUID { Value: 0 };
luid2.set_IfType(if_type as _);
luid2.set_NetLuidIndex(luid_index as _);
if luid.Value != luid2.Value {
continue;
}
// Found it!
return ffi::call_class_installer(devinfo, &devinfo_data, DIF_REMOVE);
}
Err(io::Error::new(
io::ErrorKind::NotFound,
"TAP Device not found",
))
}
/// Open an handle to an interface
pub fn open_interface(luid: &NET_LUID) -> io::Result<HANDLE> {
let guid = ffi::luid_to_guid(luid).and_then(|guid| ffi::string_from_guid(&guid))?;
let path = format!(r"\\.\Global\{}.tap", &decode_utf16(&guid));
ffi::create_file(
&encode_utf16(&path),
GENERIC_READ | GENERIC_WRITE,
FILE_SHARE_READ | FILE_SHARE_WRITE,
OPEN_EXISTING,
FILE_ATTRIBUTE_SYSTEM | FILE_FLAG_OVERLAPPED, //FILE_ATTRIBUTE_SYSTEM,
)
}
-190
View File
@@ -1,190 +0,0 @@
use std::net::Ipv4Addr;
use std::{io, time};
use winapi::shared::ifdef::NET_LUID;
use winapi::um::winioctl::*;
use winapi::um::winnt::HANDLE;
use crate::{decode_utf16, encode_utf16, ffi, netsh, route, IFace};
mod iface;
pub struct TapDevice {
index: u32,
luid: NET_LUID,
handle: HANDLE,
}
unsafe impl Send for TapDevice {}
unsafe impl Sync for TapDevice {}
impl TapDevice {
/// Retieve the mac of the interface
pub fn get_mac(&self) -> io::Result<[u8; 6]> {
let mut mac = [0; 6];
ffi::device_io_control(
self.handle,
CTL_CODE(FILE_DEVICE_UNKNOWN, 1, METHOD_BUFFERED, FILE_ANY_ACCESS),
&(),
&mut mac,
)
.map(|_| mac)
}
/// Retrieve the version of the driver
pub fn get_version(&self) -> io::Result<[u32; 3]> {
let mut version = [0; 3];
ffi::device_io_control(
self.handle,
CTL_CODE(FILE_DEVICE_UNKNOWN, 2, METHOD_BUFFERED, FILE_ANY_ACCESS),
&(),
&mut version,
)
.map(|_| version)
}
/// Retieve the mtu of the interface
pub fn get_mtu(&self) -> io::Result<u32> {
let mut mtu = 0;
ffi::device_io_control(
self.handle,
CTL_CODE(FILE_DEVICE_UNKNOWN, 3, METHOD_BUFFERED, FILE_ANY_ACCESS),
&(),
&mut mtu,
)
.map(|_| mtu)
}
/// Set the status of the interface, true for connected,
/// false for disconnected.
pub fn set_status(&self, status: bool) -> io::Result<()> {
let status: u32 = if status { 1 } else { 0 };
ffi::device_io_control(
self.handle,
CTL_CODE(FILE_DEVICE_UNKNOWN, 6, METHOD_BUFFERED, FILE_ANY_ACCESS),
&status,
&mut (),
)
}
}
impl TapDevice {
pub fn create() -> io::Result<Self> {
let luid = iface::create_interface()?;
// Even after retrieving the luid, we might need to wait
let start = time::Instant::now();
let handle = loop {
// If we surpassed 2 seconds just return
let now = time::Instant::now();
if now - start > time::Duration::from_secs(3) {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"Interface timed out",
));
}
match iface::open_interface(&luid) {
Err(_) => {
std::thread::yield_now();
continue;
}
Ok(handle) => break handle,
};
};
let index = ffi::luid_to_index(&luid).map(|index| index as u32)?;
Ok(Self {
index,
luid,
handle,
})
}
pub fn open(name: &str) -> io::Result<Self> {
let name = encode_utf16(name);
let luid = ffi::alias_to_luid(&name)?;
iface::check_interface(&luid)?;
let handle = iface::open_interface(&luid)?;
let index = ffi::luid_to_index(&luid).map(|index| index as u32)?;
Ok(Self {
index,
luid,
handle,
})
}
pub fn delete(self) -> io::Result<()> {
iface::delete_interface(&self.luid)
}
}
impl IFace for TapDevice {
fn shutdown(&self) -> io::Result<()> {
self.set_status(false)
}
fn get_index(&self) -> io::Result<u32> {
Ok(self.index)
}
fn get_name(&self) -> io::Result<String> {
ffi::luid_to_alias(&self.luid).map(|name| decode_utf16(&name))
}
fn set_name(&self, new_name: &str) -> io::Result<()> {
let name = self.get_name()?;
netsh::set_interface_name(&name, new_name)
}
fn set_ip(&self, address: Ipv4Addr, mask: Ipv4Addr) -> io::Result<()> {
let index = self.get_index()?;
netsh::set_interface_ip(index, &address, &mask)
}
fn add_route(
&self,
dest: Ipv4Addr,
netmask: Ipv4Addr,
gateway: Ipv4Addr,
metric: u16,
) -> io::Result<()> {
let index = self.get_index()?;
route::add_route(index, dest, netmask, gateway, metric)
}
fn delete_route(&self, dest: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr) -> io::Result<()> {
let index = self.get_index()?;
route::delete_route(index, dest, netmask, gateway)
}
fn set_mtu(&self, mtu: u16) -> io::Result<()> {
let index = self.get_index()?;
netsh::set_interface_mtu(index, mtu)
}
fn set_metric(&self, metric: u16) -> io::Result<()> {
let index = self.get_index()?;
netsh::set_interface_metric(index, metric)
}
}
impl TapDevice {
pub fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
ffi::read_file(self.handle, buf).map(|res| res as _)
}
pub fn write(&self, buf: &[u8]) -> io::Result<usize> {
ffi::write_file(self.handle, buf).map(|res| res as _)
}
}
impl Drop for TapDevice {
fn drop(&mut self) {
let _ = ffi::close_handle(self.handle);
let _ = iface::delete_interface(&self.luid);
}
}