mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 23:56:36 +08:00
Compare commits
15 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3ce320da5a | |||
| 7ab0db29ae | |||
| 006ea97200 | |||
| e569aedd3e | |||
| 35080aea2d | |||
| 85e588ffe9 | |||
| 03524f4a65 | |||
| 9e69e020ab | |||
| 14bbd3907d | |||
| ca8d8e92ba | |||
| 079474fa06 | |||
| 6e249a54f4 | |||
| fb798a4532 | |||
| a599f383f5 | |||
| 6bfa7f0166 |
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,407 @@
|
||||
# nftables 纯转发设计
|
||||
|
||||
**日期**: 2026-05-30
|
||||
**状态**: 待审核
|
||||
**作者**: Codex
|
||||
|
||||
## 概述
|
||||
|
||||
为 FLVX 增加一种不依赖 agent 的纯转发能力:节点可选择 `nftables` 转发模式,面板通过 SSH 在节点机器上下发和维护 nftables 规则。
|
||||
|
||||
第一阶段只支持端口级 DNAT/SNAT 纯转发。它不是 GOST 隧道能力的替代品,也不支持链路、限速、流量统计、连接数限制、Proxy Protocol、best exit 或 agent 诊断。目标是提供一个可靠、可回滚、可重建的轻量转发路径。
|
||||
|
||||
## 背景
|
||||
|
||||
当前 FLVX 的转发模型由三部分组成:
|
||||
|
||||
- `node` 表描述节点,现有本地节点通过 agent WebSocket 接收运行时命令。
|
||||
- `tunnel` 表描述入口、出口和链路类型,`type=1` 表示端口转发,`type=2` 表示隧道转发。
|
||||
- `forward` 表描述用户规则、入口端口和目标地址,运行时通过 GOST service 下发到入口节点。
|
||||
|
||||
nftables 模式的核心差异是没有 agent,因此不能复用现有 WebSocket command 通道,也不能依赖 agent 上报在线状态、流量和诊断结果。面板必须成为唯一控制面,通过 SSH 把数据库中的期望状态同步到远端 nftables。
|
||||
|
||||
## 用户决策
|
||||
|
||||
- 创建或编辑节点时选择转发模式。
|
||||
- 选择 nftables 转发后,不需要安装 agent。
|
||||
- nftables 转发不支持隧道、流量控制等能力,只支持纯转发。
|
||||
- 规则由面板端维护,并通过 SSH 下放到节点。
|
||||
|
||||
## 推荐方案
|
||||
|
||||
新增节点运行时模式:
|
||||
|
||||
| 模式 | 含义 |
|
||||
|------|------|
|
||||
| `agent` | 默认模式,保持现有 GOST agent 行为 |
|
||||
| `nftables` | 面板通过 SSH 管理 nftables 规则 |
|
||||
|
||||
业务层继续复用现有 `tunnel` 和 `forward` 概念,但对 nftables 模式加严格能力边界:
|
||||
|
||||
- nftables 节点只能创建端口转发隧道。
|
||||
- nftables 隧道不能配置出口节点或转发链。
|
||||
- 同一个隧道的入口节点必须全部是同一种运行时模式。
|
||||
- nftables 转发规则创建、更新、删除时,由后端同步 SSH 规则。
|
||||
- 面板提供节点级“测试 SSH”“重建规则”“清理 FLVX 规则”操作。
|
||||
|
||||
## 非目标
|
||||
|
||||
- 不支持 `tunnel.type=2` 隧道转发。
|
||||
- 不支持多跳链路、远程面板共享节点和 federation runtime。
|
||||
- 不支持 GOST service 能力:限速、每 IP 限速、最大连接数、Proxy Protocol、策略负载均衡。
|
||||
- 不支持 agent 流量统计、实时系统指标、节点升级、回退、agent 安装命令。
|
||||
- 不在第一阶段支持 HA 漂移、自动探活切换或复杂负载均衡。
|
||||
- 不改写用户机器上的非 FLVX nftables 规则。
|
||||
|
||||
## 数据模型
|
||||
|
||||
### node 表
|
||||
|
||||
新增字段:
|
||||
|
||||
| 字段 | 类型 | 默认 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `forward_mode` | string | `agent` | `agent` 或 `nftables` |
|
||||
|
||||
Go 模型使用 SQLite/PostgreSQL 兼容 tag:
|
||||
|
||||
```go
|
||||
ForwardMode string `gorm:"column:forward_mode;type:varchar(20);not null;default:'agent'"`
|
||||
```
|
||||
|
||||
### node_ssh_config 表
|
||||
|
||||
新增表保存 nftables 节点 SSH 配置。SSH 凭据不放进 `node` 主表,避免普通节点列表过度暴露敏感字段。
|
||||
|
||||
| 字段 | 说明 |
|
||||
|------|------|
|
||||
| `id` | 主键 |
|
||||
| `node_id` | 关联节点,唯一 |
|
||||
| `host` | SSH 主机,默认可使用 node.server_ip |
|
||||
| `port` | SSH 端口,默认 22 |
|
||||
| `username` | SSH 用户 |
|
||||
| `auth_type` | `password` 或 `private_key` |
|
||||
| `password` | 加密后密码,可为空 |
|
||||
| `private_key` | 加密后私钥,可为空 |
|
||||
| `passphrase` | 加密后私钥口令,可为空 |
|
||||
| `sudo_mode` | `none` / `sudo` |
|
||||
| `created_time` | 创建时间 |
|
||||
| `updated_time` | 更新时间 |
|
||||
|
||||
第一阶段可使用现有配置密钥派生或面板本地密钥做对称加密;如果项目尚无统一密钥管理,应至少避免在列表 API 返回完整凭据。
|
||||
|
||||
### nft_rule_binding 表
|
||||
|
||||
记录面板认为已经应用到节点的规则状态,用于更新、删除、重建和错误展示。
|
||||
|
||||
| 字段 | 说明 |
|
||||
|------|------|
|
||||
| `id` | 主键 |
|
||||
| `forward_id` | 转发规则 ID |
|
||||
| `node_id` | 下发节点 ID |
|
||||
| `in_port` | 入口端口 |
|
||||
| `protocols` | 第一阶段固定 `tcp,udp` |
|
||||
| `target_addr` | 目标地址 |
|
||||
| `bind_ip` | 可选监听 IP |
|
||||
| `rule_hash` | 当前期望规则 hash |
|
||||
| `status` | `pending` / `applied` / `error` |
|
||||
| `last_error` | 最近错误 |
|
||||
| `applied_time` | 最近成功应用时间 |
|
||||
| `created_time` | 创建时间 |
|
||||
| `updated_time` | 更新时间 |
|
||||
|
||||
绑定表不是最终事实来源。最终期望状态仍从 `forward`、`forward_port`、`tunnel` 和 `chain_tunnel` 推导,绑定表只记录应用结果。
|
||||
|
||||
## API 行为
|
||||
|
||||
### 节点创建和更新
|
||||
|
||||
`/node/create` 和 `/node/update` 新增入参:
|
||||
|
||||
```json
|
||||
{
|
||||
"forwardMode": "nftables",
|
||||
"sshConfig": {
|
||||
"host": "203.0.113.10",
|
||||
"port": 22,
|
||||
"username": "root",
|
||||
"authType": "private_key",
|
||||
"privateKey": "-----BEGIN OPENSSH PRIVATE KEY-----...",
|
||||
"passphrase": "",
|
||||
"sudoMode": "none"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- `forwardMode` 缺省时按 `agent`。
|
||||
- `agent` 节点保留现有字段和行为。
|
||||
- `nftables` 节点要求 SSH 配置完整。
|
||||
- 从 `agent` 切到 `nftables` 前,若该节点已有 agent 隧道链路或转发规则,应拒绝并提示先迁移或删除。
|
||||
- 从 `nftables` 切回 `agent` 前,若存在 nftables 规则,应拒绝并提示先清理或迁移。
|
||||
|
||||
### 隧道创建和更新
|
||||
|
||||
创建 nftables 隧道仍使用 `/tunnel/create`,但后端根据入口节点模式校验能力。
|
||||
|
||||
规则:
|
||||
|
||||
- 入口节点为 nftables 时,`type` 必须为 `1`。
|
||||
- 不允许提交 `outNodeId` 或 `chainNodes`。
|
||||
- 入口节点必须在线的现有校验不能直接套用到 nftables 节点;应改为 SSH 可用性校验或允许保存后手动测试。
|
||||
- 同一隧道入口节点不能混用 `agent` 和 `nftables`。
|
||||
- 更新隧道时不允许改变运行时模式;需要通过迁移规则到新隧道实现。
|
||||
|
||||
### 转发创建和更新
|
||||
|
||||
选择 nftables 隧道时,`/forward/create` 和 `/forward/update` 强制收窄字段:
|
||||
|
||||
- `speedId` 必须为空。
|
||||
- `ipSpeedId` 必须为空。
|
||||
- `maxConn` 和 `ipMaxConn` 必须为 0。
|
||||
- `proxyProtocol` 必须为 0。
|
||||
- 第一阶段 `remoteAddr` 只允许单目标 `host:port`。
|
||||
- `strategy` 固定为 `fifo` 或忽略。
|
||||
|
||||
创建流程:
|
||||
|
||||
1. 校验权限、隧道状态、端口占用和 nftables 能力边界。
|
||||
2. 在数据库创建 `forward` 和 `forward_port`。
|
||||
3. 通过 nftables runtime 对关联入口节点执行同步。
|
||||
4. 若同步失败,回滚数据库创建,返回 SSH/nftables 错误。
|
||||
|
||||
更新流程:
|
||||
|
||||
1. 保存旧 forward 和端口绑定。
|
||||
2. 更新数据库。
|
||||
3. 同步 nftables 规则。
|
||||
4. 若同步失败,回滚数据库状态并尝试恢复旧规则。
|
||||
|
||||
删除流程:
|
||||
|
||||
1. 先删除远端 nftables 规则。
|
||||
2. 成功后删除数据库。
|
||||
3. 如果远端删除失败,普通删除返回错误;强制删除可删除数据库并保留 binding 错误记录,提示用户稍后清理。
|
||||
|
||||
## 后端组件
|
||||
|
||||
新增 package:
|
||||
|
||||
```text
|
||||
go-backend/internal/runtime/nftables/
|
||||
```
|
||||
|
||||
建议拆分:
|
||||
|
||||
| 组件 | 职责 |
|
||||
|------|------|
|
||||
| `Manager` | 对 handler 暴露 Apply/Delete/Reconcile/Test 方法 |
|
||||
| `Planner` | 从数据库记录生成节点级期望规则 |
|
||||
| `Renderer` | 把期望规则渲染为 nftables 脚本 |
|
||||
| `SSHRunner` | 负责 SSH 连接、sudo 包装、命令执行和超时 |
|
||||
| `Parser` | 解析目标地址、协议和错误信息 |
|
||||
|
||||
handler 不直接执行 SSH,也不拼 nft 脚本;handler 只做业务校验并调用 runtime manager。
|
||||
|
||||
## nftables 规则设计
|
||||
|
||||
FLVX 只维护自己的 table,避免触碰用户已有规则:
|
||||
|
||||
```nft
|
||||
table inet flvx {
|
||||
chain prerouting {
|
||||
type nat hook prerouting priority dstnat; policy accept;
|
||||
}
|
||||
|
||||
chain postrouting {
|
||||
type nat hook postrouting priority srcnat; policy accept;
|
||||
}
|
||||
|
||||
chain forward {
|
||||
type filter hook forward priority filter; policy accept;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
每条 forward 生成 TCP 和 UDP 规则:
|
||||
|
||||
```nft
|
||||
tcp dport 12345 dnat to 198.51.100.20:443 comment "flvx forward:42 tcp"
|
||||
udp dport 12345 dnat to 198.51.100.20:443 comment "flvx forward:42 udp"
|
||||
```
|
||||
|
||||
第一阶段默认生成 masquerade:
|
||||
|
||||
```nft
|
||||
masquerade comment "flvx masquerade"
|
||||
```
|
||||
|
||||
原因是大多数纯 DNAT 场景需要回程可达;如果不做 SNAT,目标服务回包可能绕过转发节点导致连接失败。后续可增加高级开关允许用户关闭 masquerade。
|
||||
|
||||
### 原子同步策略
|
||||
|
||||
推荐节点级 reconcile,而不是逐条追加:
|
||||
|
||||
1. 从数据库查询该节点所有 nftables forward。
|
||||
2. 生成完整 `table inet flvx` 脚本。
|
||||
3. 通过 SSH 执行 `nft -f <tempfile>`。
|
||||
4. 成功后更新所有相关 `nft_rule_binding` 状态和 hash。
|
||||
|
||||
这样可以避免局部更新导致规则漂移,也能让“重建规则”与创建/更新走同一条路径。
|
||||
|
||||
## SSH 执行策略
|
||||
|
||||
基础要求:
|
||||
|
||||
- 默认超时 10-15 秒。
|
||||
- 支持密码和私钥认证。
|
||||
- 支持 `sudo nft ...`。
|
||||
- 执行前检查 `command -v nft`。
|
||||
- 执行前检查 `nft --version`,错误时提示安装 nftables。
|
||||
- 所有临时脚本写入 `/tmp/flvx-nft-<nonce>.nft`,执行后删除。
|
||||
|
||||
建议命令流程:
|
||||
|
||||
```sh
|
||||
cat > /tmp/flvx-nft-xxxx.nft <<'EOF'
|
||||
table inet flvx {
|
||||
...
|
||||
}
|
||||
EOF
|
||||
nft list table inet flvx >/dev/null 2>&1 && nft delete table inet flvx || true
|
||||
nft -f /tmp/flvx-nft-xxxx.nft
|
||||
rm -f /tmp/flvx-nft-xxxx.nft
|
||||
```
|
||||
|
||||
如果目标 nft 版本支持 `destroy table`,也可以把删除动作放进脚本:
|
||||
|
||||
```nft
|
||||
destroy table inet flvx
|
||||
table inet flvx {
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
实现时应按目标 nft 版本兼容性选择 `destroy` 或 shell 中先检测 `nft list table inet flvx`。
|
||||
|
||||
## 前端体验
|
||||
|
||||
### 节点页
|
||||
|
||||
节点表单新增“转发模式”:
|
||||
|
||||
- `Agent 节点`:默认,现有表单不变。
|
||||
- `nftables 节点`:显示 SSH 配置区块,隐藏 agent 安装相关提示。
|
||||
|
||||
nftables 节点列表操作:
|
||||
|
||||
- 测试 SSH
|
||||
- 重建规则
|
||||
- 清理 FLVX nftables 规则
|
||||
|
||||
隐藏或禁用:
|
||||
|
||||
- 安装命令
|
||||
- 升级
|
||||
- 回退
|
||||
- agent 协议开关
|
||||
- 实时 agent 指标入口
|
||||
|
||||
### 隧道页
|
||||
|
||||
隧道类型文案建议改为更明确的运行时说明:
|
||||
|
||||
- `Agent 端口转发`
|
||||
- `Agent 隧道转发`
|
||||
- `nftables 纯转发`
|
||||
|
||||
如果保持现有 `端口转发 / 隧道转发` 选择器,则在选择 nftables 入口节点后禁用隧道转发,并提示“不支持出口节点和转发链”。
|
||||
|
||||
### 转发页
|
||||
|
||||
选择 nftables 隧道后:
|
||||
|
||||
- 隐藏限速、每 IP 限速、最大连接数、Proxy Protocol。
|
||||
- 目标地址输入提示“第一阶段仅支持单目标 host:port”。
|
||||
- 创建/更新失败时显示远端 SSH 或 nftables 错误。
|
||||
|
||||
## 错误处理
|
||||
|
||||
- SSH 连接失败:返回“SSH 连接失败”,保留底层错误摘要。
|
||||
- 认证失败:返回“SSH 认证失败,请检查用户名和凭据”。
|
||||
- `nft` 不存在:返回“节点未安装 nftables”。
|
||||
- nft 脚本失败:返回 nft stderr 摘要,并记录到 `nft_rule_binding.last_error`。
|
||||
- 下发超时:标记 binding 为 `error`,允许用户重试“重建规则”。
|
||||
- 数据库成功但远端失败时,创建/更新路径应回滚数据库;批量重建路径不回滚业务规则,只记录错误。
|
||||
|
||||
## 安全边界
|
||||
|
||||
- SSH 凭据只在创建/更新时接收,列表 API 不返回明文。
|
||||
- 私钥和密码在数据库中加密保存。
|
||||
- 后端日志不得打印完整私钥、密码或 passphrase。
|
||||
- nft 脚本只由后端 renderer 生成,禁止直接拼接用户提交的自由文本。
|
||||
- `remoteAddr` 必须严格解析为 host/IP + port,端口必须为 1-65535。
|
||||
- `inPort` 仍复用现有端口占用校验。
|
||||
- comment 中只放 forward ID 和协议,不放用户输入。
|
||||
|
||||
## 与现有功能的关系
|
||||
|
||||
- `node/install` 对 nftables 节点返回错误或前端隐藏入口。
|
||||
- `node/check-status` 对 nftables 节点可返回 SSH 测试状态,而不是 agent 在线状态。
|
||||
- `forward/batch-redeploy` 对 nftables 规则执行节点级 reconcile。
|
||||
- `tunnel/batch-redeploy` 遇到 nftables 隧道时只重建相关 nftables 节点规则,不发送 GOST chain/service 命令。
|
||||
- federation 导入/共享第一阶段不支持 nftables 节点。
|
||||
- backup/import 应包含新增 node mode、SSH 配置和 binding 状态;导出时默认不导出 SSH 明文凭据。
|
||||
|
||||
## 测试计划
|
||||
|
||||
后端单元测试:
|
||||
|
||||
- nftables 节点不能创建隧道转发。
|
||||
- nftables 隧道不能包含出口节点或转发链。
|
||||
- agent 和 nftables 节点不能混在同一隧道。
|
||||
- nftables forward 拒绝限速、连接限制和 Proxy Protocol。
|
||||
- nftables forward 拒绝多目标 remoteAddr。
|
||||
- renderer 为 TCP/UDP 生成稳定脚本和 comment。
|
||||
- SSH runner 正确隐藏敏感信息并返回 stderr 摘要。
|
||||
|
||||
后端集成测试:
|
||||
|
||||
- 创建 nftables forward 时数据库和 binding 同步成功。
|
||||
- runtime 下发失败时创建回滚。
|
||||
- 更新失败时数据库和旧规则尽量恢复。
|
||||
- 删除失败时普通删除返回错误,强制删除保留清理提示。
|
||||
|
||||
前端验证:
|
||||
|
||||
- 节点表单按转发模式切换字段。
|
||||
- nftables 节点隐藏安装/升级/回退操作。
|
||||
- 隧道表单阻止 nftables 隧道转发配置。
|
||||
- 转发表单选择 nftables 隧道后隐藏不支持字段。
|
||||
|
||||
验证命令:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
## 实施顺序
|
||||
|
||||
1. 数据模型和 repository:新增字段、SSH 配置表、binding 表和查询方法。
|
||||
2. nftables runtime:实现 planner、renderer、SSH runner、manager。
|
||||
3. handler 校验:节点、隧道、转发 create/update/delete 接入 runtime。
|
||||
4. 前端节点表单:增加转发模式和 SSH 配置。
|
||||
5. 前端隧道/转发表单:按 nftables 能力收窄 UI。
|
||||
6. 批量重建和清理操作:提供运维入口。
|
||||
7. 测试与文案打磨。
|
||||
|
||||
## 第一阶段固定决策
|
||||
|
||||
本设计先固定以下选择,除非审核时调整:
|
||||
|
||||
- 第一阶段同时下发 TCP 和 UDP。
|
||||
- 第一阶段只支持单目标。
|
||||
- 第一阶段默认启用 masquerade。
|
||||
- nftables 节点的“在线状态”以 SSH 测试为准,而不是常驻连接。
|
||||
@@ -0,0 +1,284 @@
|
||||
# nftables 流量统计设计
|
||||
|
||||
**日期**: 2026-06-06
|
||||
**状态**: 待审核
|
||||
**作者**: Codex
|
||||
|
||||
## 概述
|
||||
|
||||
为 FLVX 的 `nftables` 转发模式补齐流量统计。当前 nftables 模式由面板通过 SSH 全量维护 `table inet flvx`,但没有 agent,因此不能复用 WebSocket 运行时上报。新方案由面板定时通过 SSH 拉取远端 nftables counter,计算增量后写入现有流量账本。
|
||||
|
||||
目标是让 nftables 转发在用户可见口径上尽量接近 agent 模式:
|
||||
|
||||
- forward 列表显示 `inFlow` / `outFlow`。
|
||||
- 用户、用户隧道、配额和流量策略继续生效。
|
||||
- 隧道监控继续获得分钟级 `tunnel_metric`。
|
||||
- 节点不需要安装新的 agent 或常驻进程。
|
||||
|
||||
## 背景
|
||||
|
||||
现有 agent 模式通过 `/flow/upload` 接收加密上报,handler 会把服务名解析为 `forward_id/user_id/user_tunnel_id`,再复用以下路径:
|
||||
|
||||
- `ApplyFlowUploadDeltasBatch` 更新 `forward`、`user`、`user_tunnel`。
|
||||
- `AddUserQuotaUsageBatch` 更新用户配额窗口。
|
||||
- `enforceUserQuotaIfNeeded` 和 `enforceFlowPolicies` 做约束 enforcement。
|
||||
- `recordTunnelMetricsFromForwardBatch` 写入分钟级隧道监控。
|
||||
|
||||
nftables 模式已经有 `nft_rule_binding` 记录规则应用状态,规则 comment 里包含 `forward_id`。这给 counter 到业务实体的映射提供了稳定锚点。
|
||||
|
||||
## 推荐方案
|
||||
|
||||
采用“面板 SSH 轮询 nftables counter”的方案:
|
||||
|
||||
1. 渲染 nftables 规则时,为每个 forward、协议和方向写入稳定 comment 和 `counter`。
|
||||
2. 后端定时扫描 `forward_mode = nftables` 的节点。
|
||||
3. 对每个节点通过 SSH 执行 `nft -j list table inet flvx`。
|
||||
4. 解析 JSON 规则,按 comment 得到 `forward_id/protocol/direction/bytes/packets`。
|
||||
5. 用数据库中的上次采样值计算 delta。
|
||||
6. 将 delta 转成现有 flow upload 内部结构,复用既有入账、配额、策略和监控逻辑。
|
||||
|
||||
不采用节点 crontab 或 systemd timer 回推。它会重新引入节点侧组件,削弱 nftables 模式“不安装 agent”的产品边界。
|
||||
|
||||
## 统计口径
|
||||
|
||||
正式入账使用 `forward` filter chain 的计数,不使用 NAT chain 的 DNAT 命中计数作为主口径。
|
||||
|
||||
原因:
|
||||
|
||||
- DNAT counter 表示规则命中,不一定代表后续转发成功。
|
||||
- filter forward chain 更接近实际经过内核转发的数据。
|
||||
- SNAT/masquerade 会改变包头,入账规则应在可稳定匹配目标服务地址和端口的位置统计。
|
||||
|
||||
方向定义:
|
||||
|
||||
| direction | nft 匹配 | 写入字段 |
|
||||
|-----------|----------|----------|
|
||||
| `to-target` | 外部客户端到目标服务 | `in_flow` |
|
||||
| `from-target` | 目标服务返回外部客户端 | `out_flow` |
|
||||
|
||||
用户总用量和配额仍按 `in_flow + out_flow` 计算。隧道 `traffic_ratio` 和 `flow` 倍率继续沿用 agent 模式逻辑,保证不同运行时模式的账单口径一致。
|
||||
|
||||
## nftables 规则设计
|
||||
|
||||
继续只维护 `table inet flvx`,避免触碰用户已有规则。每条 forward 对 TCP 和 UDP 各生成一组 DNAT 和统计规则。
|
||||
|
||||
示例:
|
||||
|
||||
```nft
|
||||
table inet flvx {
|
||||
chain prerouting {
|
||||
type nat hook prerouting priority dstnat; policy accept;
|
||||
tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"
|
||||
udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"
|
||||
}
|
||||
|
||||
chain postrouting {
|
||||
type nat hook postrouting priority srcnat; policy accept;
|
||||
masquerade comment "flvx masquerade"
|
||||
}
|
||||
|
||||
chain forward {
|
||||
type filter hook forward priority filter; policy accept;
|
||||
ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"
|
||||
ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"
|
||||
ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"
|
||||
ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
IPv6 目标使用 `ip6`:
|
||||
|
||||
```nft
|
||||
ip6 daddr 2001:db8::20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"
|
||||
ip6 saddr 2001:db8::20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"
|
||||
```
|
||||
|
||||
域名目标无法在 nftables 规则中动态匹配返回方向。统计第一阶段要求 nftables forward 的 `remoteAddr` host 必须是 IP 地址;如果当前纯转发实现允许域名,开启统计时应同步收紧校验。后续若要支持域名,应在规则同步时解析并固化 IP,同时明确 DNS 变化后的重建策略。
|
||||
|
||||
## Comment 格式
|
||||
|
||||
正式统计规则使用固定格式:
|
||||
|
||||
```text
|
||||
flvx forward:<forward_id> <direction> <protocol>
|
||||
```
|
||||
|
||||
字段:
|
||||
|
||||
- `forward_id`: 十进制整数。
|
||||
- `direction`: `to-target` 或 `from-target`。
|
||||
- `protocol`: `tcp` 或 `udp`。
|
||||
|
||||
DNAT 调试规则可使用 `dnat` direction,但 collector 不入账 `dnat`。后端只依赖 comment 解析,不依赖 nft handle,因为全量重建 table 会改变 handle。
|
||||
|
||||
## 数据模型
|
||||
|
||||
新增 `nft_counter_state` 表保存上次采样基线。
|
||||
|
||||
| 字段 | 说明 |
|
||||
|------|------|
|
||||
| `id` | 主键 |
|
||||
| `node_id` | nftables 节点 ID |
|
||||
| `forward_id` | 转发规则 ID |
|
||||
| `protocol` | `tcp` / `udp` |
|
||||
| `direction` | `to-target` / `from-target` |
|
||||
| `rule_hash` | 当前规则 hash |
|
||||
| `bytes` | 上次采样绝对字节数 |
|
||||
| `packets` | 上次采样绝对包数 |
|
||||
| `collected_time` | 上次采样时间 |
|
||||
| `created_time` | 创建时间 |
|
||||
| `updated_time` | 更新时间 |
|
||||
|
||||
唯一索引:
|
||||
|
||||
```text
|
||||
node_id, forward_id, protocol, direction
|
||||
```
|
||||
|
||||
GORM 模型必须定义 `TableName()`,字段 tag 保持 SQLite/PostgreSQL 兼容,不使用 `jsonb`、`serial` 等数据库专属类型。
|
||||
|
||||
## 后端组件
|
||||
|
||||
扩展 `go-backend/internal/runtime/nftables`:
|
||||
|
||||
| 组件 | 职责 |
|
||||
|------|------|
|
||||
| `CounterSample` | 表达单条 nft counter 采样 |
|
||||
| `Collector` | 对外提供 `Collect(ctx, cfg)` |
|
||||
| `SSHRunner.ListTableJSON` | 远端执行 `nft -j list table inet flvx` |
|
||||
| `ParseCounterSamples` | 解析 nft JSON 和 FLVX comment |
|
||||
|
||||
扩展 repository:
|
||||
|
||||
| 方法 | 职责 |
|
||||
|------|------|
|
||||
| `ListNftablesNodesForCollection` | 找到启用 nftables 且有 SSH 配置的节点 |
|
||||
| `GetNftCounterStatesByNode` | 读取节点上次 counter 基线 |
|
||||
| `UpsertNftCounterStates` | 批量刷新基线 |
|
||||
| `DeleteNftCounterStatesByForward` | forward 删除时清理状态 |
|
||||
|
||||
扩展 handler/job:
|
||||
|
||||
- 新增 `runNftablesTrafficCollectJob(now time.Time)`。
|
||||
- 默认每 60 秒运行一次。
|
||||
- 对节点采集设置并发上限,建议 3 到 5。
|
||||
- 单节点失败只记录日志和节点采集状态,不影响其他节点。
|
||||
|
||||
## 增量算法
|
||||
|
||||
collector 返回的是 nftables 的绝对 counter。入账前必须和上次基线做差。
|
||||
|
||||
规则:
|
||||
|
||||
- 无旧状态:只保存当前值作为基线,不入账。
|
||||
- `rule_hash` 变化:只刷新基线,不入账,避免新旧规则混算。
|
||||
- 新 bytes 大于等于旧 bytes:`delta = new - old`。
|
||||
- 新 bytes 小于旧 bytes:认为远端 table 重建、counter reset 或系统重启,只刷新基线,不入账。
|
||||
- delta 为 0:刷新采集时间,不入账。
|
||||
- 样本无法映射到有效 forward:忽略并记录 debug 日志。
|
||||
|
||||
同一 forward 的 TCP/UDP delta 要先聚合,再转换成现有账本:
|
||||
|
||||
- `to-target` bytes 聚合为原始 `bytesIn`。
|
||||
- `from-target` bytes 聚合为原始 `bytesOut`。
|
||||
- 入账时按 `traffic_ratio` 和 `tunnel.flow` 计算 scaled `InFlow` / `OutFlow`。
|
||||
- 配额使用 scaled 后的 `InFlow + OutFlow`。
|
||||
- `tunnel_metric` 使用原始 `bytesIn` / `bytesOut`。
|
||||
|
||||
## 入账路径
|
||||
|
||||
新增一个 nftables 专用的 batch builder,但输出沿用现有结构:
|
||||
|
||||
```go
|
||||
type nftTrafficDelta struct {
|
||||
ForwardID int64
|
||||
BytesIn int64
|
||||
BytesOut int64
|
||||
}
|
||||
```
|
||||
|
||||
处理流程:
|
||||
|
||||
1. 收集本轮所有 `forward_id`。
|
||||
2. 调用 `GetFlowUploadForwardMetas` 获取 `user_id/user_tunnel_id/tunnel_id/traffic_ratio/tunnel_flow`。
|
||||
3. 构造 `repo.FlowUploadCounterDelta`。
|
||||
4. 调用 `recordTunnelMetricsFromForwardBatch` 写监控。
|
||||
5. 抽出共享入账 helper,复用 `applyFlowDeltasWithFallback`、`applyQuotaUsageWithFallback`、`enforceUserQuotaIfNeeded` 和 `enforceFlowPolicies`。不要通过伪造 agent service name 去调用 agent 专用 builder。
|
||||
|
||||
不新增独立的 nftables 流量字段。`forward.in_flow/out_flow`、`user.in_flow/out_flow`、`user_tunnel.in_flow/out_flow` 仍是统一事实来源。
|
||||
|
||||
## 错误处理
|
||||
|
||||
采集错误分为三类:
|
||||
|
||||
| 类型 | 行为 |
|
||||
|------|------|
|
||||
| SSH 连接或认证失败 | 记录日志,保留下次继续采集 |
|
||||
| 远端无 `table inet flvx` | 视为规则未应用或被清理,记录 warning,不清空账本 |
|
||||
| JSON 解析失败 | 记录原始错误摘要,不入账 |
|
||||
|
||||
不要因为采集失败禁用 forward。流量统计失败和转发运行失败不是同一件事。
|
||||
|
||||
可在后续 UI 增加节点级采集状态,例如最近成功时间、最近错误。但第一步只要求后端具备日志和数据库状态即可。
|
||||
|
||||
## 与现有行为的关系
|
||||
|
||||
- agent 模式 `/flow/upload` 不变。
|
||||
- nftables 模式不新增节点侧 HTTP 回调。
|
||||
- 现有 `nft_rule_binding.rule_hash` 继续表示规则期望状态;counter state 用它判断采样是否跨规则版本。
|
||||
- `statistics_flow` 小时统计 job 不需要改,它基于用户总流量快照自然包含 nftables 入账结果。
|
||||
- 用户重置流量时不需要清空 nftables counter。重置只清业务账本;下一轮采集继续从 counter state 差值入账。
|
||||
|
||||
## 测试计划
|
||||
|
||||
后端单元测试:
|
||||
|
||||
- renderer 为 TCP/UDP、IPv4/IPv6 目标生成 `counter` 和稳定 comment。
|
||||
- comment parser 能识别合法格式,拒绝未知 direction/protocol。
|
||||
- nft JSON parser 能从 `nft -j list table` 输出中提取 bytes/packets。
|
||||
- delta 算法覆盖首次基线、正常增长、counter reset、rule_hash 变化和零增量。
|
||||
- batch builder 正确应用 `traffic_ratio` 和 `tunnel.flow`。
|
||||
|
||||
repository 测试:
|
||||
|
||||
- `nft_counter_state` 自动迁移。
|
||||
- upsert 在 SQLite 下可重复刷新。
|
||||
- forward 删除时清理 counter state。
|
||||
|
||||
handler/job 测试:
|
||||
|
||||
- 单节点采集成功会调用现有流量入账路径。
|
||||
- 单节点 SSH 失败不影响其他节点。
|
||||
- 无旧状态时不会误把历史 counter 入账。
|
||||
|
||||
验证命令:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
```
|
||||
|
||||
## 分阶段落地
|
||||
|
||||
第一阶段:
|
||||
|
||||
- 规则渲染加入 filter chain counter。
|
||||
- 实现 SSH collector、JSON parser、counter state 和后台 job。
|
||||
- 入账到现有账本和 tunnel metric。
|
||||
|
||||
第二阶段:
|
||||
|
||||
- UI 展示 nftables 采集状态。
|
||||
- 节点详情显示最近采集时间和最近错误。
|
||||
- 提供手动“采集一次”诊断按钮。
|
||||
|
||||
第三阶段:
|
||||
|
||||
- 探索域名目标的解析和重建策略。
|
||||
- 优化大量节点下的采集调度、退避和超时配置。
|
||||
|
||||
## 开放问题
|
||||
|
||||
- 采集周期默认 60 秒是否满足产品预期;如果需要更实时,可以降到 30 秒,但 SSH 压力会增加。
|
||||
- nftables 模式是否继续允许域名 remoteAddr。如果允许,需要先定义 DNS 固化和统计匹配规则。
|
||||
- 是否要在第一阶段暴露采集状态 API。推荐后端先记录,UI 后续补齐。
|
||||
@@ -240,6 +240,19 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if nftMode {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, forward.RemoteAddr, entryNodeIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(entryNodeIDs) == 0 {
|
||||
return nil, errors.New("nftables 转发缺少入口节点")
|
||||
}
|
||||
return nil, h.syncNftablesNode(entryNodeIDs[0])
|
||||
}
|
||||
ports, err := h.listForwardPorts(forward.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -23,6 +23,7 @@ import (
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/metrics"
|
||||
"go-backend/internal/monitoring"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/repo"
|
||||
"go-backend/internal/ws"
|
||||
@@ -31,11 +32,12 @@ import (
|
||||
)
|
||||
|
||||
type Handler struct {
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
metrics *metrics.IngestionService
|
||||
healthCheck *health.Checker
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
metrics *metrics.IngestionService
|
||||
healthCheck *health.Checker
|
||||
nftablesManager nftablesRuntimeManager
|
||||
|
||||
captchaMu sync.Mutex
|
||||
captchaTokens map[string]int64
|
||||
@@ -108,6 +110,7 @@ func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
metrics: metrics.NewIngestionService(repo),
|
||||
healthCheck: nil,
|
||||
nftablesManager: runtimenft.NewManager(nil),
|
||||
captchaTokens: make(map[string]int64),
|
||||
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||
nodeOnlineRedeployAt: make(map[int64]time.Time),
|
||||
@@ -190,6 +193,9 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/node/batch-upgrade", h.nodeBatchUpgrade)
|
||||
mux.HandleFunc("/api/v1/node/rollback", h.nodeRollback)
|
||||
mux.HandleFunc("/api/v1/node/releases", h.listReleases)
|
||||
mux.HandleFunc("/api/v1/node/nftables/test", h.nodeNftablesTest)
|
||||
mux.HandleFunc("/api/v1/node/nftables/reconcile", h.nodeNftablesReconcile)
|
||||
mux.HandleFunc("/api/v1/node/nftables/clear", h.nodeNftablesClear)
|
||||
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
|
||||
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
|
||||
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||
|
||||
@@ -20,7 +20,7 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
h.jobsCancel = cancel
|
||||
h.jobsStarted = true
|
||||
h.jobsWG.Add(7)
|
||||
h.jobsWG.Add(8)
|
||||
h.jobsMu.Unlock()
|
||||
|
||||
go h.runHourlyStatsLoop(ctx)
|
||||
@@ -30,6 +30,7 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
go h.runHealthChecks(ctx)
|
||||
go h.runTunnelQualityProber(ctx)
|
||||
go h.runValidateLicenseJob(ctx)
|
||||
go h.runNftablesTrafficCollectLoop(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) runValidateLicenseJob(ctx context.Context) {
|
||||
@@ -64,7 +65,7 @@ func (h *Handler) validateLicenseJob() {
|
||||
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
|
||||
|
||||
if err != nil {
|
||||
// Network error or timeout. Grace period by not revoking immediately here.
|
||||
return
|
||||
@@ -128,6 +129,21 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) {
|
||||
h.qualityProber.Start(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
ticker := time.NewTicker(time.Minute)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
h.runNftablesTrafficCollectJob(time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
|
||||
@@ -341,6 +341,11 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
inx := h.repo.NextIndex("node")
|
||||
forwardMode := defaultNodeForwardMode(asString(req["forwardMode"]))
|
||||
status := 0
|
||||
if forwardMode == "nftables" {
|
||||
status = 1
|
||||
}
|
||||
if err := h.repo.CreateNode(
|
||||
name,
|
||||
randomToken(16),
|
||||
@@ -357,7 +362,7 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
asInt(req["tls"], 0),
|
||||
asInt(req["socks"], 0),
|
||||
now,
|
||||
0,
|
||||
status,
|
||||
defaultString(asString(req["tcpListenAddr"]), "[::]"),
|
||||
defaultString(asString(req["udpListenAddr"]), "[::]"),
|
||||
inx,
|
||||
@@ -366,10 +371,20 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(asString(req["remoteToken"])),
|
||||
nullableText(asString(req["remoteConfig"])),
|
||||
nullableText(asString(req["extraIPs"])),
|
||||
forwardMode,
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
nodeID, err := h.findCreatedNodeID(name, serverIP)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.persistNodeSSHConfig(nodeID, req, forwardMode, now, false); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -417,6 +432,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwardMode := defaultNodeForwardMode(strings.TrimSpace(asString(req["forwardMode"])))
|
||||
if err := h.repo.UpdateNode(id,
|
||||
asString(req["name"]),
|
||||
serverIP,
|
||||
@@ -428,6 +444,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(strings.TrimSpace(asString(req["remark"]))),
|
||||
nullableUnixMilli(asInt64(req["expiryTime"], 0)),
|
||||
nullableText(normalizeNodeRenewalCycle(asString(req["renewalCycle"]))),
|
||||
forwardMode,
|
||||
newHTTP,
|
||||
newTLS,
|
||||
newSocks,
|
||||
@@ -438,9 +455,142 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.persistNodeSSHConfig(id, req, forwardMode, now, true); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if forwardMode == "nftables" && currentStatus != 1 {
|
||||
if err := h.repo.UpdateNodeStatus(id, 1); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) findCreatedNodeID(name, serverIP string) (int64, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return 0, errors.New("handler not initialized")
|
||||
}
|
||||
nodes, err := h.repo.ListNodes()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
for i := len(nodes) - 1; i >= 0; i-- {
|
||||
item := nodes[i]
|
||||
if asString(item["name"]) != name {
|
||||
continue
|
||||
}
|
||||
if asString(item["serverIp"]) != serverIP {
|
||||
continue
|
||||
}
|
||||
if nodeID := asInt64(item["id"], 0); nodeID > 0 {
|
||||
return nodeID, nil
|
||||
}
|
||||
}
|
||||
return 0, errors.New("节点创建成功,但未能查询到节点记录")
|
||||
}
|
||||
|
||||
func (h *Handler) persistNodeSSHConfig(nodeID int64, req map[string]interface{}, forwardMode string, now int64, preserveSecrets bool) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
if nodeID <= 0 {
|
||||
return errors.New("节点ID不能为空")
|
||||
}
|
||||
if forwardMode != "nftables" {
|
||||
return h.repo.DeleteNodeSSHConfig(nodeID)
|
||||
}
|
||||
cfgMap := asMap(req["sshConfig"])
|
||||
host := strings.TrimSpace(asString(cfgMap["host"]))
|
||||
if host == "" {
|
||||
host = strings.TrimSpace(asString(req["serverIp"]))
|
||||
}
|
||||
port := asInt(cfgMap["port"], 22)
|
||||
username := strings.TrimSpace(asString(cfgMap["username"]))
|
||||
authType := strings.TrimSpace(asString(cfgMap["authType"]))
|
||||
password := asString(cfgMap["password"])
|
||||
privateKey := asString(cfgMap["privateKey"])
|
||||
passphrase := asString(cfgMap["passphrase"])
|
||||
sudoMode := strings.TrimSpace(asString(cfgMap["sudoMode"]))
|
||||
|
||||
if preserveSecrets {
|
||||
existing, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return err
|
||||
}
|
||||
if existing != nil {
|
||||
if host == "" {
|
||||
host = strings.TrimSpace(existing.Host)
|
||||
}
|
||||
if port <= 0 {
|
||||
port = existing.Port
|
||||
}
|
||||
if username == "" {
|
||||
username = strings.TrimSpace(existing.Username)
|
||||
}
|
||||
if authType == "" {
|
||||
authType = strings.TrimSpace(existing.AuthType)
|
||||
}
|
||||
if strings.TrimSpace(password) == "" && existing.Password.Valid {
|
||||
password = existing.Password.String
|
||||
}
|
||||
if strings.TrimSpace(privateKey) == "" && existing.PrivateKey.Valid {
|
||||
privateKey = existing.PrivateKey.String
|
||||
}
|
||||
if strings.TrimSpace(passphrase) == "" && existing.Passphrase.Valid {
|
||||
passphrase = existing.Passphrase.String
|
||||
}
|
||||
if sudoMode == "" {
|
||||
sudoMode = strings.TrimSpace(existing.SudoMode)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if host == "" || username == "" {
|
||||
return errors.New("nftables 节点 SSH 配置不完整")
|
||||
}
|
||||
if port <= 0 || port > 65535 {
|
||||
return errors.New("nftables 节点 SSH 端口无效")
|
||||
}
|
||||
|
||||
authType = strings.ToLower(authType)
|
||||
switch authType {
|
||||
case "password":
|
||||
if strings.TrimSpace(password) == "" {
|
||||
return errors.New("nftables 节点 SSH 密码不能为空")
|
||||
}
|
||||
privateKey = ""
|
||||
case "private_key", "":
|
||||
authType = "private_key"
|
||||
if strings.TrimSpace(privateKey) == "" {
|
||||
return errors.New("nftables 节点 SSH 私钥不能为空")
|
||||
}
|
||||
password = ""
|
||||
default:
|
||||
return errors.New("nftables 节点 SSH 认证方式无效")
|
||||
}
|
||||
|
||||
switch strings.ToLower(sudoMode) {
|
||||
case "", "none":
|
||||
sudoMode = "none"
|
||||
case "sudo", "sudo_su":
|
||||
default:
|
||||
return errors.New("nftables 节点 sudo 模式无效")
|
||||
}
|
||||
|
||||
return h.repo.UpsertNodeSSHConfig(nodeID, repo.NftSSHConfigInput{
|
||||
Host: host,
|
||||
Port: port,
|
||||
Username: username,
|
||||
AuthType: authType,
|
||||
Password: password,
|
||||
PrivateKey: privateKey,
|
||||
Passphrase: passphrase,
|
||||
SudoMode: sudoMode,
|
||||
}, now)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -648,6 +798,16 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if strings.TrimSpace(inIP) == "" {
|
||||
inIP = buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
|
||||
}
|
||||
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
|
||||
for _, inNode := range runtimeState.InNodes {
|
||||
if inNode.NodeID > 0 {
|
||||
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateNftablesTunnelStateTx(tx, entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if len(runtimeState.InNodes) > 0 {
|
||||
firstNodeID := runtimeState.InNodes[0].NodeID
|
||||
@@ -934,6 +1094,16 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
runtimeState.TunnelID = id
|
||||
runtimeState.IPPreference = ipPreference
|
||||
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
|
||||
for _, inNode := range runtimeState.InNodes {
|
||||
if inNode.NodeID > 0 {
|
||||
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateNftablesTunnelState(entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
|
||||
|
||||
@@ -1700,6 +1870,24 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, tunnelErr)
|
||||
continue
|
||||
}
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(tunnelID); modeErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
fail++
|
||||
failures = appendBatchFailureReason(failures, tunnelID, tunnelName, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
if reconcileErr := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); reconcileErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, reconcileErr)
|
||||
continue
|
||||
}
|
||||
success++
|
||||
continue
|
||||
}
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, err)
|
||||
@@ -1909,6 +2097,17 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
port = 10000
|
||||
}
|
||||
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
|
||||
isNftTunnel, _, err := h.tunnelUsesNftables(tunnelID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if isNftTunnel {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
inIp := strings.TrimSpace(asString(req["inIp"]))
|
||||
if inIp != "" && len(entryNodes) > 1 {
|
||||
response.WriteJSON(w, response.ErrDefault("多入口隧道的转发不支持自定义监听IP"))
|
||||
@@ -2016,6 +2215,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if remoteAddr == "" {
|
||||
remoteAddr = forward.RemoteAddr
|
||||
}
|
||||
isNftTunnel, entryNodes, err := h.tunnelUsesNftables(tunnelID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if isNftTunnel {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
if actorRole != 0 && !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
@@ -2202,7 +2412,17 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
var nftNodeID int64
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, modeErr.Error()))
|
||||
return
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("nftables 转发缺少入口节点"))
|
||||
return
|
||||
}
|
||||
nftNodeID = entryNodeIDs[0]
|
||||
} else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -2210,6 +2430,12 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if nftNodeID > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -2218,7 +2444,7 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
_, _, _, err := h.resolveForwardAccess(r, id)
|
||||
forward, _, _, err := h.resolveForwardAccess(r, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, errForwardNotFound) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发不存在"))
|
||||
@@ -2234,6 +2460,13 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
_ = h.repo.DeleteNftRuleBindingsByForward(id)
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr == nil && nftMode && len(entryNodeIDs) > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -2351,7 +2584,19 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
var nftNodeID int64
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
f++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
nftNodeID = entryNodeIDs[0]
|
||||
} else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
@@ -2359,9 +2604,16 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if err := h.deleteForwardByID(id); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
continue
|
||||
}
|
||||
if nftNodeID > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
s++
|
||||
}
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||
}
|
||||
@@ -2462,6 +2714,24 @@ func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
f++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
@@ -4705,6 +4975,13 @@ func asMapSlice(v interface{}) []map[string]interface{} {
|
||||
return out
|
||||
}
|
||||
|
||||
func asMap(v interface{}) map[string]interface{} {
|
||||
if m, ok := v.(map[string]interface{}); ok && m != nil {
|
||||
return m
|
||||
}
|
||||
return map[string]interface{}{}
|
||||
}
|
||||
|
||||
func asString(v interface{}) string {
|
||||
switch t := v.(type) {
|
||||
case nil:
|
||||
@@ -4843,6 +5120,15 @@ func normalizeNodeRenewalCycle(v string) string {
|
||||
}
|
||||
}
|
||||
|
||||
func defaultNodeForwardMode(mode string) string {
|
||||
switch strings.TrimSpace(strings.ToLower(mode)) {
|
||||
case "nftables":
|
||||
return "nftables"
|
||||
default:
|
||||
return "agent"
|
||||
}
|
||||
}
|
||||
|
||||
func nullableInt(v *int64) interface{} {
|
||||
if v == nil {
|
||||
return nil
|
||||
|
||||
@@ -0,0 +1,371 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type nftablesRuntimeManager interface {
|
||||
Test(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
Reconcile(ctx context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error)
|
||||
Clear(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
CollectCounters(ctx context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error)
|
||||
}
|
||||
|
||||
func isNftablesForwardMode(mode string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(mode), runtimenft.ModeNftables)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeUsesNftables(nodeID int64) (bool, error) {
|
||||
return h.nodeUsesNftablesTx(nil, nodeID)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeUsesNftablesTx(tx *gorm.DB, nodeID int64) (bool, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return false, errors.New("handler not initialized")
|
||||
}
|
||||
var (
|
||||
mode string
|
||||
err error
|
||||
)
|
||||
if tx != nil {
|
||||
mode, err = h.repo.GetNodeForwardModeTx(tx, nodeID)
|
||||
} else {
|
||||
mode, err = h.repo.GetNodeForwardMode(nodeID)
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return isNftablesForwardMode(mode), nil
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelUsesNftables(tunnelID int64) (bool, []int64, error) {
|
||||
entryNodeIDs, err := h.tunnelEntryNodeIDs(tunnelID)
|
||||
if err != nil {
|
||||
return false, nil, err
|
||||
}
|
||||
for _, nodeID := range entryNodeIDs {
|
||||
ok, modeErr := h.nodeUsesNftables(nodeID)
|
||||
if modeErr != nil {
|
||||
return false, nil, modeErr
|
||||
}
|
||||
if ok {
|
||||
return true, entryNodeIDs, nil
|
||||
}
|
||||
}
|
||||
return false, entryNodeIDs, nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesForwardRequest(tunnel *tunnelRecord, remoteAddr string, entryNodeIDs []int64) error {
|
||||
if tunnel == nil {
|
||||
return errors.New("隧道不存在")
|
||||
}
|
||||
if tunnel.Type != 1 {
|
||||
return errors.New("nftables 节点仅支持直连隧道")
|
||||
}
|
||||
if len(entryNodeIDs) != 1 {
|
||||
return errors.New("nftables 节点仅支持单入口隧道")
|
||||
}
|
||||
target, err := runtimenft.ParseSingleTarget(remoteAddr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if net.ParseIP(strings.Trim(strings.TrimSpace(target.Host), "[]")) == nil {
|
||||
return errors.New("nftables 节点仅支持 IP 目标地址")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func sshConfigFromModel(cfg *model.NodeSSHConfig) (runtimenft.SSHConfig, error) {
|
||||
if cfg == nil {
|
||||
return runtimenft.SSHConfig{}, errors.New("节点缺少 SSH 配置")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Host) == "" || strings.TrimSpace(cfg.Username) == "" {
|
||||
return runtimenft.SSHConfig{}, errors.New("节点 SSH 配置不完整")
|
||||
}
|
||||
return runtimenft.SSHConfig{
|
||||
Host: strings.TrimSpace(cfg.Host),
|
||||
Port: cfg.Port,
|
||||
Username: strings.TrimSpace(cfg.Username),
|
||||
AuthType: strings.TrimSpace(cfg.AuthType),
|
||||
Password: cfg.Password.String,
|
||||
PrivateKey: cfg.PrivateKey.String,
|
||||
Passphrase: cfg.Passphrase.String,
|
||||
SudoMode: strings.TrimSpace(cfg.SudoMode),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesTunnelState(entryNodeIDs []int64) error {
|
||||
return h.validateNftablesTunnelStateTx(nil, entryNodeIDs)
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesTunnelStateTx(tx *gorm.DB, entryNodeIDs []int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
for _, nodeID := range entryNodeIDs {
|
||||
isNft, err := h.nodeUsesNftablesTx(tx, nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !isNft {
|
||||
continue
|
||||
}
|
||||
var cfg *model.NodeSSHConfig
|
||||
if tx != nil {
|
||||
cfg, err = h.repo.GetNodeSSHConfigTx(tx, nodeID)
|
||||
} else {
|
||||
cfg, err = h.repo.GetNodeSSHConfig(nodeID)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return errors.New("nftables 节点缺少 SSH 配置")
|
||||
}
|
||||
return err
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
|
||||
return fmt.Errorf("nftables 节点能力校验失败: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) buildNftablesNodePlan(nodeID int64) (runtimenft.NodePlan, *model.NodeSSHConfig, error) {
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
forwards, err := h.repo.ListActiveForwardsByNode(nodeID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
plan := runtimenft.NodePlan{NodeID: nodeID, Rules: make([]runtimenft.Rule, 0, len(forwards))}
|
||||
for i := range forwards {
|
||||
forward := &forwards[i]
|
||||
tunnel, err := h.getTunnelRecord(forward.TunnelID)
|
||||
if err != nil || tunnel == nil || tunnel.Status != 1 {
|
||||
continue
|
||||
}
|
||||
entryNodeIDs, err := h.tunnelEntryNodeIDs(forward.TunnelID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
if len(entryNodeIDs) != 1 || entryNodeIDs[0] != nodeID {
|
||||
continue
|
||||
}
|
||||
if err := h.validateNftablesForwardRequest(tunnel, forward.RemoteAddr, entryNodeIDs); err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
ports, err := h.listForwardPorts(forward.ID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
for _, fp := range ports {
|
||||
if fp.NodeID != nodeID {
|
||||
continue
|
||||
}
|
||||
target, err := runtimenft.ParseSingleTarget(forward.RemoteAddr)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
plan.Rules = append(plan.Rules, runtimenft.Rule{
|
||||
ForwardID: forward.ID,
|
||||
InPort: fp.Port,
|
||||
BindIP: strings.TrimSpace(fp.InIP),
|
||||
TargetHost: target.Host,
|
||||
TargetPort: target.Port,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
})
|
||||
}
|
||||
}
|
||||
return plan, cfg, nil
|
||||
}
|
||||
|
||||
func (h *Handler) syncNftablesNode(nodeID int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
plan, cfgModel, err := h.buildNftablesNodePlan(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfgModel)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
result, err := h.nftablesManager.Reconcile(context.Background(), sshCfg, plan)
|
||||
now := time.Now().UnixMilli()
|
||||
if err != nil {
|
||||
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
for _, binding := range bindings {
|
||||
_ = h.repo.MarkNftRuleBindingError(binding.ForwardID, nodeID, err.Error(), now)
|
||||
}
|
||||
return err
|
||||
}
|
||||
activeForwardIDs := make(map[int64]struct{}, len(plan.Rules))
|
||||
for _, rule := range plan.Rules {
|
||||
activeForwardIDs[rule.ForwardID] = struct{}{}
|
||||
hash := result.Hashes[rule.ForwardID]
|
||||
_ = h.repo.UpsertNftRuleBinding(modelToRuleBindingInput(nodeID, rule, hash), now)
|
||||
}
|
||||
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
for _, binding := range bindings {
|
||||
if _, ok := activeForwardIDs[binding.ForwardID]; ok {
|
||||
continue
|
||||
}
|
||||
_ = h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func modelToRuleBindingInput(nodeID int64, rule runtimenft.Rule, hash string) repo.NftRuleBindingInput {
|
||||
return repo.NftRuleBindingInput{
|
||||
ForwardID: rule.ForwardID,
|
||||
NodeID: nodeID,
|
||||
InPort: rule.InPort,
|
||||
Protocols: strings.Join(rule.Protocols, ","),
|
||||
TargetAddr: fmt.Sprintf("%s:%d", rule.TargetHost, rule.TargetPort),
|
||||
BindIP: rule.BindIP,
|
||||
RuleHash: hash,
|
||||
Status: runtimenft.StatusApplied,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) nftablesNodeIDFromRequest(r *http.Request, w http.ResponseWriter) (int64, bool) {
|
||||
nodeID := asInt64FromBodyKey(r, w, "nodeId")
|
||||
if nodeID <= 0 {
|
||||
return 0, false
|
||||
}
|
||||
return nodeID, true
|
||||
}
|
||||
|
||||
func (h *Handler) loadNftablesSSHConfig(nodeID int64) (runtimenft.SSHConfig, error) {
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return runtimenft.SSHConfig{}, errors.New("nftables 节点缺少 SSH 配置")
|
||||
}
|
||||
return runtimenft.SSHConfig{}, err
|
||||
}
|
||||
return sshConfigFromModel(cfg)
|
||||
}
|
||||
|
||||
func (h *Handler) clearNftablesNode(nodeID int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
if err := h.nftablesManager.Clear(context.Background(), sshCfg); err != nil {
|
||||
return err
|
||||
}
|
||||
bindings, listErr := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
if listErr != nil {
|
||||
return listErr
|
||||
}
|
||||
for _, binding := range bindings {
|
||||
if err := h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) reconcileNftablesNodeByRequest(nodeID int64) error {
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !usesNft {
|
||||
return errors.New("节点未启用 nftables 转发模式")
|
||||
}
|
||||
return h.syncNftablesNode(nodeID)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesTest(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if !usesNft {
|
||||
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
|
||||
return
|
||||
}
|
||||
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "nftables manager not initialized"))
|
||||
return
|
||||
}
|
||||
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesReconcile(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := h.reconcileNftablesNodeByRequest(nodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesClear(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if !usesNft {
|
||||
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
|
||||
return
|
||||
}
|
||||
if err := h.clearNftablesNode(nodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
@@ -0,0 +1,556 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type fakeNftablesManager struct {
|
||||
testErr error
|
||||
reconcileErr error
|
||||
reconcileHit int
|
||||
clearErr error
|
||||
clearHit int
|
||||
collectErr error
|
||||
collectHit int
|
||||
counterSamples []runtimenft.CounterSample
|
||||
lastConfig runtimenft.SSHConfig
|
||||
lastPlan runtimenft.NodePlan
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error {
|
||||
f.lastConfig = cfg
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) {
|
||||
f.reconcileHit++
|
||||
f.lastConfig = cfg
|
||||
f.lastPlan = plan
|
||||
if f.reconcileErr != nil {
|
||||
return runtimenft.ApplyResult{}, f.reconcileErr
|
||||
}
|
||||
return runtimenft.ApplyResult{
|
||||
NodeID: plan.NodeID,
|
||||
Script: "table inet flvx {}",
|
||||
Hashes: map[int64]string{plan.NodeID: "hash"},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error {
|
||||
f.clearHit++
|
||||
return f.clearErr
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) {
|
||||
f.collectHit++
|
||||
f.lastConfig = cfg
|
||||
if f.collectErr != nil {
|
||||
return nil, f.collectErr
|
||||
}
|
||||
return f.counterSamples, nil
|
||||
}
|
||||
|
||||
type nftablesTestFixture struct {
|
||||
handler *Handler
|
||||
nodeID int64
|
||||
}
|
||||
|
||||
func TestTunnelCreateRejectsNftablesEntryNodeWithoutSSHConfig(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
err := fixture.handler.validateNftablesTunnelState([]int64{fixture.nodeID})
|
||||
if err == nil {
|
||||
t.Fatalf("expected validation failure")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "SSH") {
|
||||
t.Fatalf("expected SSH config validation error, got %q", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateRejectsNftablesEntryNodeWhenCapabilityTestFails(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{testErr: errors.New("ssh failed")}
|
||||
h.nftablesManager = manager
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
err := h.validateNftablesTunnelState([]int64{fixture.nodeID})
|
||||
if err == nil {
|
||||
t.Fatalf("expected validation failure")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "ssh failed") {
|
||||
t.Fatalf("expected capability error in response, got %q", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncForwardServicesWithWarningsUsesNftablesRuntime(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
|
||||
warnings, err := h.syncForwardServicesWithWarnings(forward, "UpdateService", true)
|
||||
if err != nil {
|
||||
t.Fatalf("sync forward services: %v", err)
|
||||
}
|
||||
if len(warnings) != 0 {
|
||||
t.Fatalf("expected no warnings, got %v", warnings)
|
||||
}
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected nftables reconcile to run once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if manager.lastPlan.NodeID != fixture.nodeID {
|
||||
t.Fatalf("expected plan for node %d, got %+v", fixture.nodeID, manager.lastPlan)
|
||||
}
|
||||
if len(manager.lastPlan.Rules) != 1 || manager.lastPlan.Rules[0].ForwardID != forward.ID {
|
||||
t.Fatalf("unexpected plan: %+v", manager.lastPlan)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeNftablesTestEndpointRunsCapabilityCheck(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
manager := &fakeNftablesManager{}
|
||||
fixture.handler.nftablesManager = manager
|
||||
|
||||
res := postJSONToHandler(t, fixture.handler.nodeNftablesTest, map[string]int64{"nodeId": fixture.nodeID})
|
||||
assertNftablesSuccess(t, res)
|
||||
if manager.lastConfig.Host != "203.0.113.10" {
|
||||
t.Fatalf("expected SSH config to be passed to manager, got %+v", manager.lastConfig)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeNftablesReconcileEndpointPersistsBindings(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, fixture.handler, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, fixture.handler, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
fixture.handler.nftablesManager = manager
|
||||
|
||||
res := postJSONToHandler(t, fixture.handler.nodeNftablesReconcile, map[string]int64{"nodeId": fixture.nodeID})
|
||||
assertNftablesSuccess(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
bindings, err := fixture.handler.repo.ListNftRuleBindingsByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("list bindings: %v", err)
|
||||
}
|
||||
if len(bindings) != 1 || bindings[0].ForwardID != forward.ID {
|
||||
t.Fatalf("unexpected bindings: %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeNftablesClearEndpointClearsBindings(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
now := time.Now().UnixMilli()
|
||||
if err := fixture.handler.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: 99,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 24000,
|
||||
Protocols: "tcp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, now); err != nil {
|
||||
t.Fatalf("seed binding: %v", err)
|
||||
}
|
||||
manager := &fakeNftablesManager{}
|
||||
fixture.handler.nftablesManager = manager
|
||||
|
||||
res := postJSONToHandler(t, fixture.handler.nodeNftablesClear, map[string]int64{"nodeId": fixture.nodeID})
|
||||
assertNftablesSuccess(t, res)
|
||||
if manager.clearHit != 1 {
|
||||
t.Fatalf("expected clear once, got %d", manager.clearHit)
|
||||
}
|
||||
if bindings, err := fixture.handler.repo.ListNftRuleBindingsByNode(fixture.nodeID); err != nil {
|
||||
t.Fatalf("list bindings after clear: %v", err)
|
||||
} else if len(bindings) != 0 {
|
||||
t.Fatalf("expected bindings to be cleared, got %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeCreatePersistsNftablesSSHConfig(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
req := newAuthenticatedJSONRequest(t, map[string]interface{}{
|
||||
"name": "nft-node-created",
|
||||
"serverIp": "203.0.113.20",
|
||||
"serverIpV4": "203.0.113.20",
|
||||
"port": "20000-20100",
|
||||
"forwardMode": "nftables",
|
||||
"sshConfig": map[string]interface{}{
|
||||
"host": "203.0.113.21",
|
||||
"port": 2222,
|
||||
"username": "root",
|
||||
"authType": "private_key",
|
||||
"privateKey": "TEST-PRIVATE-KEY",
|
||||
"passphrase": "secret",
|
||||
"sudoMode": "sudo",
|
||||
},
|
||||
})
|
||||
res := httptest.NewRecorder()
|
||||
fixture.handler.nodeCreate(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
|
||||
nodes, err := fixture.handler.repo.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("list nodes: %v", err)
|
||||
}
|
||||
var createdNodeID int64
|
||||
for _, item := range nodes {
|
||||
if item["name"] == "nft-node-created" {
|
||||
createdNodeID = item["id"].(int64)
|
||||
break
|
||||
}
|
||||
}
|
||||
if createdNodeID <= 0 {
|
||||
t.Fatalf("expected created node to exist")
|
||||
}
|
||||
createdNode, err := fixture.handler.repo.GetNodeRecord(createdNodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load created node: %v", err)
|
||||
}
|
||||
if createdNode == nil {
|
||||
t.Fatal("expected created node record, got nil")
|
||||
}
|
||||
if createdNode.Status != 1 {
|
||||
t.Fatalf("expected nftables node to be online, got status %d", createdNode.Status)
|
||||
}
|
||||
cfg, err := fixture.handler.repo.GetNodeSSHConfig(createdNodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load ssh config: %v", err)
|
||||
}
|
||||
if cfg.Host != "203.0.113.21" || cfg.Port != 2222 || cfg.Username != "root" || cfg.AuthType != "private_key" {
|
||||
t.Fatalf("unexpected ssh config: %+v", cfg)
|
||||
}
|
||||
if !cfg.PrivateKey.Valid || cfg.PrivateKey.String != "TEST-PRIVATE-KEY" {
|
||||
t.Fatalf("expected private key to persist, got %+v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeUpdatePreservesExistingNftablesSecretsWhenFieldsOmitted(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string]interface{}{
|
||||
"id": fixture.nodeID,
|
||||
"name": "nft-node-updated",
|
||||
"serverIp": "198.51.100.10",
|
||||
"serverIpV4": "198.51.100.10",
|
||||
"port": "1000-65535",
|
||||
"forwardMode": "nftables",
|
||||
"sshConfig": map[string]interface{}{
|
||||
"host": "203.0.113.30",
|
||||
"port": 22,
|
||||
"username": "admin",
|
||||
"authType": "password",
|
||||
"sudoMode": "none",
|
||||
},
|
||||
})
|
||||
res := httptest.NewRecorder()
|
||||
fixture.handler.nodeUpdate(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
|
||||
cfg, err := fixture.handler.repo.GetNodeSSHConfig(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load ssh config: %v", err)
|
||||
}
|
||||
if cfg.Host != "203.0.113.30" || cfg.Username != "admin" || cfg.AuthType != "password" {
|
||||
t.Fatalf("unexpected ssh config after update: %+v", cfg)
|
||||
}
|
||||
if !cfg.Password.Valid || cfg.Password.String != "secret" {
|
||||
t.Fatalf("expected password secret to be preserved, got %+v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateNftablesForwardRequestRejectsHostnameTarget(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
tunnel, err := h.getTunnelRecord(tunnelID)
|
||||
if err != nil {
|
||||
t.Fatalf("load tunnel: %v", err)
|
||||
}
|
||||
|
||||
err = h.validateNftablesForwardRequest(tunnel, "example.com:443", []int64{fixture.nodeID})
|
||||
if err == nil {
|
||||
t.Fatalf("expected hostname target to be rejected")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "IP") {
|
||||
t.Fatalf("expected IP literal validation error, got %q", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardDeleteReconcilesNftablesAfterDBDelete(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string]int64{"id": forward.ID})
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardDelete(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if len(manager.lastPlan.Rules) != 0 {
|
||||
t.Fatalf("expected reconcile after DB delete to render no rules, got %+v", manager.lastPlan.Rules)
|
||||
}
|
||||
if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) {
|
||||
t.Fatalf("expected forward to be deleted, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardForceDeleteRemovesNftablesBindingAndReconciles(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: forward.ID,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 20000,
|
||||
Protocols: "tcp,udp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("seed binding: %v", err)
|
||||
}
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string]int64{"id": forward.ID})
|
||||
req.URL.Path = "/api/v1/forward/force-delete"
|
||||
res := httptest.NewRecorder()
|
||||
mux := http.NewServeMux()
|
||||
h.Register(mux)
|
||||
mux.ServeHTTP(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) {
|
||||
t.Fatalf("expected forward to be deleted, got %v", err)
|
||||
}
|
||||
if bindings, err := h.repo.ListNftRuleBindingsByNode(fixture.nodeID); err != nil {
|
||||
t.Fatalf("list bindings after delete: %v", err)
|
||||
} else if len(bindings) != 0 {
|
||||
t.Fatalf("expected no bindings after delete, got %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardBatchDeleteReconcilesNftablesAfterDBDelete(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {forward.ID}})
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardBatchDelete(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if len(manager.lastPlan.Rules) != 0 {
|
||||
t.Fatalf("expected reconcile after DB delete to render no rules, got %+v", manager.lastPlan.Rules)
|
||||
}
|
||||
if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) {
|
||||
t.Fatalf("expected forward to be deleted, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardBatchRedeployUsesNftablesReconcile(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {forward.ID}})
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardBatchRedeploy(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelBatchRedeployUsesNftablesReconcile(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {tunnelID}})
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelBatchRedeploy(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
}
|
||||
|
||||
func setupNftablesHandler(t *testing.T) nftablesTestFixture {
|
||||
t.Helper()
|
||||
|
||||
dbPath := filepath.Join(t.TempDir(), "handler-nftables.sqlite")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
|
||||
h := New(r, "test-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
if _, err := r.CreateUser("admin", "hash", 0, now+86400000, 1, 1, 100, 1, 0, now); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
if err := r.CreateNode("nft-node", "secret", "198.51.100.10", nil, nil, "1000-65535", nil, nil, nil, nil, nil, 0, 0, 0, now, 1, "", "", 1, 0, nil, nil, nil, nil, "nftables"); err != nil {
|
||||
t.Fatalf("create node: %v", err)
|
||||
}
|
||||
node, err := r.GetNodeRecord(1)
|
||||
if err != nil || node == nil {
|
||||
t.Fatalf("get node: %v", err)
|
||||
}
|
||||
return nftablesTestFixture{handler: h, nodeID: node.ID}
|
||||
}
|
||||
|
||||
func seedNftablesSSHConfig(t *testing.T, h *Handler, nodeID int64) {
|
||||
t.Helper()
|
||||
if err := h.repo.UpsertNodeSSHConfig(nodeID, repo.NftSSHConfigInput{
|
||||
Host: "203.0.113.10",
|
||||
Port: 22,
|
||||
Username: "root",
|
||||
AuthType: "password",
|
||||
Password: "secret",
|
||||
SudoMode: "none",
|
||||
}, time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("upsert ssh config: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func seedTunnelForNftables(t *testing.T, h *Handler, name string, nodeID int64) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
tx := h.repo.BeginTx()
|
||||
if tx == nil {
|
||||
t.Fatal("begin tx: nil transaction")
|
||||
}
|
||||
if tx.Error != nil {
|
||||
t.Fatalf("begin tx: %v", tx.Error)
|
||||
}
|
||||
tunnelID, err := h.repo.CreateTunnelTx(tx, name, 1, 1, 1, now, 1, nil, 1, "", "", 0)
|
||||
if err != nil {
|
||||
_ = tx.Rollback().Error
|
||||
t.Fatalf("create tunnel: %v", err)
|
||||
}
|
||||
if err := h.repo.CreateChainTunnelTx(tx, tunnelID, "1", nodeID, sql.NullInt64{}, "", 1, "tls", ""); err != nil {
|
||||
_ = tx.Rollback().Error
|
||||
t.Fatalf("create chain tunnel: %v", err)
|
||||
}
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
_ = tx.Rollback().Error
|
||||
t.Fatalf("commit tx: %v", err)
|
||||
}
|
||||
return tunnelID
|
||||
}
|
||||
|
||||
func seedForwardForNftables(t *testing.T, h *Handler, tunnelID, nodeID int64, remoteAddr string) *forwardRecord {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
forwardID, err := h.repo.CreateForwardTx(
|
||||
1, "admin", "nft-forward", tunnelID, remoteAddr, "fifo", now, 1,
|
||||
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("create forward: %v", err)
|
||||
}
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("get forward: %v", err)
|
||||
}
|
||||
return forward
|
||||
}
|
||||
|
||||
func postJSONToHandler(t *testing.T, fn func(http.ResponseWriter, *http.Request), payload any) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
|
||||
res := httptest.NewRecorder()
|
||||
fn(res, req)
|
||||
return res
|
||||
}
|
||||
|
||||
func newAuthenticatedJSONRequest(t *testing.T, payload any) *http.Request {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
|
||||
token, err := auth.GenerateToken(1, "admin", 0, "test-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("create token: %v", err)
|
||||
}
|
||||
req.Header.Set("Authorization", token)
|
||||
claims, ok := auth.ValidateToken(token, "test-secret")
|
||||
if !ok {
|
||||
t.Fatalf("validate token failed")
|
||||
}
|
||||
return req.WithContext(context.WithValue(req.Context(), middleware.ClaimsContextKey, claims))
|
||||
}
|
||||
|
||||
func assertNftablesSuccess(t *testing.T, res *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
}
|
||||
|
||||
func assertNftablesSuccessWithBody(t *testing.T, res *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected HTTP %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
if err := json.NewDecoder(res.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success, got %+v", payload)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,400 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"math"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type nftTrafficDelta struct {
|
||||
ForwardID int64
|
||||
BytesIn int64
|
||||
BytesOut int64
|
||||
}
|
||||
|
||||
type nftCounterStateKey struct {
|
||||
forwardID int64
|
||||
protocol string
|
||||
direction string
|
||||
}
|
||||
|
||||
func (h *Handler) runNftablesTrafficCollectJob(now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
nodes, err := h.repo.ListNftablesNodesForCollection()
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=list_nodes err=%v", err)
|
||||
return
|
||||
}
|
||||
for i := range nodes {
|
||||
node := &nodes[i]
|
||||
h.collectNftablesNodeTraffic(node.NodeID, &node.Config, now)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) collectNftablesNodeTraffic(nodeID int64, cfgModel *model.NodeSSHConfig, now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
log.Printf("nftables traffic collection failed op=collect node_id=%d err=%v", nodeID, "nftables manager not initialized")
|
||||
return
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfgModel)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=ssh_config node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
samples, err := h.nftablesManager.CollectCounters(context.Background(), sshCfg)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=collect node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
oldStates, err := h.repo.GetNftCounterStatesByNode(nodeID)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=list_states node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=list_bindings node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
hashes := make(map[int64]string, len(bindings))
|
||||
for _, binding := range bindings {
|
||||
if strings.ToLower(strings.TrimSpace(binding.Status)) != runtimenft.StatusApplied {
|
||||
continue
|
||||
}
|
||||
ruleHash := strings.TrimSpace(binding.RuleHash)
|
||||
if ruleHash == "" {
|
||||
continue
|
||||
}
|
||||
hashes[binding.ForwardID] = ruleHash
|
||||
}
|
||||
|
||||
nowMs := now.UnixMilli()
|
||||
boundSamples := filterNftCounterSamplesWithBinding(samples, hashes)
|
||||
deltas, newStates := buildNftCounterDeltas(nodeID, boundSamples, oldStates, hashes, nowMs)
|
||||
if len(newStates) == 0 {
|
||||
if len(deltas) != 0 {
|
||||
log.Printf("nftables traffic collection skipped suspicious deltas without states node_id=%d deltas=%d", nodeID, len(deltas))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
var metas map[int64]repo.FlowUploadForwardMeta
|
||||
forwardIDs := make([]int64, 0, len(deltas))
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID > 0 {
|
||||
forwardIDs = append(forwardIDs, delta.ForwardID)
|
||||
}
|
||||
}
|
||||
if len(deltas) != 0 {
|
||||
metas, err = h.repo.GetFlowUploadForwardMetas(forwardIDs)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=load_flow_metas node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
if missingForwardID, ok := firstNftDeltaMissingMeta(deltas, metas); ok {
|
||||
log.Printf("nftables traffic collection skipped state advance op=missing_flow_meta node_id=%d forward_id=%d", nodeID, missingForwardID)
|
||||
return
|
||||
}
|
||||
}
|
||||
if len(deltas) == 0 {
|
||||
if err := h.repo.UpsertNftCounterStates(newStates, nowMs); err != nil {
|
||||
log.Printf("nftables traffic collection failed op=upsert_states node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
batch := buildNftFlowUploadBatch(deltas, metas)
|
||||
if missingForwardID, ok := firstNftBatchMissingDelta(deltas, batch); ok {
|
||||
log.Printf("nftables traffic collection skipped state advance op=unaccounted_delta node_id=%d forward_id=%d", nodeID, missingForwardID)
|
||||
return
|
||||
}
|
||||
quotaViews, err := h.repo.ApplyNftTrafficAccounting(batch.flowDeltas, batch.quotaUsage, newStates, now)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=accounting node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
h.recordTunnelMetricsFromForwardBatch(nodeID, batch.forwardTraffic, metas, nowMs)
|
||||
for userID, quota := range quotaViews {
|
||||
h.enforceUserQuotaIfNeeded(userID, quota)
|
||||
}
|
||||
for _, target := range batch.policyTargets {
|
||||
if target.UserID <= 0 || target.UserTunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
h.enforceFlowPolicies(target.UserID, target.UserTunnelID)
|
||||
}
|
||||
}
|
||||
|
||||
func firstNftBatchMissingDelta(deltas []nftTrafficDelta, batch flowUploadBatch) (int64, bool) {
|
||||
flowSeen := make(map[int64]struct{}, len(batch.flowDeltas))
|
||||
for _, delta := range batch.flowDeltas {
|
||||
flowSeen[delta.ForwardID] = struct{}{}
|
||||
}
|
||||
|
||||
expectedRaw := make(map[int64]tunnelTrafficDelta, len(batch.forwardTraffic))
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID <= 0 || (delta.BytesIn == 0 && delta.BytesOut == 0) {
|
||||
continue
|
||||
}
|
||||
if delta.BytesIn < 0 || delta.BytesOut < 0 {
|
||||
return delta.ForwardID, true
|
||||
}
|
||||
raw := expectedRaw[delta.ForwardID]
|
||||
if raw.bytesIn > math.MaxInt64-delta.BytesIn || raw.bytesOut > math.MaxInt64-delta.BytesOut {
|
||||
return delta.ForwardID, true
|
||||
}
|
||||
raw.bytesIn += delta.BytesIn
|
||||
raw.bytesOut += delta.BytesOut
|
||||
expectedRaw[delta.ForwardID] = raw
|
||||
}
|
||||
|
||||
for forwardID, expected := range expectedRaw {
|
||||
actual, ok := batch.forwardTraffic[forwardID]
|
||||
if !ok || actual.bytesIn != expected.bytesIn || actual.bytesOut != expected.bytesOut {
|
||||
return forwardID, true
|
||||
}
|
||||
if expected.bytesIn != 0 || expected.bytesOut != 0 {
|
||||
if _, ok := flowSeen[forwardID]; !ok {
|
||||
return forwardID, true
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func firstNftDeltaMissingMeta(deltas []nftTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta) (int64, bool) {
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := metas[delta.ForwardID]; !ok {
|
||||
return delta.ForwardID, true
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func filterNftCounterSamplesWithBinding(samples []runtimenft.CounterSample, hashes map[int64]string) []runtimenft.CounterSample {
|
||||
if len(samples) == 0 || len(hashes) == 0 {
|
||||
return nil
|
||||
}
|
||||
filtered := make([]runtimenft.CounterSample, 0, len(samples))
|
||||
for _, sample := range samples {
|
||||
if _, ok := hashes[sample.ForwardID]; !ok {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, sample)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func nftCounterKey(forwardID int64, protocol, direction string) nftCounterStateKey {
|
||||
return nftCounterStateKey{
|
||||
forwardID: forwardID,
|
||||
protocol: strings.ToLower(strings.TrimSpace(protocol)),
|
||||
direction: strings.ToLower(strings.TrimSpace(direction)),
|
||||
}
|
||||
}
|
||||
|
||||
func buildNftCounterDeltas(nodeID int64, samples []runtimenft.CounterSample, oldStates []model.NftCounterState, hashes map[int64]string, nowMs int64) ([]nftTrafficDelta, []repo.NftCounterStateInput) {
|
||||
oldByKey := make(map[nftCounterStateKey]model.NftCounterState, len(oldStates))
|
||||
for _, old := range oldStates {
|
||||
if old.NodeID != nodeID {
|
||||
continue
|
||||
}
|
||||
oldByKey[nftCounterKey(old.ForwardID, old.Protocol, old.Direction)] = old
|
||||
}
|
||||
|
||||
stateInputs := make([]repo.NftCounterStateInput, 0, len(samples))
|
||||
deltaByForward := make(map[int64]nftTrafficDelta)
|
||||
for _, sample := range samples {
|
||||
direction := strings.ToLower(strings.TrimSpace(sample.Direction))
|
||||
if direction != runtimenft.CounterDirectionToTarget && direction != runtimenft.CounterDirectionFromTarget {
|
||||
continue
|
||||
}
|
||||
|
||||
protocol := strings.ToLower(strings.TrimSpace(sample.Protocol))
|
||||
if protocol != "tcp" && protocol != "udp" {
|
||||
continue
|
||||
}
|
||||
if sample.Bytes > uint64(math.MaxInt64) || sample.Packets > uint64(math.MaxInt64) {
|
||||
continue
|
||||
}
|
||||
ruleHash := strings.TrimSpace(hashes[sample.ForwardID])
|
||||
stateInput := repo.NftCounterStateInput{
|
||||
NodeID: nodeID,
|
||||
ForwardID: sample.ForwardID,
|
||||
Protocol: protocol,
|
||||
Direction: direction,
|
||||
RuleHash: ruleHash,
|
||||
Bytes: sample.Bytes,
|
||||
Packets: sample.Packets,
|
||||
CollectedTime: nowMs,
|
||||
}
|
||||
|
||||
old, exists := oldByKey[nftCounterKey(sample.ForwardID, protocol, direction)]
|
||||
if !exists || old.RuleHash != ruleHash {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
if old.Bytes < 0 {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
oldBytes := uint64(old.Bytes)
|
||||
if sample.Bytes < oldBytes {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
rawDelta := sample.Bytes - oldBytes
|
||||
if rawDelta == 0 {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
|
||||
delta := deltaByForward[sample.ForwardID]
|
||||
delta.ForwardID = sample.ForwardID
|
||||
rawDeltaInt := int64(rawDelta)
|
||||
if direction == runtimenft.CounterDirectionToTarget {
|
||||
if delta.BytesIn > math.MaxInt64-rawDeltaInt {
|
||||
continue
|
||||
}
|
||||
delta.BytesIn += rawDeltaInt
|
||||
} else {
|
||||
if delta.BytesOut > math.MaxInt64-rawDeltaInt {
|
||||
continue
|
||||
}
|
||||
delta.BytesOut += rawDeltaInt
|
||||
}
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
deltaByForward[sample.ForwardID] = delta
|
||||
}
|
||||
|
||||
forwardIDs := make([]int64, 0, len(deltaByForward))
|
||||
for forwardID := range deltaByForward {
|
||||
forwardIDs = append(forwardIDs, forwardID)
|
||||
}
|
||||
sort.Slice(forwardIDs, func(i, j int) bool { return forwardIDs[i] < forwardIDs[j] })
|
||||
|
||||
deltas := make([]nftTrafficDelta, 0, len(forwardIDs))
|
||||
for _, forwardID := range forwardIDs {
|
||||
delta := deltaByForward[forwardID]
|
||||
if delta.BytesIn == 0 && delta.BytesOut == 0 {
|
||||
continue
|
||||
}
|
||||
deltas = append(deltas, delta)
|
||||
}
|
||||
return deltas, stateInputs
|
||||
}
|
||||
|
||||
func buildNftFlowUploadBatch(deltas []nftTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta) flowUploadBatch {
|
||||
batch := flowUploadBatch{
|
||||
quotaUsage: make(map[int64]int64),
|
||||
forwardTraffic: make(map[int64]tunnelTrafficDelta),
|
||||
orphanServices: make(map[string]struct{}),
|
||||
peerShareForwardItems: make(map[string]flowItem),
|
||||
peerShareRuntimeItems: make(map[int64]flowItem),
|
||||
}
|
||||
policySeen := map[flowPolicyTarget]struct{}{}
|
||||
flowSeen := map[int64]int{}
|
||||
|
||||
for _, delta := range deltas {
|
||||
meta, exists := metas[delta.ForwardID]
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
|
||||
raw := batch.forwardTraffic[delta.ForwardID]
|
||||
if delta.BytesIn < 0 || delta.BytesOut < 0 || raw.bytesIn > math.MaxInt64-delta.BytesIn || raw.bytesOut > math.MaxInt64-delta.BytesOut {
|
||||
continue
|
||||
}
|
||||
|
||||
scaledIn, ok := scaleNftTrafficBytes(delta.BytesIn, meta.TrafficRatio, meta.TunnelFlow)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
scaledOut, ok := scaleNftTrafficBytes(delta.BytesOut, meta.TrafficRatio, meta.TunnelFlow)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if scaledIn > math.MaxInt64-scaledOut {
|
||||
continue
|
||||
}
|
||||
quotaDelta := scaledIn + scaledOut
|
||||
if batch.quotaUsage[meta.UserID] > math.MaxInt64-quotaDelta {
|
||||
continue
|
||||
}
|
||||
|
||||
flowIdx, flowExists := flowSeen[delta.ForwardID]
|
||||
if flowExists && (batch.flowDeltas[flowIdx].InFlow > math.MaxInt64-scaledIn || batch.flowDeltas[flowIdx].OutFlow > math.MaxInt64-scaledOut) {
|
||||
continue
|
||||
}
|
||||
|
||||
raw.bytesIn += delta.BytesIn
|
||||
raw.bytesOut += delta.BytesOut
|
||||
batch.forwardTraffic[delta.ForwardID] = raw
|
||||
|
||||
if flowExists {
|
||||
batch.flowDeltas[flowIdx].InFlow += scaledIn
|
||||
batch.flowDeltas[flowIdx].OutFlow += scaledOut
|
||||
} else {
|
||||
flowSeen[delta.ForwardID] = len(batch.flowDeltas)
|
||||
batch.flowDeltas = append(batch.flowDeltas, repo.FlowUploadCounterDelta{
|
||||
ForwardID: delta.ForwardID,
|
||||
UserID: meta.UserID,
|
||||
UserTunnelID: meta.UserTunnelID,
|
||||
InFlow: scaledIn,
|
||||
OutFlow: scaledOut,
|
||||
})
|
||||
}
|
||||
batch.quotaUsage[meta.UserID] += quotaDelta
|
||||
|
||||
target := flowPolicyTarget{UserID: meta.UserID, UserTunnelID: meta.UserTunnelID}
|
||||
if _, seen := policySeen[target]; !seen {
|
||||
policySeen[target] = struct{}{}
|
||||
batch.policyTargets = append(batch.policyTargets, target)
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(batch.policyTargets, func(i, j int) bool {
|
||||
if batch.policyTargets[i].UserID == batch.policyTargets[j].UserID {
|
||||
return batch.policyTargets[i].UserTunnelID < batch.policyTargets[j].UserTunnelID
|
||||
}
|
||||
return batch.policyTargets[i].UserID < batch.policyTargets[j].UserID
|
||||
})
|
||||
|
||||
return batch
|
||||
}
|
||||
|
||||
func scaleNftTrafficBytes(bytes int64, ratio float64, tunnelFlow int64) (int64, bool) {
|
||||
if bytes < 0 || ratio < 0 || tunnelFlow < 0 {
|
||||
return 0, false
|
||||
}
|
||||
var scaled int64
|
||||
if ratio == 1 {
|
||||
scaled = bytes
|
||||
} else {
|
||||
scaledFloat := float64(bytes) * ratio
|
||||
if math.IsNaN(scaledFloat) || math.IsInf(scaledFloat, 0) || scaledFloat < 0 || scaledFloat >= math.Pow(2, 63) {
|
||||
return 0, false
|
||||
}
|
||||
scaled = int64(scaledFloat)
|
||||
}
|
||||
if tunnelFlow != 0 && scaled > math.MaxInt64/tunnelFlow {
|
||||
return 0, false
|
||||
}
|
||||
return scaled * tunnelFlow, true
|
||||
}
|
||||
@@ -0,0 +1,680 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildNftCounterDeltasSavesFirstBaselineWithoutDelta(t *testing.T) {
|
||||
nowMs := int64(1700000000123)
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}, nil, map[int64]string{42: "hash-a"}, nowMs)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected no deltas for first baseline, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one state input, got %d", len(states))
|
||||
}
|
||||
state := states[0]
|
||||
if state.NodeID != 11 || state.ForwardID != 42 || state.Protocol != "tcp" || state.Direction != runtimenft.CounterDirectionToTarget {
|
||||
t.Fatalf("unexpected state identity: %#v", state)
|
||||
}
|
||||
if state.RuleHash != "hash-a" || state.Bytes != 1000 || state.Packets != 10 || state.CollectedTime != nowMs {
|
||||
t.Fatalf("unexpected state values: %#v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasNormalGrowthProducesDirectionalBytes(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: 15},
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2600, Packets: 26},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, RuleHash: "hash-a", Bytes: 2000, Packets: 20},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 1 {
|
||||
t.Fatalf("expected one aggregated delta, got %#v", deltas)
|
||||
}
|
||||
if deltas[0].ForwardID != 42 || deltas[0].BytesIn != 500 || deltas[0].BytesOut != 600 {
|
||||
t.Fatalf("unexpected delta: %#v", deltas[0])
|
||||
}
|
||||
if len(states) != 2 {
|
||||
t.Fatalf("expected two state inputs, got %d", len(states))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasResetRefreshesBaselineWithoutDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 25, Packets: 2},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected reset to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 || states[0].Bytes != 25 || states[0].RuleHash != "hash-a" {
|
||||
t.Fatalf("expected refreshed baseline state, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasRuleHashChangeRefreshesBaselineWithoutDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: 15},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-b"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected rule hash change to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 || states[0].Bytes != 1500 || states[0].RuleHash != "hash-b" {
|
||||
t.Fatalf("expected refreshed hash baseline state, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasEqualBytesRefreshesBaselineWithoutDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 11},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected equal bytes to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 || states[0].Bytes != 1000 || states[0].Packets != 11 || states[0].RuleHash != "hash-a" {
|
||||
t.Fatalf("expected refreshed baseline state, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasAggregatesProtocolsAndDirections(t *testing.T) {
|
||||
deltas, _ := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1100, Packets: 11},
|
||||
{ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 2200, Packets: 22},
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 3300, Packets: 33},
|
||||
{ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 4400, Packets: 44},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 2000},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, RuleHash: "hash-a", Bytes: 3000},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionFromTarget, RuleHash: "hash-a", Bytes: 4000},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 1 {
|
||||
t.Fatalf("expected one aggregated delta, got %#v", deltas)
|
||||
}
|
||||
if deltas[0].ForwardID != 42 || deltas[0].BytesIn != 300 || deltas[0].BytesOut != 700 {
|
||||
t.Fatalf("unexpected aggregated delta: %#v", deltas[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsInvalidProtocolBeforeStateAndDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "icmp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: 15},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "icmp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected invalid protocol to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected invalid protocol to produce no state inputs, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsOversizedPacketsBeforeStateAndDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: uint64(math.MaxInt64) + 1},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected oversized packets to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected oversized packets to produce no state inputs, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsOversizedBytesBeforeStateAndDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64) + 1, Packets: 10},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected oversized bytes to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected oversized bytes to produce no state inputs, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsOverflowingAggregateSampleWithoutState(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64), Packets: 10},
|
||||
{ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 10, Packets: 1},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1, Packets: 1},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1, Packets: 1},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 1 {
|
||||
t.Fatalf("expected only non-overflowing aggregate delta, got %#v", deltas)
|
||||
}
|
||||
if deltas[0].ForwardID != 42 || deltas[0].BytesIn != math.MaxInt64-1 || deltas[0].BytesOut != 0 {
|
||||
t.Fatalf("unexpected aggregate delta: %#v", deltas[0])
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected only the accounted safe sample to advance baseline, got %#v", states)
|
||||
}
|
||||
if states[0].ForwardID != 42 || states[0].Protocol != "tcp" || states[0].Bytes != uint64(math.MaxInt64) {
|
||||
t.Fatalf("expected safe sample state input to be preserved, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsUnknownDirectionAndOversizedDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: "sideways", Bytes: 1500, Packets: 15},
|
||||
{ForwardID: 43, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64) + 1, Packets: 1},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-b", Bytes: 100},
|
||||
}, map[int64]string{42: "hash-a", 43: "hash-b"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected no delta for skipped/oversized samples, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state inputs for skipped/oversized samples, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchScalesFlowAndPreservesRawTunnelTraffic(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: 80, BytesOut: 110},
|
||||
{ForwardID: 21, BytesIn: 7, BytesOut: 11},
|
||||
{ForwardID: 20, BytesIn: 20, BytesOut: 10},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TunnelID: 1, TrafficRatio: 2, TunnelFlow: 3},
|
||||
21: {ForwardID: 21, UserID: 2, UserTunnelID: 10, TunnelID: 1, TrafficRatio: 1.5, TunnelFlow: 2},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 2 {
|
||||
t.Fatalf("expected two flow deltas, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if batch.flowDeltas[0].ForwardID != 20 || batch.flowDeltas[0].InFlow != 600 || batch.flowDeltas[0].OutFlow != 720 {
|
||||
t.Fatalf("unexpected first flow delta: %#v", batch.flowDeltas[0])
|
||||
}
|
||||
if batch.flowDeltas[1].ForwardID != 21 || batch.flowDeltas[1].InFlow != 20 || batch.flowDeltas[1].OutFlow != 32 {
|
||||
t.Fatalf("unexpected second flow delta: %#v", batch.flowDeltas[1])
|
||||
}
|
||||
if batch.quotaUsage[2] != 1372 {
|
||||
t.Fatalf("expected quota usage 1372, got %d", batch.quotaUsage[2])
|
||||
}
|
||||
if len(batch.policyTargets) != 1 || batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("expected deduped policy target, got %#v", batch.policyTargets)
|
||||
}
|
||||
if traffic := batch.forwardTraffic[20]; traffic.bytesIn != 100 || traffic.bytesOut != 120 {
|
||||
t.Fatalf("expected raw traffic for forward 20, got %#v", traffic)
|
||||
}
|
||||
if traffic := batch.forwardTraffic[21]; traffic.bytesIn != 7 || traffic.bytesOut != 11 {
|
||||
t.Fatalf("expected raw traffic for forward 21, got %#v", traffic)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsOverflowingScaledFlow(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 2, TunnelFlow: 2},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 0 {
|
||||
t.Fatalf("expected overflowing scaled flow to be skipped, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if len(batch.quotaUsage) != 0 {
|
||||
t.Fatalf("expected no quota usage for overflowing scaled flow, got %#v", batch.quotaUsage)
|
||||
}
|
||||
if len(batch.policyTargets) != 0 {
|
||||
t.Fatalf("expected no policy targets for overflowing scaled flow, got %#v", batch.policyTargets)
|
||||
}
|
||||
if len(batch.forwardTraffic) != 0 {
|
||||
t.Fatalf("expected no raw traffic for overflowing scaled flow, got %#v", batch.forwardTraffic)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsRawForwardTrafficOverflow(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
{ForwardID: 20, BytesIn: 1, BytesOut: 0},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 0.5, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
traffic := batch.forwardTraffic[20]
|
||||
if traffic.bytesIn != math.MaxInt64 || traffic.bytesOut != 0 {
|
||||
t.Fatalf("expected overflowing raw delta to be skipped without negative traffic, got %#v", traffic)
|
||||
}
|
||||
if len(batch.flowDeltas) != 1 || batch.flowDeltas[0].ForwardID != 20 {
|
||||
t.Fatalf("expected only the safe flow delta, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if len(batch.policyTargets) != 1 || batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("expected policy target only from safe delta, got %#v", batch.policyTargets)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsQuotaOverflow(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
{ForwardID: 21, BytesIn: 1, BytesOut: 0},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 1, TunnelFlow: 1},
|
||||
21: {ForwardID: 21, UserID: 2, UserTunnelID: 10, TrafficRatio: 1, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 1 || batch.flowDeltas[0].ForwardID != 20 || batch.flowDeltas[0].InFlow != math.MaxInt64 {
|
||||
t.Fatalf("expected only non-overflowing quota delta, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if batch.quotaUsage[2] != math.MaxInt64 {
|
||||
t.Fatalf("expected quota usage to remain at max int64, got %#v", batch.quotaUsage)
|
||||
}
|
||||
if len(batch.policyTargets) != 1 || batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("expected one policy target from non-overflowing delta, got %#v", batch.policyTargets)
|
||||
}
|
||||
if _, ok := batch.forwardTraffic[21]; ok {
|
||||
t.Fatalf("expected quota-overflowing delta to be skipped from raw traffic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsMissingMeta(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: 80, BytesOut: 110},
|
||||
{ForwardID: 99, BytesIn: 1, BytesOut: 2},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 1, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 1 || batch.flowDeltas[0].ForwardID != 20 {
|
||||
t.Fatalf("expected only forward 20 delta, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if _, ok := batch.forwardTraffic[99]; ok {
|
||||
t.Fatalf("expected missing meta forward to be skipped from raw traffic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftBatchCoversDeltasRequiresRawAndFlowEntries(t *testing.T) {
|
||||
deltas := []nftTrafficDelta{{ForwardID: 20, BytesIn: 1, BytesOut: 0}}
|
||||
batch := flowUploadBatch{
|
||||
forwardTraffic: map[int64]tunnelTrafficDelta{20: {bytesIn: 1}},
|
||||
flowDeltas: []repo.FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 1}},
|
||||
}
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); ok || missing != 0 {
|
||||
t.Fatalf("expected batch to cover delta, missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
|
||||
delete(batch.forwardTraffic, 20)
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
|
||||
t.Fatalf("expected missing raw traffic for forward 20, got missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
|
||||
batch.forwardTraffic[20] = tunnelTrafficDelta{bytesIn: 1}
|
||||
batch.flowDeltas = nil
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
|
||||
t.Fatalf("expected missing flow delta for forward 20, got missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftBatchCoversDeltasRequiresAggregateRawTotals(t *testing.T) {
|
||||
deltas := []nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
{ForwardID: 20, BytesIn: 1, BytesOut: 0},
|
||||
}
|
||||
batch := buildNftFlowUploadBatch(deltas, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 0.5, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
|
||||
t.Fatalf("expected aggregate raw overflow/mismatch for forward 20, got missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficFirstBaselineSavesStateWithoutFlow(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2000, Packets: 20},
|
||||
}}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
if manager.collectHit != 1 {
|
||||
t.Fatalf("expected one collection, got %d", manager.collectHit)
|
||||
}
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 2 {
|
||||
t.Fatalf("expected two baseline states, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no forward flow on baseline, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT out_flow FROM user WHERE id = 1`); got != 0 {
|
||||
t.Fatalf("expected no user flow on baseline, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficGrowthAppliesFlowAndUpdatesState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2000, Packets: 20},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2600, Packets: 26},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 400 {
|
||||
t.Fatalf("expected forward in_flow=400, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT out_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 600 {
|
||||
t.Fatalf("expected forward out_flow=600, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM user WHERE id = 1`); got != 400 {
|
||||
t.Fatalf("expected user in_flow=400, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT out_flow FROM user_tunnel WHERE id = ?`, fixture.userTunnelID); got != 600 {
|
||||
t.Fatalf("expected user_tunnel out_flow=600, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT COALESCE((SELECT daily_used_bytes FROM user_quota WHERE user_id = 1), 0)`); got != 1000 {
|
||||
t.Fatalf("expected daily quota usage=1000, got %d", got)
|
||||
}
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 2 {
|
||||
t.Fatalf("expected two states after growth, got %+v", states)
|
||||
}
|
||||
for _, state := range states {
|
||||
if state.Direction == runtimenft.CounterDirectionToTarget && state.Bytes != 1400 {
|
||||
t.Fatalf("expected to-target state bytes 1400, got %+v", state)
|
||||
}
|
||||
if state.Direction == runtimenft.CounterDirectionFromTarget && state.Bytes != 2600 {
|
||||
t.Fatalf("expected from-target state bytes 2600, got %+v", state)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficSkippedBatchDeltaDoesNotAdvanceState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
if err := h.repo.DB().Exec(`UPDATE tunnel SET traffic_ratio = 2 WHERE id = (SELECT tunnel_id FROM forward WHERE id = ?)`, fixture.forwardID).Error; err != nil {
|
||||
t.Fatalf("update tunnel ratio: %v", err)
|
||||
}
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 0, Packets: 0},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64), Packets: 1},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one state, got %+v", states)
|
||||
}
|
||||
if states[0].Bytes != 0 || states[0].Packets != 0 {
|
||||
t.Fatalf("expected state to remain at old baseline after skipped batch delta, got %+v", states[0])
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no forward flow for skipped batch delta, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficMetadataErrorDoesNotAdvanceState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
if err := h.repo.DB().Exec(`DROP TABLE tunnel`).Error; err != nil {
|
||||
t.Fatalf("drop tunnel table: %v", err)
|
||||
}
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one baseline state, got %+v", states)
|
||||
}
|
||||
if states[0].Bytes != 1000 || states[0].Packets != 10 {
|
||||
t.Fatalf("expected state to remain at first baseline after metadata failure, got %+v", states[0])
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow after metadata failure, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficMissingMetaDoesNotAdvanceState(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
forwardID := int64(4242)
|
||||
nowMs := time.Now().UnixMilli()
|
||||
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: forwardID,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 20000,
|
||||
Protocols: "tcp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
RuleHash: "hash-a",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, nowMs); err != nil {
|
||||
t.Fatalf("seed stale applied binding: %v", err)
|
||||
}
|
||||
if err := h.repo.UpsertNftCounterStates([]repo.NftCounterStateInput{{
|
||||
NodeID: fixture.nodeID,
|
||||
ForwardID: forwardID,
|
||||
Protocol: "tcp",
|
||||
Direction: runtimenft.CounterDirectionToTarget,
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1000,
|
||||
Packets: 10,
|
||||
CollectedTime: nowMs,
|
||||
}}, nowMs); err != nil {
|
||||
t.Fatalf("seed counter state: %v", err)
|
||||
}
|
||||
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
|
||||
}}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one state, got %+v", states)
|
||||
}
|
||||
if states[0].Bytes != 1000 || states[0].Packets != 10 {
|
||||
t.Fatalf("expected state to remain at old baseline when meta is missing, got %+v", states[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficSkipsSamplesWithoutBinding(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
if err := h.repo.DeleteNftRuleBindingsByForward(fixture.forwardID); err != nil {
|
||||
t.Fatalf("delete nft binding: %v", err)
|
||||
}
|
||||
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state for unbound sample, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow for unbound sample, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficSkipsNonAppliedBinding(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
if err := h.repo.MarkNftRuleBindingError(fixture.forwardID, fixture.nodeID, "apply failed", time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("mark binding error: %v", err)
|
||||
}
|
||||
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state for non-applied binding, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow for non-applied binding, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficCollectionErrorDoesNotWriteState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
h.nftablesManager = &fakeNftablesManager{collectErr: errors.New("ssh failed")}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state on collection error, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow on collection error, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
type nftablesCollectionFixture struct {
|
||||
handler *Handler
|
||||
nodeID int64
|
||||
forwardID int64
|
||||
userTunnelID int64
|
||||
}
|
||||
|
||||
func setupNftablesCollectionFixture(t *testing.T) nftablesCollectionFixture {
|
||||
t.Helper()
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-traffic-tunnel", fixture.nodeID)
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(1, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelID).Error; err != nil {
|
||||
t.Fatalf("seed user_tunnel: %v", err)
|
||||
}
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: forward.ID,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 20000,
|
||||
Protocols: "tcp,udp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
RuleHash: "hash-a",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, now); err != nil {
|
||||
t.Fatalf("seed nft binding: %v", err)
|
||||
}
|
||||
userTunnelID := mustHandlerCount(t, h, `SELECT id FROM user_tunnel WHERE user_id = 1 AND tunnel_id = ?`, tunnelID)
|
||||
return nftablesCollectionFixture{
|
||||
handler: h,
|
||||
nodeID: fixture.nodeID,
|
||||
forwardID: forward.ID,
|
||||
userTunnelID: userTunnelID,
|
||||
}
|
||||
}
|
||||
|
||||
func mustCollectionSSHConfig(t *testing.T, h *Handler, nodeID int64) *model.NodeSSHConfig {
|
||||
t.Helper()
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load ssh config: %v", err)
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
|
||||
func mustHandlerCount(t *testing.T, h *Handler, query string, args ...interface{}) int64 {
|
||||
t.Helper()
|
||||
var value int64
|
||||
if err := h.repo.DB().Raw(query, args...).Row().Scan(&value); err != nil {
|
||||
t.Fatalf("query %q failed: %v", query, err)
|
||||
}
|
||||
return value
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type CounterSample struct {
|
||||
ForwardID int64
|
||||
Direction string
|
||||
Protocol string
|
||||
Bytes uint64
|
||||
Packets uint64
|
||||
}
|
||||
|
||||
func ParseCounterComment(comment string) (CounterSample, bool) {
|
||||
parts := strings.Split(comment, " ")
|
||||
if len(parts) != 4 || parts[0] != "flvx" {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
if !strings.HasPrefix(parts[1], "forward:") {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
forwardText := strings.TrimPrefix(parts[1], "forward:")
|
||||
forwardID, err := strconv.ParseInt(forwardText, 10, 64)
|
||||
if err != nil || forwardID <= 0 {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
direction := parts[2]
|
||||
if direction != CounterDirectionToTarget && direction != CounterDirectionFromTarget {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
protocol := parts[3]
|
||||
if protocol != "tcp" && protocol != "udp" {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
return CounterSample{
|
||||
ForwardID: forwardID,
|
||||
Direction: direction,
|
||||
Protocol: protocol,
|
||||
}, true
|
||||
}
|
||||
|
||||
func ParseCounterSamples(raw []byte) ([]CounterSample, error) {
|
||||
var doc nftListTable
|
||||
if err := json.Unmarshal(raw, &doc); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
samples := make([]CounterSample, 0)
|
||||
for _, item := range doc.Nftables {
|
||||
ruleRaw, ok := item["rule"]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
var rule nftCounterRule
|
||||
if err := json.Unmarshal(ruleRaw, &rule); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rule.Table != "flvx" || rule.Chain != "forward" {
|
||||
continue
|
||||
}
|
||||
|
||||
sample, ok, err := parseCounterRule(rule)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
samples = append(samples, sample)
|
||||
}
|
||||
return samples, nil
|
||||
}
|
||||
|
||||
type nftListTable struct {
|
||||
Nftables []map[string]json.RawMessage `json:"nftables"`
|
||||
}
|
||||
|
||||
type nftCounterRule struct {
|
||||
Table string `json:"table"`
|
||||
Chain string `json:"chain"`
|
||||
Comment string `json:"comment"`
|
||||
Expr []map[string]json.RawMessage `json:"expr"`
|
||||
}
|
||||
|
||||
type nftCounter struct {
|
||||
Bytes uint64 `json:"bytes"`
|
||||
Packets uint64 `json:"packets"`
|
||||
}
|
||||
|
||||
func parseCounterRule(rule nftCounterRule) (CounterSample, bool, error) {
|
||||
var (
|
||||
counter nftCounter
|
||||
hasCounter bool
|
||||
comment = rule.Comment
|
||||
)
|
||||
|
||||
for _, expr := range rule.Expr {
|
||||
if rawCounter, ok := expr["counter"]; ok {
|
||||
if err := json.Unmarshal(rawCounter, &counter); err != nil {
|
||||
return CounterSample{}, false, err
|
||||
}
|
||||
hasCounter = true
|
||||
continue
|
||||
}
|
||||
if rawComment, ok := expr["comment"]; ok && strings.TrimSpace(comment) == "" {
|
||||
if err := json.Unmarshal(rawComment, &comment); err != nil {
|
||||
return CounterSample{}, false, err
|
||||
}
|
||||
}
|
||||
}
|
||||
if !hasCounter {
|
||||
return CounterSample{}, false, nil
|
||||
}
|
||||
|
||||
sample, ok := ParseCounterComment(comment)
|
||||
if !ok {
|
||||
return CounterSample{}, false, nil
|
||||
}
|
||||
sample.Bytes = counter.Bytes
|
||||
sample.Packets = counter.Packets
|
||||
return sample, true, nil
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
package nftables
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseCounterCommentAcceptsValidToTargetTCP(t *testing.T) {
|
||||
sample, ok := ParseCounterComment("flvx forward:42 to-target tcp")
|
||||
if !ok {
|
||||
t.Fatal("expected comment to parse")
|
||||
}
|
||||
if sample.ForwardID != 42 ||
|
||||
sample.Direction != CounterDirectionToTarget ||
|
||||
sample.Protocol != "tcp" {
|
||||
t.Fatalf("unexpected sample: %+v", sample)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterCommentRejectsDNAT(t *testing.T) {
|
||||
if sample, ok := ParseCounterComment("flvx forward:42 dnat tcp"); ok {
|
||||
t.Fatalf("expected dnat comment to be rejected, got %+v", sample)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesParsesForwardBillableCounters(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"metainfo": {"json_schema_version": 1}},
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"handle": 10,
|
||||
"comment": "flvx forward:42 to-target tcp",
|
||||
"expr": [
|
||||
{"match": {"left": {"payload": {"protocol": "ip", "field": "daddr"}}, "op": "==", "right": "198.51.100.20"}},
|
||||
{"counter": {"packets": 7, "bytes": 4096}}
|
||||
]
|
||||
}},
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"handle": 11,
|
||||
"comment": "flvx forward:42 from-target udp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 9, "bytes": 8192}}
|
||||
]
|
||||
}},
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "prerouting",
|
||||
"handle": 12,
|
||||
"comment": "flvx forward:42 dnat tcp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 100, "bytes": 65536}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
|
||||
samples, err := ParseCounterSamples(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseCounterSamples: %v", err)
|
||||
}
|
||||
if len(samples) != 2 {
|
||||
t.Fatalf("expected 2 samples, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
|
||||
want := []CounterSample{
|
||||
{ForwardID: 42, Direction: CounterDirectionToTarget, Protocol: "tcp", Bytes: 4096, Packets: 7},
|
||||
{ForwardID: 42, Direction: CounterDirectionFromTarget, Protocol: "udp", Bytes: 8192, Packets: 9},
|
||||
}
|
||||
for i := range want {
|
||||
if samples[i] != want[i] {
|
||||
t.Fatalf("sample %d: expected %+v, got %+v", i, want[i], samples[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesUsesRuleLevelComment(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"comment": "flvx forward:77 to-target udp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 3, "bytes": 2048}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
|
||||
samples, err := ParseCounterSamples(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseCounterSamples: %v", err)
|
||||
}
|
||||
if len(samples) != 1 {
|
||||
t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
want := CounterSample{
|
||||
ForwardID: 77,
|
||||
Direction: CounterDirectionToTarget,
|
||||
Protocol: "udp",
|
||||
Bytes: 2048,
|
||||
Packets: 3,
|
||||
}
|
||||
if samples[0] != want {
|
||||
t.Fatalf("expected %+v, got %+v", want, samples[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesUsesExprLevelComment(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"expr": [
|
||||
{"counter": {"packets": 4, "bytes": 3072}},
|
||||
{"comment": "flvx forward:78 from-target tcp"}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
|
||||
samples, err := ParseCounterSamples(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseCounterSamples: %v", err)
|
||||
}
|
||||
if len(samples) != 1 {
|
||||
t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
want := CounterSample{
|
||||
ForwardID: 78,
|
||||
Direction: CounterDirectionFromTarget,
|
||||
Protocol: "tcp",
|
||||
Bytes: 3072,
|
||||
Packets: 4,
|
||||
}
|
||||
if samples[0] != want {
|
||||
t.Fatalf("expected %+v, got %+v", want, samples[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesMalformedJSONReturnsError(t *testing.T) {
|
||||
if _, err := ParseCounterSamples([]byte(`{"nftables": [`)); err == nil {
|
||||
t.Fatal("expected malformed JSON error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesMalformedRuleJSONReturnsError(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"comment": "flvx forward:42 to-target tcp",
|
||||
"expr": [
|
||||
{"counter": {"packets": "bad", "bytes": 4096}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
if _, err := ParseCounterSamples(raw); err == nil {
|
||||
t.Fatal("expected malformed rule JSON error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
type Manager struct {
|
||||
runner Runner
|
||||
}
|
||||
|
||||
func NewManager(runner Runner) *Manager {
|
||||
if runner == nil {
|
||||
runner = NewSSHRunner()
|
||||
}
|
||||
return &Manager{runner: runner}
|
||||
}
|
||||
|
||||
func (m *Manager) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return err
|
||||
}
|
||||
return m.runner.Test(ctx, cfg)
|
||||
}
|
||||
|
||||
func (m *Manager) Reconcile(ctx context.Context, cfg SSHConfig, plan NodePlan) (ApplyResult, error) {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return ApplyResult{}, err
|
||||
}
|
||||
result := ApplyResult{
|
||||
NodeID: plan.NodeID,
|
||||
Script: RenderTable(plan),
|
||||
Hashes: PlanHashes(plan),
|
||||
}
|
||||
if err := m.runner.ApplyScript(ctx, cfg, result.Script); err != nil {
|
||||
return ApplyResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (m *Manager) Clear(ctx context.Context, cfg SSHConfig) error {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return err
|
||||
}
|
||||
script := RenderTable(NodePlan{})
|
||||
return m.runner.ApplyScript(ctx, cfg, script)
|
||||
}
|
||||
|
||||
func (m *Manager) CollectCounters(ctx context.Context, cfg SSHConfig) ([]CounterSample, error) {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
raw, err := m.runner.ListTableJSON(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return ParseCounterSamples(raw)
|
||||
}
|
||||
|
||||
func (m *Manager) ensureInitialized() error {
|
||||
if m == nil || m.runner == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,162 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type fakeRunner struct {
|
||||
scripts []string
|
||||
err error
|
||||
testErr error
|
||||
listJSON []byte
|
||||
listJSONErr error
|
||||
}
|
||||
|
||||
func (f *fakeRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
f.scripts = append(f.scripts, script)
|
||||
return f.err
|
||||
}
|
||||
|
||||
func (f *fakeRunner) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
|
||||
return f.listJSON, f.listJSONErr
|
||||
}
|
||||
|
||||
func TestManagerReconcileAppliesRenderedScript(t *testing.T) {
|
||||
runner := &fakeRunner{}
|
||||
manager := NewManager(runner)
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}},
|
||||
}
|
||||
|
||||
result, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, plan)
|
||||
if err != nil {
|
||||
t.Fatalf("Reconcile: %v", err)
|
||||
}
|
||||
if len(runner.scripts) != 1 {
|
||||
t.Fatalf("expected 1 script, got %d", len(runner.scripts))
|
||||
}
|
||||
if !strings.Contains(runner.scripts[0], `flvx forward:42 dnat tcp`) {
|
||||
t.Fatalf("script missing forward comment:\n%s", runner.scripts[0])
|
||||
}
|
||||
if result.NodeID != 7 || result.Hashes[42] == "" {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerReconcileReturnsRunnerError(t *testing.T) {
|
||||
runner := &fakeRunner{err: errors.New("ssh failed")}
|
||||
manager := NewManager(runner)
|
||||
|
||||
_, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, NodePlan{NodeID: 7})
|
||||
if !errors.Is(err, runner.err) {
|
||||
t.Fatalf("expected original runner error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerClearAppliesEmptyTable(t *testing.T) {
|
||||
runner := &fakeRunner{}
|
||||
manager := NewManager(runner)
|
||||
|
||||
if err := manager.Clear(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}); err != nil {
|
||||
t.Fatalf("Clear: %v", err)
|
||||
}
|
||||
if len(runner.scripts) != 1 {
|
||||
t.Fatalf("expected 1 script, got %d", len(runner.scripts))
|
||||
}
|
||||
if strings.Contains(runner.scripts[0], "masquerade comment") {
|
||||
t.Fatalf("empty table should not include masquerade:\n%s", runner.scripts[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerTestPassesThroughRunnerError(t *testing.T) {
|
||||
runner := &fakeRunner{testErr: errors.New("probe failed")}
|
||||
manager := NewManager(runner)
|
||||
|
||||
err := manager.Test(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"})
|
||||
if !errors.Is(err, runner.testErr) {
|
||||
t.Fatalf("expected original runner error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCollectCountersParsesRunnerTableJSON(t *testing.T) {
|
||||
runner := &fakeRunner{listJSON: []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"comment": "flvx forward:77 to-target tcp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 3, "bytes": 2048}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)}
|
||||
manager := NewManager(runner)
|
||||
|
||||
samples, err := manager.CollectCounters(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"})
|
||||
if err != nil {
|
||||
t.Fatalf("CollectCounters: %v", err)
|
||||
}
|
||||
if len(samples) != 1 {
|
||||
t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
want := CounterSample{
|
||||
ForwardID: 77,
|
||||
Direction: CounterDirectionToTarget,
|
||||
Protocol: "tcp",
|
||||
Bytes: 2048,
|
||||
Packets: 3,
|
||||
}
|
||||
if samples[0] != want {
|
||||
t.Fatalf("expected %+v, got %+v", want, samples[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
|
||||
cfg := SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}
|
||||
plan := NodePlan{NodeID: 7}
|
||||
expected := errors.New("nftables manager not initialized")
|
||||
|
||||
var nilManager *Manager
|
||||
if err := nilManager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager Test, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := nilManager.Reconcile(context.Background(), cfg, plan); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager Reconcile, got %v", err)
|
||||
}
|
||||
|
||||
if err := nilManager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager Clear, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := nilManager.CollectCounters(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager CollectCounters, got %v", err)
|
||||
}
|
||||
|
||||
manager := &Manager{}
|
||||
if err := manager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from Test, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := manager.Reconcile(context.Background(), cfg, plan); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from Reconcile, got %v", err)
|
||||
}
|
||||
|
||||
if err := manager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from Clear, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := manager.CollectCounters(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from CollectCounters, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func ParseSingleTarget(raw string) (Target, error) {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" {
|
||||
return Target{}, fmt.Errorf("目标地址不能为空")
|
||||
}
|
||||
if strings.Contains(value, ",") || strings.Contains(value, "\n") {
|
||||
return Target{}, fmt.Errorf("nftables 纯转发第一阶段仅支持单目标")
|
||||
}
|
||||
if hasScheme(value) {
|
||||
return Target{}, fmt.Errorf("目标地址必须是 host:port,不能包含 URL scheme")
|
||||
}
|
||||
host, portText, err := net.SplitHostPort(value)
|
||||
if err != nil {
|
||||
return Target{}, fmt.Errorf("目标地址必须是 host:port")
|
||||
}
|
||||
host = strings.TrimSpace(strings.Trim(host, "[]"))
|
||||
if host == "" {
|
||||
return Target{}, fmt.Errorf("目标主机不能为空")
|
||||
}
|
||||
port, err := strconv.Atoi(portText)
|
||||
if err != nil || port < 1 || port > 65535 {
|
||||
return Target{}, fmt.Errorf("目标端口必须在 1-65535 之间")
|
||||
}
|
||||
return Target{Host: host, Port: port}, nil
|
||||
}
|
||||
|
||||
func hasScheme(value string) bool {
|
||||
parsed, err := url.Parse(value)
|
||||
if err != nil || parsed.Scheme == "" {
|
||||
return false
|
||||
}
|
||||
if strings.Contains(value, "://") {
|
||||
return true
|
||||
}
|
||||
colon := strings.IndexByte(value, ':')
|
||||
if colon <= 0 || strings.Contains(parsed.Scheme, ".") {
|
||||
return false
|
||||
}
|
||||
suffix := value[colon+1:]
|
||||
return strings.IndexByte(suffix, ':') == -1
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package nftables
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseSingleTargetAcceptsHostPortAndIPv6(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
raw string
|
||||
host string
|
||||
port int
|
||||
}{
|
||||
{name: "hostname", raw: "example.com:443", host: "example.com", port: 443},
|
||||
{name: "ipv4", raw: "198.51.100.20:8443", host: "198.51.100.20", port: 8443},
|
||||
{name: "ipv6", raw: "[2001:db8::1]:443", host: "2001:db8::1", port: 443},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
target, err := ParseSingleTarget(tt.raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseSingleTarget: %v", err)
|
||||
}
|
||||
if target.Host != tt.host || target.Port != tt.port {
|
||||
t.Fatalf("expected %s/%d, got %+v", tt.host, tt.port, target)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseSingleTargetRejectsUnsupportedValues(t *testing.T) {
|
||||
for _, raw := range []string{
|
||||
"",
|
||||
"example.com",
|
||||
"example.com:0",
|
||||
"example.com:65536",
|
||||
"a:1,b:2",
|
||||
"http://example.com:443",
|
||||
"https:443",
|
||||
"mailto:443",
|
||||
} {
|
||||
t.Run(raw, func(t *testing.T) {
|
||||
if _, err := ParseSingleTarget(raw); err == nil {
|
||||
t.Fatalf("expected error for %q", raw)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,146 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func RenderTable(plan NodePlan) string {
|
||||
var b strings.Builder
|
||||
b.WriteString("table inet flvx {\n")
|
||||
b.WriteString(" chain prerouting {\n")
|
||||
b.WriteString(" type nat hook prerouting priority dstnat; policy accept;\n")
|
||||
for _, rule := range sortedRules(plan.Rules) {
|
||||
family := nftAddressFamily(rule.TargetHost)
|
||||
dnatFamily := ""
|
||||
if family != "" {
|
||||
dnatFamily = family + " "
|
||||
}
|
||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
||||
b.WriteString(fmt.Sprintf(" %s dport %d counter dnat %sto %s comment %q\n",
|
||||
protocol,
|
||||
rule.InPort,
|
||||
dnatFamily,
|
||||
formatDNATTarget(rule.TargetHost, rule.TargetPort),
|
||||
counterComment(rule.ForwardID, CounterDirectionDNAT, protocol),
|
||||
))
|
||||
}
|
||||
}
|
||||
b.WriteString(" }\n\n")
|
||||
b.WriteString(" chain postrouting {\n")
|
||||
b.WriteString(" type nat hook postrouting priority srcnat; policy accept;\n")
|
||||
if len(plan.Rules) > 0 {
|
||||
b.WriteString(" masquerade comment \"flvx masquerade\"\n")
|
||||
}
|
||||
b.WriteString(" }\n\n")
|
||||
b.WriteString(" chain forward {\n")
|
||||
b.WriteString(" type filter hook forward priority filter; policy accept;\n")
|
||||
for _, rule := range sortedRules(plan.Rules) {
|
||||
family := nftAddressFamily(rule.TargetHost)
|
||||
if family == "" {
|
||||
continue
|
||||
}
|
||||
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
|
||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
||||
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
|
||||
rule.InPort,
|
||||
family,
|
||||
targetHost,
|
||||
protocol,
|
||||
rule.TargetPort,
|
||||
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
|
||||
))
|
||||
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
|
||||
rule.InPort,
|
||||
family,
|
||||
targetHost,
|
||||
protocol,
|
||||
rule.TargetPort,
|
||||
counterComment(rule.ForwardID, CounterDirectionFromTarget, protocol),
|
||||
))
|
||||
}
|
||||
}
|
||||
b.WriteString(" }\n")
|
||||
b.WriteString("}\n")
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func counterComment(forwardID int64, direction, protocol string) string {
|
||||
return fmt.Sprintf("flvx forward:%d %s %s", forwardID, direction, protocol)
|
||||
}
|
||||
|
||||
func RuleHash(rule Rule) string {
|
||||
protocols := normalizedProtocols(rule.Protocols)
|
||||
sum := sha256.Sum256([]byte(fmt.Sprintf("%d|%d|%s|%d|%s",
|
||||
rule.ForwardID,
|
||||
rule.InPort,
|
||||
strings.TrimSpace(rule.TargetHost),
|
||||
rule.TargetPort,
|
||||
strings.Join(protocols, ","),
|
||||
)))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func PlanHashes(plan NodePlan) map[int64]string {
|
||||
hashes := make(map[int64]string, len(plan.Rules))
|
||||
for _, rule := range plan.Rules {
|
||||
hashes[rule.ForwardID] = RuleHash(rule)
|
||||
}
|
||||
return hashes
|
||||
}
|
||||
|
||||
func sortedRules(rules []Rule) []Rule {
|
||||
out := append([]Rule(nil), rules...)
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
if out[i].InPort == out[j].InPort {
|
||||
return out[i].ForwardID < out[j].ForwardID
|
||||
}
|
||||
return out[i].InPort < out[j].InPort
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
func normalizedProtocols(protocols []string) []string {
|
||||
seen := map[string]struct{}{}
|
||||
out := make([]string, 0, 2)
|
||||
for _, protocol := range protocols {
|
||||
p := strings.ToLower(strings.TrimSpace(protocol))
|
||||
if p != "tcp" && p != "udp" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[p]; ok {
|
||||
continue
|
||||
}
|
||||
seen[p] = struct{}{}
|
||||
out = append(out, p)
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return []string{"tcp", "udp"}
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func formatDNATTarget(host string, port int) string {
|
||||
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
|
||||
if ip := net.ParseIP(trimmed); ip != nil && ip.To4() == nil {
|
||||
return fmt.Sprintf("[%s]:%d", trimmed, port)
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", trimmed, port)
|
||||
}
|
||||
|
||||
func nftAddressFamily(host string) string {
|
||||
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
|
||||
ip := net.ParseIP(trimmed)
|
||||
if ip == nil {
|
||||
return ""
|
||||
}
|
||||
if ip.To4() == nil {
|
||||
return "ip6"
|
||||
}
|
||||
return "ip"
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRenderTableIncludesDNATAndMasquerade(t *testing.T) {
|
||||
script := RenderTable(NodePlan{
|
||||
NodeID: 10,
|
||||
Rules: []Rule{
|
||||
{
|
||||
ForwardID: 42,
|
||||
InPort: 24000,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
expectedParts := []string{
|
||||
"table inet flvx",
|
||||
"type nat hook prerouting priority dstnat; policy accept;",
|
||||
"type nat hook postrouting priority srcnat; policy accept;",
|
||||
"tcp dport 24000 counter dnat ip to 198.51.100.20:443 comment \"flvx forward:42 dnat tcp\"",
|
||||
"udp dport 24000 counter dnat ip to 198.51.100.20:443 comment \"flvx forward:42 dnat udp\"",
|
||||
"masquerade comment \"flvx masquerade\"",
|
||||
}
|
||||
for _, part := range expectedParts {
|
||||
if !strings.Contains(script, part) {
|
||||
t.Fatalf("script missing %q:\n%s", part, script)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableBracketsIPv6Target(t *testing.T) {
|
||||
script := RenderTable(NodePlan{
|
||||
NodeID: 10,
|
||||
Rules: []Rule{
|
||||
{ForwardID: 42, InPort: 24000, TargetHost: "2001:db8::1", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
},
|
||||
})
|
||||
if !strings.Contains(script, "dnat ip6 to [2001:db8::1]:443") {
|
||||
t.Fatalf("expected bracketed IPv6 dnat target, got:\n%s", script)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 42,
|
||||
InPort: 12345,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`,
|
||||
`udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"`,
|
||||
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
|
||||
`ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
|
||||
`ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 43,
|
||||
InPort: 12346,
|
||||
TargetHost: "2001:db8::20",
|
||||
TargetPort: 8443,
|
||||
Protocols: []string{"tcp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`,
|
||||
`ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
|
||||
`ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableAccountingCountersIncludeOriginalPort(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{
|
||||
{ForwardID: 42, InPort: 12345, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
{ForwardID: 43, InPort: 12346, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTablePreservesHostnameDNATAndSkipsAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 44,
|
||||
InPort: 12347,
|
||||
TargetHost: "example.com",
|
||||
TargetPort: 9443,
|
||||
Protocols: []string{"tcp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
want := `tcp dport 12347 counter dnat to example.com:9443 comment "flvx forward:44 dnat tcp"`
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
unwantedLines := []string{
|
||||
`dnat ip to example.com`,
|
||||
`ip daddr example.com`,
|
||||
`ip saddr example.com`,
|
||||
}
|
||||
for _, unwanted := range unwantedLines {
|
||||
if strings.Contains(got, unwanted) {
|
||||
t.Fatalf("RenderTable() unexpectedly contains %q\n%s", unwanted, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHashIsStable(t *testing.T) {
|
||||
rule := Rule{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}
|
||||
if RuleHash(rule) != RuleHash(rule) {
|
||||
t.Fatalf("expected stable rule hash")
|
||||
}
|
||||
if RuleHash(rule) == RuleHash(Rule{ForwardID: 42, InPort: 24001, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}) {
|
||||
t.Fatalf("expected hash to change when port changes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHashIgnoresBindIPWhenRenderingDoesNotUseIt(t *testing.T) {
|
||||
base := Rule{
|
||||
ForwardID: 42,
|
||||
InPort: 24000,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}
|
||||
withBind := base
|
||||
withBind.BindIP = "192.0.2.10"
|
||||
|
||||
if RuleHash(base) != RuleHash(withBind) {
|
||||
t.Fatalf("expected bind IP to be ignored by hash when it is not rendered")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,206 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/ssh"
|
||||
)
|
||||
|
||||
type Runner interface {
|
||||
ApplyScript(ctx context.Context, cfg SSHConfig, script string) error
|
||||
Test(ctx context.Context, cfg SSHConfig) error
|
||||
ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error)
|
||||
}
|
||||
|
||||
type SSHRunner struct {
|
||||
Timeout time.Duration
|
||||
}
|
||||
|
||||
func NewSSHRunner() *SSHRunner {
|
||||
return &SSHRunner{Timeout: 15 * time.Second}
|
||||
}
|
||||
|
||||
func (r *SSHRunner) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
return r.run(ctx, cfg, "command -v nft >/dev/null 2>&1 && nft --version >/dev/null 2>&1")
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
nft := nftBinary(cfg)
|
||||
command := "tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft) || exit 1\n" +
|
||||
"cleanup() {\n" +
|
||||
" rm -f \"$tmp\"\n" +
|
||||
"}\n" +
|
||||
"trap cleanup EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
nft + " -c -f \"$tmp\"\n" +
|
||||
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
|
||||
" " + nft + " delete table inet flvx\n" +
|
||||
"fi\n" +
|
||||
nft + " -f \"$tmp\""
|
||||
return r.run(ctx, cfg, command)
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
|
||||
return r.runOutput(ctx, cfg, nftBinary(cfg)+" -j list table inet flvx")
|
||||
}
|
||||
|
||||
func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) error {
|
||||
_, err := r.runOutput(ctx, cfg, command)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *SSHRunner) runOutput(ctx context.Context, cfg SSHConfig, command string) ([]byte, error) {
|
||||
timeout := r.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 15 * time.Second
|
||||
}
|
||||
runCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
clientConfig, err := buildSSHClientConfig(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
addr := net.JoinHostPort(strings.TrimSpace(cfg.Host), fmt.Sprintf("%d", normalizedSSHPort(cfg.Port)))
|
||||
dialer := net.Dialer{Timeout: timeout}
|
||||
conn, err := dialer.DialContext(runCtx, "tcp", addr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 连接失败: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 认证失败: %w", err)
|
||||
}
|
||||
client := ssh.NewClient(sshConn, chans, reqs)
|
||||
defer client.Close()
|
||||
|
||||
session, err := client.NewSession()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 会话创建失败: %w", err)
|
||||
}
|
||||
defer session.Close()
|
||||
|
||||
var stdout bytes.Buffer
|
||||
var stderr bytes.Buffer
|
||||
session.Stdout = &stdout
|
||||
session.Stderr = &stderr
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- session.Run(command)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-runCtx.Done():
|
||||
_ = session.Close()
|
||||
return nil, fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
message := strings.TrimSpace(stderr.String())
|
||||
if message != "" {
|
||||
return nil, fmt.Errorf("远程执行失败: %s: %w", message, err)
|
||||
}
|
||||
return nil, fmt.Errorf("远程执行失败: %w", err)
|
||||
}
|
||||
return stdout.Bytes(), nil
|
||||
}
|
||||
}
|
||||
|
||||
func buildSSHClientConfig(cfg SSHConfig) (*ssh.ClientConfig, error) {
|
||||
if strings.TrimSpace(cfg.Host) == "" {
|
||||
return nil, fmt.Errorf("SSH 主机不能为空")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Username) == "" {
|
||||
return nil, fmt.Errorf("SSH 用户名不能为空")
|
||||
}
|
||||
|
||||
auth, err := authMethods(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(auth) == 0 {
|
||||
return nil, fmt.Errorf("SSH 认证方式不能为空")
|
||||
}
|
||||
|
||||
return &ssh.ClientConfig{
|
||||
User: strings.TrimSpace(cfg.Username),
|
||||
Auth: auth,
|
||||
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
||||
Timeout: 15 * time.Second,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func authMethods(cfg SSHConfig) ([]ssh.AuthMethod, error) {
|
||||
switch strings.ToLower(strings.TrimSpace(cfg.AuthType)) {
|
||||
case "":
|
||||
if strings.TrimSpace(cfg.PrivateKey) == "" {
|
||||
return nil, fmt.Errorf("SSH 私钥不能为空")
|
||||
}
|
||||
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
|
||||
case "password":
|
||||
if cfg.Password == "" {
|
||||
return nil, fmt.Errorf("SSH 密码不能为空")
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.Password(cfg.Password)}, nil
|
||||
case "private_key":
|
||||
if strings.TrimSpace(cfg.PrivateKey) == "" {
|
||||
return nil, fmt.Errorf("SSH 私钥不能为空")
|
||||
}
|
||||
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("不支持的 SSH 认证方式: %s", cfg.AuthType)
|
||||
}
|
||||
}
|
||||
|
||||
func parsePrivateKey(privateKey, passphrase string) (ssh.Signer, error) {
|
||||
if passphrase != "" {
|
||||
signer, err := ssh.ParsePrivateKeyWithPassphrase([]byte(privateKey), []byte(passphrase))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
|
||||
}
|
||||
return signer, nil
|
||||
}
|
||||
signer, err := ssh.ParsePrivateKey([]byte(privateKey))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
|
||||
}
|
||||
return signer, nil
|
||||
}
|
||||
|
||||
func nftCommand(cfg SSHConfig, command string) string {
|
||||
return "sh -lc " + sshQuote(command)
|
||||
}
|
||||
|
||||
func nftBinary(cfg SSHConfig) string {
|
||||
if strings.EqualFold(strings.TrimSpace(cfg.SudoMode), "sudo") {
|
||||
return "sudo -n nft"
|
||||
}
|
||||
return "nft"
|
||||
}
|
||||
|
||||
func sshQuote(value string) string {
|
||||
return "'" + strings.ReplaceAll(value, "'", "'\"'\"'") + "'"
|
||||
}
|
||||
|
||||
func normalizedSSHPort(port int) int {
|
||||
if port <= 0 {
|
||||
return 22
|
||||
}
|
||||
return port
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
|
||||
privateKey := mustGeneratePrivateKey(t)
|
||||
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
|
||||
if err != nil {
|
||||
t.Fatalf("authMethods: %v", err)
|
||||
}
|
||||
if len(methods) != 1 {
|
||||
t.Fatalf("expected 1 auth method, got %d", len(methods))
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthMethodsDefaultPrivateKeyRequiresKey(t *testing.T) {
|
||||
_, err := authMethods(SSHConfig{})
|
||||
if err == nil || !strings.Contains(err.Error(), "SSH 私钥不能为空") {
|
||||
t.Fatalf("expected private key required error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustGeneratePrivateKey(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
key, err := rsa.GenerateKey(rand.Reader, 1024)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateKey: %v", err)
|
||||
}
|
||||
block := &pem.Block{
|
||||
Type: "RSA PRIVATE KEY",
|
||||
Bytes: x509.MarshalPKCS1PrivateKey(key),
|
||||
}
|
||||
return string(pem.EncodeToMemory(block))
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package nftables
|
||||
|
||||
const (
|
||||
ModeAgent = "agent"
|
||||
ModeNftables = "nftables"
|
||||
|
||||
StatusPending = "pending"
|
||||
StatusApplied = "applied"
|
||||
StatusError = "error"
|
||||
|
||||
CounterDirectionDNAT = "dnat"
|
||||
CounterDirectionToTarget = "to-target"
|
||||
CounterDirectionFromTarget = "from-target"
|
||||
)
|
||||
|
||||
type Target struct {
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
|
||||
type Rule struct {
|
||||
ForwardID int64
|
||||
InPort int
|
||||
BindIP string
|
||||
TargetHost string
|
||||
TargetPort int
|
||||
Protocols []string
|
||||
}
|
||||
|
||||
type NodePlan struct {
|
||||
NodeID int64
|
||||
Rules []Rule
|
||||
}
|
||||
|
||||
type SSHConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
AuthType string
|
||||
Password string
|
||||
PrivateKey string
|
||||
Passphrase string
|
||||
SudoMode string
|
||||
}
|
||||
|
||||
type ApplyResult struct {
|
||||
NodeID int64
|
||||
Script string
|
||||
Hashes map[int64]string
|
||||
}
|
||||
@@ -87,6 +87,7 @@ type Node struct {
|
||||
UDPListenAddr string `gorm:"column:udp_listen_addr;type:varchar(100);not null;default:'[::]'"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
IsRemote int `gorm:"column:is_remote;default:0"`
|
||||
ForwardMode string `gorm:"column:forward_mode;type:varchar(20);not null;default:'agent'"`
|
||||
RemoteURL sql.NullString `gorm:"column:remote_url;type:text"`
|
||||
RemoteToken sql.NullString `gorm:"column:remote_token;type:text"`
|
||||
RemoteConfig sql.NullString `gorm:"column:remote_config;type:text"`
|
||||
@@ -95,6 +96,57 @@ type Node struct {
|
||||
|
||||
func (Node) TableName() string { return "node" }
|
||||
|
||||
type NodeSSHConfig struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex"`
|
||||
Host string `gorm:"type:varchar(255);not null"`
|
||||
Port int `gorm:"not null;default:22"`
|
||||
Username string `gorm:"type:varchar(100);not null"`
|
||||
AuthType string `gorm:"column:auth_type;type:varchar(20);not null"`
|
||||
Password sql.NullString `gorm:"type:text"`
|
||||
PrivateKey sql.NullString `gorm:"column:private_key;type:text"`
|
||||
Passphrase sql.NullString `gorm:"type:text"`
|
||||
SudoMode string `gorm:"column:sudo_mode;type:varchar(20);not null;default:'none'"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (NodeSSHConfig) TableName() string { return "node_ssh_config" }
|
||||
|
||||
type NftRuleBinding struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
ForwardID int64 `gorm:"column:forward_id;not null;uniqueIndex:idx_nft_rule_binding_forward_node;index"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_nft_rule_binding_forward_node;index"`
|
||||
InPort int `gorm:"column:in_port;not null"`
|
||||
Protocols string `gorm:"type:varchar(20);not null;default:'tcp,udp'"`
|
||||
TargetAddr string `gorm:"column:target_addr;type:text;not null"`
|
||||
BindIP string `gorm:"column:bind_ip;type:text;not null;default:''"`
|
||||
RuleHash string `gorm:"column:rule_hash;type:varchar(128);not null;default:''"`
|
||||
Status string `gorm:"type:varchar(20);not null;default:'pending'"`
|
||||
LastError string `gorm:"column:last_error;type:text;not null;default:''"`
|
||||
AppliedTime int64 `gorm:"column:applied_time;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (NftRuleBinding) TableName() string { return "nft_rule_binding" }
|
||||
|
||||
type NftCounterState struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||
ForwardID int64 `gorm:"column:forward_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||
Protocol string `gorm:"type:varchar(10);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||
Direction string `gorm:"type:varchar(20);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||
RuleHash string `gorm:"column:rule_hash;type:varchar(128);not null;default:''"`
|
||||
Bytes int64 `gorm:"not null;default:0"`
|
||||
Packets int64 `gorm:"not null;default:0"`
|
||||
CollectedTime int64 `gorm:"column:collected_time;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (NftCounterState) TableName() string { return "nft_counter_state" }
|
||||
|
||||
type SpeedLimit struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
@@ -602,6 +654,7 @@ type NodeRecord struct {
|
||||
UDPListenAddr string
|
||||
InterfaceName string
|
||||
IsRemote int
|
||||
ForwardMode string
|
||||
RemoteURL string
|
||||
RemoteToken string
|
||||
RemoteConfig string
|
||||
|
||||
@@ -104,6 +104,19 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
|
||||
return nil
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
return applyFlowUploadDeltasTx(tx, deltas)
|
||||
})
|
||||
}
|
||||
|
||||
func applyFlowUploadDeltasTx(tx *gorm.DB, deltas []FlowUploadCounterDelta) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
if len(deltas) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
forwardTotals := make(map[int64][2]int64, len(deltas))
|
||||
userTotals := make(map[int64][2]int64, len(deltas))
|
||||
userTunnelTotals := make(map[int64][2]int64, len(deltas))
|
||||
@@ -128,36 +141,34 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
|
||||
}
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) {
|
||||
total := forwardTotals[forwardID]
|
||||
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) {
|
||||
total := forwardTotals[forwardID]
|
||||
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, userID := range sortedFlowUploadTargetIDs(userTotals) {
|
||||
total := userTotals[userID]
|
||||
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, userID := range sortedFlowUploadTargetIDs(userTotals) {
|
||||
total := userTotals[userID]
|
||||
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) {
|
||||
total := userTunnelTotals[userTunnelID]
|
||||
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) {
|
||||
total := userTunnelTotals[userTunnelID]
|
||||
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ─── Open / Close ────────────────────────────────────────────────────
|
||||
@@ -189,7 +200,6 @@ func Open(path string) (*Repository, error) {
|
||||
_ = sqlDB.Close()
|
||||
return nil, fmt.Errorf("prepare sqlite legacy schema: %w", err)
|
||||
}
|
||||
|
||||
if err := autoMigrateAll(db); err != nil {
|
||||
_ = sqlDB.Close()
|
||||
return nil, fmt.Errorf("auto migrate: %w", err)
|
||||
@@ -269,12 +279,21 @@ func (r *Repository) Close() error {
|
||||
}
|
||||
|
||||
func autoMigrateAll(db *gorm.DB) error {
|
||||
if db.Dialector.Name() == "sqlite" {
|
||||
if err := prepareSQLiteNftablesColumns(db); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
models := []interface{}{
|
||||
&model.User{},
|
||||
&model.UserQuota{},
|
||||
&model.Forward{},
|
||||
&model.ForwardPort{},
|
||||
&model.Node{},
|
||||
&model.NodeSSHConfig{},
|
||||
&model.NftRuleBinding{},
|
||||
&model.NftCounterState{},
|
||||
&model.SpeedLimit{},
|
||||
&model.StatisticsFlow{},
|
||||
&model.Tunnel{},
|
||||
@@ -418,6 +437,21 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func prepareSQLiteNftablesColumns(db *gorm.DB) error {
|
||||
if db == nil || db.Dialector.Name() != "sqlite" {
|
||||
return nil
|
||||
}
|
||||
if !db.Migrator().HasTable(&model.Node{}) {
|
||||
return nil
|
||||
}
|
||||
if !db.Migrator().HasColumn(&model.Node{}, "forward_mode") {
|
||||
if err := db.Exec("ALTER TABLE node ADD COLUMN forward_mode varchar(20) NOT NULL DEFAULT 'agent'").Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func seedData(db *gorm.DB) {
|
||||
var adminCount int64
|
||||
if err := db.Model(&model.User{}).Where("id = ?", 1).Count(&adminCount).Error; err == nil && adminCount == 0 {
|
||||
@@ -793,6 +827,28 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
|
||||
if err := r.db.Order("inx ASC, id ASC").Find(&nodes).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeIDs := make([]int64, 0, len(nodes))
|
||||
for _, n := range nodes {
|
||||
if defaultNodeForwardMode(n.ForwardMode) == "nftables" {
|
||||
nodeIDs = append(nodeIDs, n.ID)
|
||||
}
|
||||
}
|
||||
sshConfigByNodeID := make(map[int64]map[string]interface{}, len(nodeIDs))
|
||||
if len(nodeIDs) > 0 {
|
||||
var configs []model.NodeSSHConfig
|
||||
if err := r.db.Where("node_id IN ?", nodeIDs).Find(&configs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, cfg := range configs {
|
||||
sshConfigByNodeID[cfg.NodeID] = map[string]interface{}{
|
||||
"host": cfg.Host,
|
||||
"port": cfg.Port,
|
||||
"username": cfg.Username,
|
||||
"authType": cfg.AuthType,
|
||||
"sudoMode": cfg.SudoMode,
|
||||
}
|
||||
}
|
||||
}
|
||||
items := make([]map[string]interface{}, 0, len(nodes))
|
||||
for _, n := range nodes {
|
||||
items = append(items, map[string]interface{}{
|
||||
@@ -810,11 +866,13 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
|
||||
"version": nullableString(n.Version),
|
||||
"http": n.HTTP, "tls": n.TLS, "socks": n.Socks,
|
||||
"status": n.Status, "isRemote": n.IsRemote,
|
||||
"forwardMode": defaultNodeForwardMode(n.ForwardMode),
|
||||
"remoteUrl": nullableString(n.RemoteURL),
|
||||
"remoteToken": nullableString(n.RemoteToken),
|
||||
"remoteConfig": nullableString(n.RemoteConfig),
|
||||
"expiryReminderDismissed": n.ExpiryReminderDismissed,
|
||||
"interfaceName": nullableString(n.InterfaceName),
|
||||
"sshConfig": sshConfigByNodeID[n.ID],
|
||||
})
|
||||
}
|
||||
return items, nil
|
||||
|
||||
@@ -233,7 +233,8 @@ func nodeRecordFromModel(n *model.Node) *model.NodeRecord {
|
||||
Status: n.Status,
|
||||
PortRange: n.Port,
|
||||
TCPListenAddr: n.TCPListenAddr, UDPListenAddr: n.UDPListenAddr,
|
||||
IsRemote: n.IsRemote,
|
||||
IsRemote: n.IsRemote,
|
||||
ForwardMode: defaultNodeForwardMode(n.ForwardMode),
|
||||
}
|
||||
if n.ServerIPV4.Valid {
|
||||
rec.ServerIPv4 = strings.TrimSpace(n.ServerIPV4.String)
|
||||
|
||||
@@ -11,6 +11,8 @@ import (
|
||||
|
||||
type FlowUploadForwardMeta struct {
|
||||
ForwardID int64
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
TunnelID int64
|
||||
TrafficRatio float64
|
||||
TunnelFlow int64
|
||||
@@ -59,6 +61,8 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
|
||||
type row struct {
|
||||
ForwardID int64 `gorm:"column:forward_id"`
|
||||
UserID int64 `gorm:"column:user_id"`
|
||||
UserTunnelID int64 `gorm:"column:user_tunnel_id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id"`
|
||||
TrafficRatio float64 `gorm:"column:traffic_ratio"`
|
||||
TunnelFlow int64 `gorm:"column:tunnel_flow"`
|
||||
@@ -68,8 +72,9 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
for _, chunk := range chunkFlowUploadForwardIDs(ids) {
|
||||
var rows []row
|
||||
err := r.db.Table("forward AS f").
|
||||
Select("f.id AS forward_id, f.tunnel_id AS tunnel_id, t.traffic_ratio AS traffic_ratio, t.flow AS tunnel_flow").
|
||||
Select("f.id AS forward_id, f.user_id AS user_id, COALESCE(ut.id, 0) AS user_tunnel_id, f.tunnel_id AS tunnel_id, t.traffic_ratio AS traffic_ratio, t.flow AS tunnel_flow").
|
||||
Joins("LEFT JOIN tunnel t ON t.id = f.tunnel_id").
|
||||
Joins("LEFT JOIN user_tunnel ut ON ut.user_id = f.user_id AND ut.tunnel_id = f.tunnel_id").
|
||||
Where("f.id IN ?", chunk).
|
||||
Scan(&rows).Error
|
||||
if err != nil {
|
||||
@@ -84,6 +89,8 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
}
|
||||
out[row.ForwardID] = FlowUploadForwardMeta{
|
||||
ForwardID: row.ForwardID,
|
||||
UserID: row.UserID,
|
||||
UserTunnelID: row.UserTunnelID,
|
||||
TunnelID: row.TunnelID,
|
||||
TrafficRatio: row.TrafficRatio,
|
||||
TunnelFlow: row.TunnelFlow,
|
||||
|
||||
@@ -64,7 +64,7 @@ func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("get metas: %v", err)
|
||||
}
|
||||
if metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
|
||||
if metas[20].UserID != 2 || metas[20].UserTunnelID != 10 || metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
|
||||
t.Fatalf("unexpected meta for forward 20: %#v", metas[20])
|
||||
}
|
||||
if _, ok := metas[99]; ok {
|
||||
@@ -86,6 +86,99 @@ func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNftTrafficAccountingAppliesFlowQuotaAndStates(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
seedFlowBatchRows(t, r, nowMs)
|
||||
|
||||
quotaViews, err := r.ApplyNftTrafficAccounting(
|
||||
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
|
||||
map[int64]int64{2: 1140},
|
||||
[]NftCounterStateInput{{
|
||||
NodeID: 11,
|
||||
ForwardID: 20,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1400,
|
||||
Packets: 14,
|
||||
CollectedTime: nowMs,
|
||||
}},
|
||||
now,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("ApplyNftTrafficAccounting: %v", err)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 480 {
|
||||
t.Fatalf("expected forward in_flow=480, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 660 {
|
||||
t.Fatalf("expected user out_flow=660, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 480 {
|
||||
t.Fatalf("expected user_tunnel in_flow=480, got %d", got)
|
||||
}
|
||||
if quotaViews[2] == nil || quotaViews[2].DailyUsedBytes != 1140 || quotaViews[2].MonthlyUsedBytes != 1140 {
|
||||
t.Fatalf("unexpected quota view: %#v", quotaViews[2])
|
||||
}
|
||||
states, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(states) != 1 || states[0].ForwardID != 20 || states[0].Bytes != 1400 {
|
||||
t.Fatalf("unexpected nft counter state: %+v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNftTrafficAccountingRollsBackFlowAndQuotaWhenStateWriteFails(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting-rollback.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
seedFlowBatchRows(t, r, nowMs)
|
||||
if err := r.DB().Exec(`DROP TABLE nft_counter_state`).Error; err != nil {
|
||||
t.Fatalf("drop nft_counter_state: %v", err)
|
||||
}
|
||||
|
||||
_, err = r.ApplyNftTrafficAccounting(
|
||||
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
|
||||
map[int64]int64{2: 1140},
|
||||
[]NftCounterStateInput{{
|
||||
NodeID: 11,
|
||||
ForwardID: 20,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1400,
|
||||
Packets: 14,
|
||||
CollectedTime: nowMs,
|
||||
}},
|
||||
now,
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatalf("expected ApplyNftTrafficAccounting to fail")
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 0 {
|
||||
t.Fatalf("expected forward flow rollback, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 0 {
|
||||
t.Fatalf("expected user flow rollback, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT COALESCE((SELECT daily_used_bytes FROM user_quota WHERE user_id = 2), 0)`); got != 0 {
|
||||
t.Fatalf("expected quota rollback, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db"))
|
||||
if err != nil {
|
||||
@@ -106,7 +199,7 @@ func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T
|
||||
if !ok {
|
||||
t.Fatalf("expected metadata for forward with missing tunnel row")
|
||||
}
|
||||
if meta.ForwardID != 25 || meta.TunnelID != 99 || meta.TrafficRatio != 1 || meta.TunnelFlow != 1 {
|
||||
if meta.ForwardID != 25 || meta.UserID != 2 || meta.UserTunnelID != 0 || meta.TunnelID != 99 || meta.TrafficRatio != 1 || meta.TunnelFlow != 1 {
|
||||
t.Fatalf("unexpected fallback meta: %#v", meta)
|
||||
}
|
||||
}
|
||||
@@ -167,3 +260,19 @@ func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...inter
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func seedFlowBatchRows(t *testing.T, r *Repository, now int64) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 't1', 2.0, 1, 'tls', 3, ?, ?, 1, NULL, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(20, 2, 'u2', 'f20', 1, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -220,7 +220,7 @@ func (r *Repository) GetUserDefaultsForTunnel(userID int64) (flow int64, num int
|
||||
return user.Flow, user.Num, user.ExpTime, user.FlowResetTime, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}) error {
|
||||
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}, forwardMode string) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -247,6 +247,7 @@ func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serve
|
||||
UDPListenAddr: udpAddr,
|
||||
Inx: inx,
|
||||
IsRemote: isRemote,
|
||||
ForwardMode: defaultNodeForwardMode(forwardMode),
|
||||
RemoteURL: nullStringFromInterface(remoteURL),
|
||||
RemoteToken: nullStringFromInterface(remoteToken),
|
||||
RemoteConfig: nullStringFromInterface(remoteConfig),
|
||||
@@ -266,31 +267,44 @@ func (r *Repository) GetNodeStatusFields(nodeID int64) (status, httpFlag, tlsFla
|
||||
return node.Status, node.HTTP, node.TLS, node.Socks, nil
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
|
||||
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, expiryTime, renewalCycle interface{}, forwardMode string, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
updates := map[string]interface{}{
|
||||
"name": name,
|
||||
"remark": nullStringFromInterface(remark),
|
||||
"expiry_time": nullInt64FromInterface(expiryTime),
|
||||
"renewal_cycle": nullStringFromInterface(renewalCycle),
|
||||
"server_ip": serverIP,
|
||||
"server_ip_v4": nullStringFromInterface(serverIPV4),
|
||||
"server_ip_v6": nullStringFromInterface(serverIPV6),
|
||||
"extra_ips": nullStringFromInterface(extraIPs),
|
||||
"port": stringFromInterface(port),
|
||||
"interface_name": nullStringFromInterface(interfaceName),
|
||||
"http": httpFlag,
|
||||
"tls": tlsFlag,
|
||||
"socks": socksFlag,
|
||||
"tcp_listen_addr": tcpAddr,
|
||||
"udp_listen_addr": udpAddr,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
"expiry_reminder_dismissed": 0,
|
||||
}
|
||||
if strings.TrimSpace(forwardMode) != "" {
|
||||
updates["forward_mode"] = defaultNodeForwardMode(forwardMode)
|
||||
}
|
||||
return r.db.Model(&model.Node{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"remark": nullStringFromInterface(remark),
|
||||
"expiry_time": nullInt64FromInterface(expiryTime),
|
||||
"renewal_cycle": nullStringFromInterface(renewalCycle),
|
||||
"server_ip": serverIP,
|
||||
"server_ip_v4": nullStringFromInterface(serverIPV4),
|
||||
"server_ip_v6": nullStringFromInterface(serverIPV6),
|
||||
"extra_ips": nullStringFromInterface(extraIPs),
|
||||
"port": stringFromInterface(port),
|
||||
"interface_name": nullStringFromInterface(interfaceName),
|
||||
"http": httpFlag,
|
||||
"tls": tlsFlag,
|
||||
"socks": socksFlag,
|
||||
"tcp_listen_addr": tcpAddr,
|
||||
"udp_listen_addr": udpAddr,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
"expiry_reminder_dismissed": 0,
|
||||
}).Error
|
||||
Updates(updates).Error
|
||||
}
|
||||
|
||||
func defaultNodeForwardMode(mode string) string {
|
||||
switch strings.TrimSpace(strings.ToLower(mode)) {
|
||||
case "nftables":
|
||||
return "nftables"
|
||||
default:
|
||||
return "agent"
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeSecret(nodeID int64) (string, error) {
|
||||
@@ -755,6 +769,9 @@ func (r *Repository) DeleteForwardCascade(forwardID int64) error {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.NftCounterState{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.ForwardPort{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const (
|
||||
nftCounterProtocolTCP = "tcp"
|
||||
nftCounterProtocolUDP = "udp"
|
||||
|
||||
nftCounterDirectionToTarget = "to-target"
|
||||
nftCounterDirectionFromTarget = "from-target"
|
||||
)
|
||||
|
||||
type NftCounterStateInput struct {
|
||||
NodeID int64
|
||||
ForwardID int64
|
||||
Protocol string
|
||||
Direction string
|
||||
RuleHash string
|
||||
Bytes uint64
|
||||
Packets uint64
|
||||
CollectedTime int64
|
||||
}
|
||||
|
||||
type NftablesCollectionNode struct {
|
||||
NodeID int64
|
||||
Config model.NodeSSHConfig
|
||||
}
|
||||
|
||||
func (r *Repository) ListNftablesNodesForCollection() ([]NftablesCollectionNode, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
type collectionRow struct {
|
||||
NodeID int64 `gorm:"column:node_id"`
|
||||
ConfigID int64 `gorm:"column:config_id"`
|
||||
Host string `gorm:"column:host"`
|
||||
Port int `gorm:"column:port"`
|
||||
Username string `gorm:"column:username"`
|
||||
AuthType string `gorm:"column:auth_type"`
|
||||
Password string `gorm:"column:password"`
|
||||
PrivateKey string `gorm:"column:private_key"`
|
||||
Passphrase string `gorm:"column:passphrase"`
|
||||
SudoMode string `gorm:"column:sudo_mode"`
|
||||
CreatedTime int64 `gorm:"column:created_time"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time"`
|
||||
}
|
||||
|
||||
var rows []collectionRow
|
||||
if err := r.db.Table("node").
|
||||
Select("node.id AS node_id, node_ssh_config.id AS config_id, node_ssh_config.host, node_ssh_config.port, node_ssh_config.username, node_ssh_config.auth_type, node_ssh_config.password, node_ssh_config.private_key, node_ssh_config.passphrase, node_ssh_config.sudo_mode, node_ssh_config.created_time, node_ssh_config.updated_time").
|
||||
Joins("JOIN node_ssh_config ON node_ssh_config.node_id = node.id").
|
||||
Where("node.status = ? AND LOWER(TRIM(node.forward_mode)) = ?", 1, "nftables").
|
||||
Order("node.id ASC").
|
||||
Scan(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nodes := make([]NftablesCollectionNode, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
nodes = append(nodes, NftablesCollectionNode{
|
||||
NodeID: row.NodeID,
|
||||
Config: model.NodeSSHConfig{
|
||||
ID: row.ConfigID,
|
||||
NodeID: row.NodeID,
|
||||
Host: row.Host,
|
||||
Port: row.Port,
|
||||
Username: row.Username,
|
||||
AuthType: row.AuthType,
|
||||
Password: nullStringFromInterface(row.Password),
|
||||
PrivateKey: nullStringFromInterface(row.PrivateKey),
|
||||
Passphrase: nullStringFromInterface(row.Passphrase),
|
||||
SudoMode: row.SudoMode,
|
||||
CreatedTime: row.CreatedTime,
|
||||
UpdatedTime: row.UpdatedTime,
|
||||
},
|
||||
})
|
||||
}
|
||||
return nodes, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetNftCounterStatesByNode(nodeID int64) ([]model.NftCounterState, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var rows []model.NftCounterState
|
||||
err := r.db.Where("node_id = ?", nodeID).
|
||||
Order("forward_id ASC, protocol ASC, direction ASC").
|
||||
Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertNftCounterStates(inputs []NftCounterStateInput, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if len(inputs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
return upsertNftCounterStatesTx(tx, inputs, now)
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) ApplyNftTrafficAccounting(deltas []FlowUploadCounterDelta, quotaUsage map[int64]int64, states []NftCounterStateInput, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
quotaViews := map[int64]*model.UserQuotaView{}
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := applyFlowUploadDeltasTx(tx, deltas); err != nil {
|
||||
return err
|
||||
}
|
||||
var err error
|
||||
quotaViews, err = r.addUserQuotaUsageBatchTx(tx, quotaUsage, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return upsertNftCounterStatesTx(tx, states, now.UnixMilli())
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return quotaViews, nil
|
||||
}
|
||||
|
||||
func upsertNftCounterStatesTx(tx *gorm.DB, inputs []NftCounterStateInput, now int64) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
for _, input := range inputs {
|
||||
row, ok := nftCounterStateFromInput(input, now)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{
|
||||
{Name: "node_id"},
|
||||
{Name: "forward_id"},
|
||||
{Name: "protocol"},
|
||||
{Name: "direction"},
|
||||
},
|
||||
DoUpdates: clause.Assignments(map[string]interface{}{
|
||||
"rule_hash": row.RuleHash,
|
||||
"bytes": row.Bytes,
|
||||
"packets": row.Packets,
|
||||
"collected_time": row.CollectedTime,
|
||||
"updated_time": row.UpdatedTime,
|
||||
}),
|
||||
}).Create(&row).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNftCounterStatesByForward(forwardID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Where("forward_id = ?", forwardID).Delete(&model.NftCounterState{}).Error
|
||||
}
|
||||
|
||||
func nftCounterStateFromInput(input NftCounterStateInput, now int64) (model.NftCounterState, bool) {
|
||||
protocol := strings.ToLower(strings.TrimSpace(input.Protocol))
|
||||
direction := strings.ToLower(strings.TrimSpace(input.Direction))
|
||||
if input.NodeID <= 0 || input.ForwardID <= 0 || !isValidNftCounterProtocol(protocol) || !isValidNftCounterDirection(direction) {
|
||||
return model.NftCounterState{}, false
|
||||
}
|
||||
if input.Bytes > uint64(math.MaxInt64) || input.Packets > uint64(math.MaxInt64) {
|
||||
return model.NftCounterState{}, false
|
||||
}
|
||||
return model.NftCounterState{
|
||||
NodeID: input.NodeID,
|
||||
ForwardID: input.ForwardID,
|
||||
Protocol: protocol,
|
||||
Direction: direction,
|
||||
RuleHash: strings.TrimSpace(input.RuleHash),
|
||||
Bytes: int64(input.Bytes),
|
||||
Packets: int64(input.Packets),
|
||||
CollectedTime: input.CollectedTime,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}, true
|
||||
}
|
||||
|
||||
func isValidNftCounterProtocol(protocol string) bool {
|
||||
return protocol == nftCounterProtocolTCP || protocol == nftCounterProtocolUDP
|
||||
}
|
||||
|
||||
func isValidNftCounterDirection(direction string) bool {
|
||||
return direction == nftCounterDirectionToTarget || direction == nftCounterDirectionFromTarget
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"math"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestNftCounterStateUpsertUpdatesExistingKey(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
first := []NftCounterStateInput{
|
||||
{
|
||||
NodeID: 11,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 100,
|
||||
Packets: 10,
|
||||
CollectedTime: 1000,
|
||||
},
|
||||
{
|
||||
NodeID: 0,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
Bytes: 999,
|
||||
},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(first, 2000); err != nil {
|
||||
t.Fatalf("first UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
second := []NftCounterStateInput{
|
||||
{
|
||||
NodeID: 11,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-b",
|
||||
Bytes: 250,
|
||||
Packets: 25,
|
||||
CollectedTime: 3000,
|
||||
},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(second, 4000); err != nil {
|
||||
t.Fatalf("second UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected one counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
got := rows[0]
|
||||
if got.ForwardID != 42 || got.Protocol != "tcp" || got.Direction != "to-target" {
|
||||
t.Fatalf("unexpected counter state key: %+v", got)
|
||||
}
|
||||
if got.RuleHash != "hash-b" || got.Bytes != 250 || got.Packets != 25 || got.CollectedTime != 3000 {
|
||||
t.Fatalf("counter state was not updated: %+v", got)
|
||||
}
|
||||
if got.CreatedTime != 2000 || got.UpdatedTime != 4000 {
|
||||
t.Fatalf("unexpected timestamps after upsert: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateDeleteByForwardRemovesOnlyMatchingRows(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "a", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "b", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||
{NodeID: 12, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "c", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
if err := r.DeleteNftCounterStatesByForward(42); err != nil {
|
||||
t.Fatalf("DeleteNftCounterStatesByForward: %v", err)
|
||||
}
|
||||
|
||||
node11, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode(11): %v", err)
|
||||
}
|
||||
if len(node11) != 1 || node11[0].ForwardID != 43 {
|
||||
t.Fatalf("expected only forward 43 for node 11, got %+v", node11)
|
||||
}
|
||||
node12, err := r.GetNftCounterStatesByNode(12)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode(12): %v", err)
|
||||
}
|
||||
if len(node12) != 0 {
|
||||
t.Fatalf("expected forward 42 state removed from node 12, got %+v", node12)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteForwardCascadeRemovesNftCounterStateOnlyForDeletedForward(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwards := []model.Forward{
|
||||
{ID: 42, UserID: 1, UserName: "admin", Name: "forward-a", TunnelID: 10, RemoteAddr: "203.0.113.1:80", Strategy: "fifo", CreatedTime: now, UpdatedTime: now, Status: 1},
|
||||
{ID: 43, UserID: 1, UserName: "admin", Name: "forward-b", TunnelID: 10, RemoteAddr: "203.0.113.2:80", Strategy: "fifo", CreatedTime: now, UpdatedTime: now, Status: 1},
|
||||
}
|
||||
if err := r.DB().Create(&forwards).Error; err != nil {
|
||||
t.Fatalf("seed forwards: %v", err)
|
||||
}
|
||||
if err := r.UpsertNftCounterStates([]NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "a", Bytes: 100, Packets: 10, CollectedTime: now},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "b", Bytes: 200, Packets: 20, CollectedTime: now},
|
||||
}, now); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DeleteForwardCascade(42); err != nil {
|
||||
t.Fatalf("DeleteForwardCascade: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 || rows[0].ForwardID != 43 {
|
||||
t.Fatalf("expected only forward 43 counter state to remain, got %+v", rows)
|
||||
}
|
||||
var deletedForwardCount int64
|
||||
if err := r.DB().Model(&model.Forward{}).Where("id = ?", int64(42)).Count(&deletedForwardCount).Error; err != nil {
|
||||
t.Fatalf("count deleted forward: %v", err)
|
||||
}
|
||||
if deletedForwardCount != 0 {
|
||||
t.Fatalf("expected forward 42 deleted, count=%d", deletedForwardCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateUpsertSkipsInvalidProtocolAndDirection(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "icmp", Direction: "to-target", RuleHash: "bad-protocol", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "tcp", Direction: "sideways", RuleHash: "bad-direction", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 44, Protocol: " UDP ", Direction: " FROM-TARGET ", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
if rows[0].ForwardID != 44 || rows[0].Protocol != "udp" || rows[0].Direction != "from-target" {
|
||||
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateUpsertSkipsCountersAboveInt64(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "too-large", Bytes: uint64(math.MaxInt64) + 1, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
if rows[0].ForwardID != 43 || rows[0].Bytes != 300 || rows[0].Packets != 30 {
|
||||
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestListNftablesNodesForCollectionReturnsActiveNftablesWithSSHOrdered(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-collection.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
seedCollectionNode(t, r, 1, "agent", 1, now)
|
||||
seedCollectionNode(t, r, 2, " nftables ", 1, now)
|
||||
seedCollectionNode(t, r, 3, "NFTABLES", 0, now)
|
||||
seedCollectionNode(t, r, 4, "nftables", 1, now)
|
||||
seedCollectionNode(t, r, 5, "nftables", 1, now)
|
||||
|
||||
if err := r.UpsertNodeSSHConfig(4, NftSSHConfigInput{
|
||||
Host: "203.0.113.4",
|
||||
Port: 2222,
|
||||
Username: "root",
|
||||
AuthType: "password",
|
||||
Password: "secret-4",
|
||||
SudoMode: "none",
|
||||
}, now); err != nil {
|
||||
t.Fatalf("upsert ssh config 4: %v", err)
|
||||
}
|
||||
if err := r.UpsertNodeSSHConfig(2, NftSSHConfigInput{
|
||||
Host: "203.0.113.2",
|
||||
Port: 22,
|
||||
Username: "admin",
|
||||
AuthType: "private_key",
|
||||
SudoMode: "sudo",
|
||||
}, now); err != nil {
|
||||
t.Fatalf("upsert ssh config 2: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNftablesNodesForCollection()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftablesNodesForCollection: %v", err)
|
||||
}
|
||||
if len(nodes) != 2 {
|
||||
t.Fatalf("expected 2 collection nodes, got %d: %+v", len(nodes), nodes)
|
||||
}
|
||||
if nodes[0].NodeID != 2 || nodes[1].NodeID != 4 {
|
||||
t.Fatalf("expected nodes ordered by id [2 4], got [%d %d]", nodes[0].NodeID, nodes[1].NodeID)
|
||||
}
|
||||
if nodes[0].Config.NodeID != 2 || nodes[0].Config.Host != "203.0.113.2" || nodes[0].Config.Username != "admin" {
|
||||
t.Fatalf("unexpected first config: %+v", nodes[0].Config)
|
||||
}
|
||||
if nodes[1].Config.NodeID != 4 || nodes[1].Config.Port != 2222 || nodes[1].Config.Password.String != "secret-4" {
|
||||
t.Fatalf("unexpected second config: %+v", nodes[1].Config)
|
||||
}
|
||||
}
|
||||
|
||||
func seedCollectionNode(t *testing.T, r *Repository, id int64, forwardMode string, status int, now int64) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, port, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, forward_mode)
|
||||
VALUES(?, ?, 'secret', ?, '1000-2000', ?, ?, ?, '[::]', '[::]', 0, ?)
|
||||
`, id, "node", "198.51.100.1", now, now, status, forwardMode).Error; err != nil {
|
||||
t.Fatalf("insert node %d: %v", id, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
type NftSSHConfigInput struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
AuthType string
|
||||
Password string
|
||||
PrivateKey string
|
||||
Passphrase string
|
||||
SudoMode string
|
||||
}
|
||||
|
||||
type NftRuleBindingInput struct {
|
||||
ForwardID int64
|
||||
NodeID int64
|
||||
InPort int
|
||||
Protocols string
|
||||
TargetAddr string
|
||||
BindIP string
|
||||
RuleHash string
|
||||
Status string
|
||||
LastError string
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertNodeSSHConfig(nodeID int64, cfg NftSSHConfigInput, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if nodeID <= 0 {
|
||||
return errors.New("node id is required")
|
||||
}
|
||||
port := cfg.Port
|
||||
if port <= 0 {
|
||||
port = 22
|
||||
}
|
||||
authType := strings.TrimSpace(strings.ToLower(cfg.AuthType))
|
||||
if authType == "" {
|
||||
authType = "private_key"
|
||||
}
|
||||
sudoMode := strings.TrimSpace(strings.ToLower(cfg.SudoMode))
|
||||
if sudoMode == "" {
|
||||
sudoMode = "none"
|
||||
}
|
||||
row := model.NodeSSHConfig{
|
||||
NodeID: nodeID,
|
||||
Host: strings.TrimSpace(cfg.Host),
|
||||
Port: port,
|
||||
Username: strings.TrimSpace(cfg.Username),
|
||||
AuthType: authType,
|
||||
Password: nullStringFromInterface(cfg.Password),
|
||||
PrivateKey: nullStringFromInterface(cfg.PrivateKey),
|
||||
Passphrase: nullStringFromInterface(cfg.Passphrase),
|
||||
SudoMode: sudoMode,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
return r.db.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "node_id"}},
|
||||
DoUpdates: clause.Assignments(map[string]interface{}{
|
||||
"host": row.Host,
|
||||
"port": row.Port,
|
||||
"username": row.Username,
|
||||
"auth_type": row.AuthType,
|
||||
"password": row.Password,
|
||||
"private_key": row.PrivateKey,
|
||||
"passphrase": row.Passphrase,
|
||||
"sudo_mode": row.SudoMode,
|
||||
"updated_time": row.UpdatedTime,
|
||||
}),
|
||||
}).Create(&row).Error
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeSSHConfig(nodeID int64) (*model.NodeSSHConfig, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
return r.GetNodeSSHConfigTx(r.db, nodeID)
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeSSHConfigTx(tx *gorm.DB, nodeID int64) (*model.NodeSSHConfig, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
var cfg model.NodeSSHConfig
|
||||
if err := tx.Where("node_id = ?", nodeID).First(&cfg).Error; err != nil {
|
||||
return nil, normalizeNotFoundErr(err)
|
||||
}
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNodeSSHConfig(nodeID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Where("node_id = ?", nodeID).Delete(&model.NodeSSHConfig{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertNftRuleBinding(input NftRuleBindingInput, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
row := model.NftRuleBinding{
|
||||
ForwardID: input.ForwardID,
|
||||
NodeID: input.NodeID,
|
||||
InPort: input.InPort,
|
||||
Protocols: defaultString(strings.TrimSpace(input.Protocols), "tcp,udp"),
|
||||
TargetAddr: strings.TrimSpace(input.TargetAddr),
|
||||
BindIP: strings.TrimSpace(input.BindIP),
|
||||
RuleHash: strings.TrimSpace(input.RuleHash),
|
||||
Status: defaultString(strings.TrimSpace(input.Status), "pending"),
|
||||
LastError: strings.TrimSpace(input.LastError),
|
||||
AppliedTime: now,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
return r.db.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "forward_id"}, {Name: "node_id"}},
|
||||
DoUpdates: clause.Assignments(map[string]interface{}{
|
||||
"in_port": row.InPort,
|
||||
"protocols": row.Protocols,
|
||||
"target_addr": row.TargetAddr,
|
||||
"bind_ip": row.BindIP,
|
||||
"rule_hash": row.RuleHash,
|
||||
"status": row.Status,
|
||||
"last_error": row.LastError,
|
||||
"applied_time": row.AppliedTime,
|
||||
"updated_time": row.UpdatedTime,
|
||||
}),
|
||||
}).Create(&row).Error
|
||||
}
|
||||
|
||||
func (r *Repository) MarkNftRuleBindingError(forwardID, nodeID int64, message string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.NftRuleBinding{}).
|
||||
Where("forward_id = ? AND node_id = ?", forwardID, nodeID).
|
||||
Updates(map[string]interface{}{
|
||||
"status": "error",
|
||||
"last_error": strings.TrimSpace(message),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ListNftRuleBindingsByNode(nodeID int64) ([]model.NftRuleBinding, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var rows []model.NftRuleBinding
|
||||
err := r.db.Where("node_id = ?", nodeID).Order("forward_id ASC").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNftRuleBindingsByForward(forwardID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Where("forward_id = ?", forwardID).Delete(&model.NftRuleBinding{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeForwardMode(nodeID int64) (string, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return "", errors.New("repository not initialized")
|
||||
}
|
||||
return r.GetNodeForwardModeTx(r.db, nodeID)
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeForwardModeTx(tx *gorm.DB, nodeID int64) (string, error) {
|
||||
if tx == nil {
|
||||
return "", errors.New("database unavailable")
|
||||
}
|
||||
var row struct {
|
||||
ForwardMode sql.NullString `gorm:"column:forward_mode"`
|
||||
}
|
||||
err := tx.Model(&model.Node{}).Select("forward_mode").Where("id = ?", nodeID).First(&row).Error
|
||||
if err != nil {
|
||||
return "", normalizeNotFoundErr(err)
|
||||
}
|
||||
return defaultNodeForwardMode(row.ForwardMode.String), nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListActiveForwardsByNode(nodeID int64) ([]model.ForwardRecord, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var forwards []model.Forward
|
||||
err := r.db.Model(&model.Forward{}).
|
||||
Joins("JOIN forward_port ON forward_port.forward_id = forward.id").
|
||||
Where("forward_port.node_id = ? AND forward.status = 1", nodeID).
|
||||
Order("forward.id ASC").
|
||||
Distinct("forward.*").
|
||||
Find(&forwards).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
if strings.TrimSpace(rows[i].Strategy) == "" {
|
||||
rows[i].Strategy = "fifo"
|
||||
}
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func defaultString(value, fallback string) string {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fallback
|
||||
}
|
||||
return value
|
||||
}
|
||||
@@ -0,0 +1,214 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestNftablesNodeModeSSHConfigAndBindingPersistence(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreateNode(
|
||||
"nft-node",
|
||||
"secret",
|
||||
"203.0.113.10",
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
now,
|
||||
1,
|
||||
"[::]",
|
||||
"[::]",
|
||||
1,
|
||||
0,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"nftables",
|
||||
); err != nil {
|
||||
t.Fatalf("CreateNode: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes: %v", err)
|
||||
}
|
||||
if len(nodes) != 1 {
|
||||
t.Fatalf("expected 1 node, got %d", len(nodes))
|
||||
}
|
||||
nodeID := nodes[0]["id"].(int64)
|
||||
if got := nodes[0]["forwardMode"]; got != "nftables" {
|
||||
t.Fatalf("expected forwardMode nftables, got %#v", got)
|
||||
}
|
||||
|
||||
cfg := NftSSHConfigInput{
|
||||
Host: "203.0.113.10",
|
||||
Port: 22,
|
||||
Username: "root",
|
||||
AuthType: "private_key",
|
||||
PrivateKey: "encrypted-private-key",
|
||||
SudoMode: "none",
|
||||
}
|
||||
if err := r.UpsertNodeSSHConfig(nodeID, cfg, now); err != nil {
|
||||
t.Fatalf("UpsertNodeSSHConfig: %v", err)
|
||||
}
|
||||
loaded, err := r.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNodeSSHConfig: %v", err)
|
||||
}
|
||||
if loaded.Host != cfg.Host || loaded.Port != cfg.Port || loaded.Username != cfg.Username || loaded.AuthType != cfg.AuthType {
|
||||
t.Fatalf("unexpected ssh config: %+v", loaded)
|
||||
}
|
||||
|
||||
binding := NftRuleBindingInput{
|
||||
ForwardID: 42,
|
||||
NodeID: nodeID,
|
||||
InPort: 24000,
|
||||
Protocols: "tcp,udp",
|
||||
TargetAddr: "198.51.100.20:443",
|
||||
BindIP: "",
|
||||
RuleHash: "hash-a",
|
||||
Status: "applied",
|
||||
LastError: "",
|
||||
}
|
||||
if err := r.UpsertNftRuleBinding(binding, now); err != nil {
|
||||
t.Fatalf("UpsertNftRuleBinding: %v", err)
|
||||
}
|
||||
bindings, err := r.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftRuleBindingsByNode: %v", err)
|
||||
}
|
||||
if len(bindings) != 1 {
|
||||
t.Fatalf("expected 1 binding, got %d", len(bindings))
|
||||
}
|
||||
if bindings[0].ForwardID != 42 || bindings[0].RuleHash != "hash-a" || bindings[0].Status != "applied" {
|
||||
t.Fatalf("unexpected binding: %+v", bindings[0])
|
||||
}
|
||||
|
||||
if err := r.MarkNftRuleBindingError(42, nodeID, "nft failed", now+1); err != nil {
|
||||
t.Fatalf("MarkNftRuleBindingError: %v", err)
|
||||
}
|
||||
bindings, err = r.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftRuleBindingsByNode after error: %v", err)
|
||||
}
|
||||
if bindings[0].Status != "error" || !strings.Contains(bindings[0].LastError, "nft failed") {
|
||||
t.Fatalf("expected error binding, got %+v", bindings[0])
|
||||
}
|
||||
|
||||
if err := r.DeleteNftRuleBindingsByForward(42); err != nil {
|
||||
t.Fatalf("DeleteNftRuleBindingsByForward: %v", err)
|
||||
}
|
||||
bindings, err = r.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftRuleBindingsByNode after delete: %v", err)
|
||||
}
|
||||
if len(bindings) != 0 {
|
||||
t.Fatalf("expected no bindings after delete, got %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateNodeWithoutForwardModePreservesExistingMode(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreateNode(
|
||||
"nft-node",
|
||||
"secret",
|
||||
"203.0.113.11",
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
now,
|
||||
1,
|
||||
"[::]",
|
||||
"[::]",
|
||||
1,
|
||||
0,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"nftables",
|
||||
); err != nil {
|
||||
t.Fatalf("CreateNode: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes: %v", err)
|
||||
}
|
||||
if len(nodes) != 1 {
|
||||
t.Fatalf("expected 1 node, got %d", len(nodes))
|
||||
}
|
||||
nodeID := nodes[0]["id"].(int64)
|
||||
|
||||
if err := r.UpdateNode(
|
||||
nodeID,
|
||||
"nft-node-updated",
|
||||
"203.0.113.11",
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"",
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
"[::]",
|
||||
"[::]",
|
||||
now+1,
|
||||
); err != nil {
|
||||
t.Fatalf("UpdateNode: %v", err)
|
||||
}
|
||||
|
||||
gotNode, err := r.GetNodeRecord(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNodeRecord: %v", err)
|
||||
}
|
||||
if gotNode == nil {
|
||||
t.Fatal("expected node record, got nil")
|
||||
}
|
||||
if gotNode.ForwardMode != "nftables" {
|
||||
t.Fatalf("expected mapped forward mode nftables, got %q", gotNode.ForwardMode)
|
||||
}
|
||||
|
||||
nodes, err = r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes after update: %v", err)
|
||||
}
|
||||
if got := nodes[0]["forwardMode"]; got != "nftables" {
|
||||
t.Fatalf("expected persisted forwardMode nftables after update, got %#v", got)
|
||||
}
|
||||
}
|
||||
@@ -264,39 +264,11 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
var result map[int64]*model.UserQuotaView
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
userIDs := make([]int64, 0, len(usages))
|
||||
for userID := range usages {
|
||||
if userID > 0 {
|
||||
userIDs = append(userIDs, userID)
|
||||
}
|
||||
}
|
||||
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
|
||||
|
||||
for _, userID := range userIDs {
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
if usages[userID] > 0 {
|
||||
q.DailyUsedBytes += usages[userID]
|
||||
q.MonthlyUsedBytes += usages[userID]
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
|
||||
}
|
||||
return nil
|
||||
var err error
|
||||
result, err = r.addUserQuotaUsageBatchTx(tx, usages, now)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -304,6 +276,48 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *Repository) addUserQuotaUsageBatchTx(tx *gorm.DB, usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
if len(usages) == 0 {
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
userIDs := make([]int64, 0, len(usages))
|
||||
for userID := range usages {
|
||||
if userID > 0 {
|
||||
userIDs = append(userIDs, userID)
|
||||
}
|
||||
}
|
||||
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
|
||||
|
||||
for _, userID := range userIDs {
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
if usages[userID] > 0 {
|
||||
q.DailyUsedBytes += usages[userID]
|
||||
q.MonthlyUsedBytes += usages[userID]
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
|
||||
Generated
+1
-20
@@ -1496,28 +1496,24 @@ packages:
|
||||
engines: {node: ^20.19.0 || >=22.12.0}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
'@rolldown/binding-linux-arm64-musl@1.0.0-beta.53':
|
||||
resolution: {integrity: sha512-bGe5EBB8FVjHBR1mOLOPEFg1Lp3//7geqWkU5NIhxe+yH0W8FVrQ6WRYOap4SUTKdklD/dC4qPLREkMMQ855FA==}
|
||||
engines: {node: ^20.19.0 || >=22.12.0}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
'@rolldown/binding-linux-x64-gnu@1.0.0-beta.53':
|
||||
resolution: {integrity: sha512-qL+63WKVQs1CMvFedlPt0U9PiEKJOAL/bsHMKUDS6Vp2Q+YAv/QLPu8rcvkfIMvQ0FPU2WL0aX4eWwF6e/GAnA==}
|
||||
engines: {node: ^20.19.0 || >=22.12.0}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
'@rolldown/binding-linux-x64-musl@1.0.0-beta.53':
|
||||
resolution: {integrity: sha512-VGl9JIGjoJh3H8Mb+7xnVqODajBmrdOOb9lxWXdcmxyI+zjB2sux69br0hZJDTyLJfvBoYm439zPACYbCjGRmw==}
|
||||
engines: {node: ^20.19.0 || >=22.12.0}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
'@rolldown/binding-openharmony-arm64@1.0.0-beta.53':
|
||||
resolution: {integrity: sha512-B4iIserJXuSnNzA5xBLFUIjTfhNy7d9sq4FUMQY3GhQWGVhS2RWWzzDnkSU6MUt7/aHUrep0CdQfXUJI9D3W7A==}
|
||||
@@ -1683,56 +1679,48 @@ packages:
|
||||
engines: {node: '>= 10'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
'@tailwindcss/oxide-linux-arm64-gnu@4.2.4':
|
||||
resolution: {integrity: sha512-+E4wxJ0ZGOzSH325reXTWB48l42i93kQqMvDyz5gqfRzRZ7faNhnmvlV4EPGJU3QJM/3Ab5jhJ5pCRUsKn6OQw==}
|
||||
engines: {node: '>= 20'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
'@tailwindcss/oxide-linux-arm64-musl@4.1.11':
|
||||
resolution: {integrity: sha512-m/NVRFNGlEHJrNVk3O6I9ggVuNjXHIPoD6bqay/pubtYC9QIdAMpS+cswZQPBLvVvEF6GtSNONbDkZrjWZXYNQ==}
|
||||
engines: {node: '>= 10'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
'@tailwindcss/oxide-linux-arm64-musl@4.2.4':
|
||||
resolution: {integrity: sha512-bBADEGAbo4ASnppIziaQJelekCxdMaxisrk+fB7Thit72IBnALp9K6ffA2G4ruj90G9XRS2VQ6q2bCKbfFV82g==}
|
||||
engines: {node: '>= 20'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
'@tailwindcss/oxide-linux-x64-gnu@4.1.11':
|
||||
resolution: {integrity: sha512-YW6sblI7xukSD2TdbbaeQVDysIm/UPJtObHJHKxDEcW2exAtY47j52f8jZXkqE1krdnkhCMGqP3dbniu1Te2Fg==}
|
||||
engines: {node: '>= 10'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
'@tailwindcss/oxide-linux-x64-gnu@4.2.4':
|
||||
resolution: {integrity: sha512-7Mx25E4WTfnht0TVRTyC00j3i0M+EeFe7wguMDTlX4mRxafznw0CA8WJkFjWYH5BlgELd1kSjuU2JiPnNZbJDA==}
|
||||
engines: {node: '>= 20'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
'@tailwindcss/oxide-linux-x64-musl@4.1.11':
|
||||
resolution: {integrity: sha512-e3C/RRhGunWYNC3aSF7exsQkdXzQ/M+aYuZHKnw4U7KQwTJotnWsGOIVih0s2qQzmEzOFIJ3+xt7iq67K/p56Q==}
|
||||
engines: {node: '>= 10'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
'@tailwindcss/oxide-linux-x64-musl@4.2.4':
|
||||
resolution: {integrity: sha512-2wwJRF7nyhOR0hhHoChc04xngV3iS+akccHTGtz965FwF0up4b2lOdo6kI1EbDaEXKgvcrFBYcYQQ/rrnWFVfA==}
|
||||
engines: {node: '>= 20'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
'@tailwindcss/oxide-wasm32-wasi@4.1.11':
|
||||
resolution: {integrity: sha512-Xo1+/GU0JEN/C/dvcammKHzeM6NqKovG+6921MR6oadee5XPBaKOumrJCXvopJ/Qb5TH7LX/UAywbqrP4lax0g==}
|
||||
@@ -1954,6 +1942,7 @@ packages:
|
||||
|
||||
'@ungap/structured-clone@1.3.0':
|
||||
resolution: {integrity: sha512-WmoN8qaIAo7WTYWbAZuG8PYEhn5fkz7dZrqTBZ7dtt//lL2Gwms1IcnQ5yHqjDfX8Ft5j4YzDM23f87zBfDe9g==}
|
||||
deprecated: Potential CWE-502 - Update to 1.3.1 or higher
|
||||
|
||||
'@vitejs/plugin-react@5.2.0':
|
||||
resolution: {integrity: sha512-YmKkfhOAi3wsB1PhJq5Scj3GXMn3WvtQ/JC0xoopuHoXSdmtdStOpFrYaT1kie2YgFBcIe64ROzMYRjCrYOdYw==}
|
||||
@@ -3077,56 +3066,48 @@ packages:
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
lightningcss-linux-arm64-gnu@1.32.0:
|
||||
resolution: {integrity: sha512-0nnMyoyOLRJXfbMOilaSRcLH3Jw5z9HDNGfT/gwCPgaDjnx0i8w7vBzFLFR1f6CMLKF8gVbebmkUN3fa/kQJpQ==}
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
lightningcss-linux-arm64-musl@1.30.1:
|
||||
resolution: {integrity: sha512-jmUQVx4331m6LIX+0wUhBbmMX7TCfjF5FoOH6SD1CttzuYlGNVpA7QnrmLxrsub43ClTINfGSYyHe2HWeLl5CQ==}
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
lightningcss-linux-arm64-musl@1.32.0:
|
||||
resolution: {integrity: sha512-UpQkoenr4UJEzgVIYpI80lDFvRmPVg6oqboNHfoH4CQIfNA+HOrZ7Mo7KZP02dC6LjghPQJeBsvXhJod/wnIBg==}
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [arm64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
lightningcss-linux-x64-gnu@1.30.1:
|
||||
resolution: {integrity: sha512-piWx3z4wN8J8z3+O5kO74+yr6ze/dKmPnI7vLqfSqI8bccaTGY5xiSGVIJBDd5K5BHlvVLpUB3S2YCfelyJ1bw==}
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
lightningcss-linux-x64-gnu@1.32.0:
|
||||
resolution: {integrity: sha512-V7Qr52IhZmdKPVr+Vtw8o+WLsQJYCTd8loIfpDaMRWGUZfBOYEJeyJIkqGIDMZPwPx24pUMfwSxxI8phr/MbOA==}
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [glibc]
|
||||
|
||||
lightningcss-linux-x64-musl@1.30.1:
|
||||
resolution: {integrity: sha512-rRomAK7eIkL+tHY0YPxbc5Dra2gXlI63HL+v1Pdi1a3sC+tJTcFrHX+E86sulgAXeI7rSzDYhPSeHHjqFhqfeQ==}
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
lightningcss-linux-x64-musl@1.32.0:
|
||||
resolution: {integrity: sha512-bYcLp+Vb0awsiXg/80uCRezCYHNg1/l3mt0gzHnWV9XP1W5sKa5/TCdGWaR/zBM2PeF/HbsQv/j2URNOiVuxWg==}
|
||||
engines: {node: '>= 12.0.0'}
|
||||
cpu: [x64]
|
||||
os: [linux]
|
||||
libc: [musl]
|
||||
|
||||
lightningcss-win32-arm64-msvc@1.30.1:
|
||||
resolution: {integrity: sha512-mSL4rqPi4iXq5YVqzSsJgMVFENoa4nGTT/GjO2c0Yl9OuQfPsIfncvLrEW6RbbB24WtZ3xP/2CCmI3tNkNV4oA==}
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
allowBuilds:
|
||||
'@tailwindcss/oxide': true
|
||||
@@ -129,6 +129,12 @@ export const getNodeReleases = (channel: ReleaseChannel = "stable") =>
|
||||
Network.post<NodeReleaseApiItem[]>("/node/releases", { channel });
|
||||
export const rollbackNode = (id: number) =>
|
||||
Network.post("/node/rollback", { id });
|
||||
export const testNodeNftables = (nodeId: number) =>
|
||||
Network.post("/node/nftables/test", { nodeId });
|
||||
export const reconcileNodeNftables = (nodeId: number) =>
|
||||
Network.post("/node/nftables/reconcile", { nodeId });
|
||||
export const clearNodeNftables = (nodeId: number) =>
|
||||
Network.post("/node/nftables/clear", { nodeId });
|
||||
|
||||
// 隧道CRUD操作 - 全部使用POST请求
|
||||
export const createTunnel = (data: TunnelMutationPayload) =>
|
||||
|
||||
@@ -2,6 +2,8 @@ export interface NodeApiItem {
|
||||
id: number;
|
||||
name: string;
|
||||
status: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
sshConfig?: NodeSshConfigApiItem | null;
|
||||
inx?: number;
|
||||
remark?: string;
|
||||
expiryTime?: number;
|
||||
@@ -44,6 +46,7 @@ export interface TunnelApiItem {
|
||||
name: string;
|
||||
type: number;
|
||||
status: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
flow?: number;
|
||||
trafficRatio?: number;
|
||||
inIp?: string;
|
||||
@@ -63,6 +66,7 @@ export interface ForwardApiItem {
|
||||
id: number;
|
||||
name: string;
|
||||
status: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
tunnelName?: string;
|
||||
tunnelTrafficRatio?: number;
|
||||
inIp?: string;
|
||||
@@ -303,6 +307,8 @@ export interface NodeMutationPayload {
|
||||
id?: number | null;
|
||||
name?: string;
|
||||
status?: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
sshConfig?: NodeSshConfigMutationPayload | null;
|
||||
inx?: number;
|
||||
remark?: string;
|
||||
expiryTime?: number;
|
||||
@@ -320,6 +326,29 @@ export interface NodeMutationPayload {
|
||||
socks?: number;
|
||||
}
|
||||
|
||||
export interface NodeSshConfigApiItem {
|
||||
host?: string;
|
||||
port?: number;
|
||||
username?: string;
|
||||
authType?: "password" | "private_key";
|
||||
password?: string;
|
||||
privateKey?: string;
|
||||
passphrase?: string;
|
||||
sudoMode?: "none" | "sudo" | "sudo_su";
|
||||
[key: string]: unknown;
|
||||
}
|
||||
|
||||
export interface NodeSshConfigMutationPayload {
|
||||
host?: string;
|
||||
port?: number;
|
||||
username?: string;
|
||||
authType?: "password" | "private_key";
|
||||
password?: string;
|
||||
privateKey?: string;
|
||||
passphrase?: string;
|
||||
sudoMode?: "none" | "sudo" | "sudo_su";
|
||||
}
|
||||
|
||||
export interface TunnelChainNodePayload {
|
||||
nodeId: number;
|
||||
protocol?: string;
|
||||
@@ -334,6 +363,7 @@ export interface TunnelMutationPayload {
|
||||
name?: string;
|
||||
type?: number;
|
||||
status?: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
flow?: number;
|
||||
trafficRatio?: number;
|
||||
inIp?: string;
|
||||
@@ -381,6 +411,7 @@ export interface ForwardMutationPayload {
|
||||
id?: number;
|
||||
name?: string;
|
||||
status?: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
tunnelId?: number | null;
|
||||
inIp?: string;
|
||||
inPort?: number | null;
|
||||
|
||||
@@ -106,6 +106,7 @@ import { JwtUtil } from "@/utils/jwt";
|
||||
interface Forward {
|
||||
id: number;
|
||||
name: string;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
tunnelId: number;
|
||||
tunnelName: string;
|
||||
tunnelTrafficRatio?: number;
|
||||
@@ -134,6 +135,7 @@ interface Forward {
|
||||
interface Tunnel {
|
||||
id: number;
|
||||
name: string;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
type?: number;
|
||||
inIp?: string;
|
||||
inNodeId?: Array<{ nodeId: number }>;
|
||||
|
||||
@@ -109,6 +109,17 @@ interface Node {
|
||||
socks?: number; // 0 关 1 开
|
||||
status: number;
|
||||
isRemote?: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
sshConfig?: {
|
||||
host?: string;
|
||||
port?: number;
|
||||
username?: string;
|
||||
authType?: "password" | "private_key";
|
||||
password?: string;
|
||||
privateKey?: string;
|
||||
passphrase?: string;
|
||||
sudoMode?: "none" | "sudo" | "sudo_su";
|
||||
} | null;
|
||||
remoteUrl?: string;
|
||||
syncError?: string;
|
||||
connectionStatus: "online" | "offline";
|
||||
@@ -132,6 +143,15 @@ interface NodeForm {
|
||||
udpListenAddr: string;
|
||||
interfaceName: string;
|
||||
extraIPs: string;
|
||||
forwardMode: "agent" | "nftables";
|
||||
sshHost: string;
|
||||
sshPort: string;
|
||||
sshUsername: string;
|
||||
sshAuthType: "password" | "private_key";
|
||||
sshPassword: string;
|
||||
sshPrivateKey: string;
|
||||
sshPassphrase: string;
|
||||
sshSudoMode: "none" | "sudo" | "sudo_su";
|
||||
http: number; // 0 关 1 开
|
||||
tls: number; // 0 关 1 开
|
||||
socks: number; // 0 关 1 开
|
||||
@@ -344,11 +364,25 @@ export default function NodePage() {
|
||||
udpListenAddr: "[::]",
|
||||
interfaceName: "",
|
||||
extraIPs: "",
|
||||
forwardMode: "agent",
|
||||
sshHost: "",
|
||||
sshPort: "22",
|
||||
sshUsername: "",
|
||||
sshAuthType: "private_key",
|
||||
sshPassword: "",
|
||||
sshPrivateKey: "",
|
||||
sshPassphrase: "",
|
||||
sshSudoMode: "none",
|
||||
http: 0,
|
||||
tls: 0,
|
||||
socks: 0,
|
||||
});
|
||||
const [errors, setErrors] = useState<Record<string, string>>({});
|
||||
const isNftablesMode = form.forwardMode === "nftables";
|
||||
const protocolControlsDisabled = protocolDisabled || isNftablesMode;
|
||||
const protocolControlsReason = isNftablesMode
|
||||
? "nftables 模式不支持 agent 协议开关"
|
||||
: protocolDisabledReason || "等待节点上线后再设置";
|
||||
|
||||
const [selectMode, setSelectMode] = useState(false);
|
||||
const [selectedIds, setSelectedIds] = useState<Set<number>>(new Set());
|
||||
@@ -856,6 +890,27 @@ export default function NodePage() {
|
||||
newErrors.port = portValidation.error || "端口格式错误";
|
||||
}
|
||||
|
||||
if (form.forwardMode === "nftables") {
|
||||
if (!form.sshHost.trim()) {
|
||||
newErrors.sshHost = "请输入 SSH 主机";
|
||||
}
|
||||
const sshPort = Number(form.sshPort);
|
||||
|
||||
if (!Number.isInteger(sshPort) || sshPort < 1 || sshPort > 65535) {
|
||||
newErrors.sshPort = "SSH 端口必须为 1-65535";
|
||||
}
|
||||
if (!form.sshUsername.trim()) {
|
||||
newErrors.sshUsername = "请输入 SSH 用户名";
|
||||
}
|
||||
if (form.sshAuthType === "password") {
|
||||
if (!isEdit && !form.sshPassword.trim()) {
|
||||
newErrors.sshPassword = "请输入 SSH 密码";
|
||||
}
|
||||
} else if (!isEdit && !form.sshPrivateKey.trim()) {
|
||||
newErrors.sshPrivateKey = "请输入 SSH 私钥";
|
||||
}
|
||||
}
|
||||
|
||||
setErrors(newErrors);
|
||||
|
||||
return Object.keys(newErrors).length === 0;
|
||||
@@ -898,15 +953,28 @@ export default function NodePage() {
|
||||
udpListenAddr: node.udpListenAddr || "[::]",
|
||||
interfaceName: (node as any).interfaceName || "",
|
||||
extraIPs: node.extraIPs || "",
|
||||
forwardMode: node.forwardMode === "nftables" ? "nftables" : "agent",
|
||||
sshHost: node.sshConfig?.host || normalizedHost || legacy,
|
||||
sshPort: String(node.sshConfig?.port || 22),
|
||||
sshUsername: node.sshConfig?.username || "",
|
||||
sshAuthType: node.sshConfig?.authType || "private_key",
|
||||
sshPassword: node.sshConfig?.password || "",
|
||||
sshPrivateKey: node.sshConfig?.privateKey || "",
|
||||
sshPassphrase: node.sshConfig?.passphrase || "",
|
||||
sshSudoMode: node.sshConfig?.sudoMode || "none",
|
||||
http: typeof node.http === "number" ? node.http : 1,
|
||||
tls: typeof node.tls === "number" ? node.tls : 1,
|
||||
socks: typeof node.socks === "number" ? node.socks : 1,
|
||||
});
|
||||
const offline = node.connectionStatus !== "online";
|
||||
|
||||
setProtocolDisabled(offline);
|
||||
setProtocolDisabled(offline || node.forwardMode === "nftables");
|
||||
setProtocolDisabledReason(
|
||||
offline ? "节点未在线,等待节点上线后再设置" : "",
|
||||
node.forwardMode === "nftables"
|
||||
? "nftables 模式不支持 agent 协议开关"
|
||||
: offline
|
||||
? "节点未在线,等待节点上线后再设置"
|
||||
: "",
|
||||
);
|
||||
setDialogVisible(true);
|
||||
};
|
||||
@@ -1166,12 +1234,33 @@ export default function NodePage() {
|
||||
try {
|
||||
const apiCall = isEdit ? updateNode : createNode;
|
||||
const { serverHost, ...rest } = form;
|
||||
const sshConfig =
|
||||
form.forwardMode === "nftables"
|
||||
? {
|
||||
host: form.sshHost.trim(),
|
||||
port: Number(form.sshPort || 22),
|
||||
username: form.sshUsername.trim(),
|
||||
authType: form.sshAuthType,
|
||||
password:
|
||||
form.sshAuthType === "password"
|
||||
? form.sshPassword.trim()
|
||||
: undefined,
|
||||
privateKey:
|
||||
form.sshAuthType === "private_key"
|
||||
? form.sshPrivateKey
|
||||
: undefined,
|
||||
passphrase: form.sshPassphrase,
|
||||
sudoMode: form.sshSudoMode,
|
||||
}
|
||||
: null;
|
||||
const data = {
|
||||
...rest,
|
||||
remark: form.remark.trim(),
|
||||
expiryTime: form.expiryTime,
|
||||
renewalCycle: form.renewalCycle,
|
||||
extraIPs: form.extraIPs,
|
||||
forwardMode: form.forwardMode,
|
||||
sshConfig,
|
||||
serverIp:
|
||||
form.serverIpV4?.trim() ||
|
||||
form.serverIpV6?.trim() ||
|
||||
@@ -1206,6 +1295,8 @@ export default function NodePage() {
|
||||
tcpListenAddr: form.tcpListenAddr,
|
||||
udpListenAddr: form.udpListenAddr,
|
||||
interfaceName: form.interfaceName,
|
||||
forwardMode: form.forwardMode,
|
||||
sshConfig,
|
||||
http: form.http,
|
||||
tls: form.tls,
|
||||
socks: form.socks,
|
||||
@@ -1242,6 +1333,15 @@ export default function NodePage() {
|
||||
udpListenAddr: "[::]",
|
||||
interfaceName: "",
|
||||
extraIPs: "",
|
||||
forwardMode: "agent",
|
||||
sshHost: "",
|
||||
sshPort: "22",
|
||||
sshUsername: "",
|
||||
sshAuthType: "password",
|
||||
sshPassword: "",
|
||||
sshPrivateKey: "",
|
||||
sshPassphrase: "",
|
||||
sshSudoMode: "none",
|
||||
http: 0,
|
||||
tls: 0,
|
||||
socks: 0,
|
||||
@@ -2581,6 +2681,43 @@ export default function NodePage() {
|
||||
}
|
||||
/>
|
||||
|
||||
<Select
|
||||
label="转发模式"
|
||||
selectedKeys={[form.forwardMode]}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selected = Array.from(keys)[0] as
|
||||
| NodeForm["forwardMode"]
|
||||
| undefined;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
forwardMode: selected || "agent",
|
||||
}));
|
||||
}}
|
||||
>
|
||||
<SelectItem key="agent" textValue="agent">
|
||||
agent 模式
|
||||
</SelectItem>
|
||||
<SelectItem key="nftables" textValue="nftables">
|
||||
nftables 模式
|
||||
</SelectItem>
|
||||
</Select>
|
||||
|
||||
{form.forwardMode === "nftables" ? (
|
||||
<Alert
|
||||
color="warning"
|
||||
description="nftables 模式仅支持纯转发,不支持隧道、流量控制和 agent 安装,规则将由面板通过 SSH 下发。"
|
||||
variant="flat"
|
||||
/>
|
||||
) : (
|
||||
<Alert
|
||||
color="primary"
|
||||
description="agent 模式会继续使用节点 agent 执行转发与管理操作。"
|
||||
variant="flat"
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* 高级配置 */}
|
||||
<Accordion className="px-0" variant="light">
|
||||
<AccordionItem
|
||||
@@ -2594,6 +2731,177 @@ export default function NodePage() {
|
||||
}
|
||||
>
|
||||
<div className="space-y-4 pb-2">
|
||||
{form.forwardMode === "nftables" ? (
|
||||
<div className="space-y-4 rounded-xl border border-warning-200 bg-warning-50/60 p-4 dark:border-warning-500/30 dark:bg-warning-950/20">
|
||||
<div>
|
||||
<div className="text-sm font-medium text-warning-700 dark:text-warning-300">
|
||||
SSH 配置
|
||||
</div>
|
||||
<div className="text-xs text-default-500">
|
||||
面板会通过 SSH 下发 nftables 规则,不会安装 agent。
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="grid grid-cols-1 md:grid-cols-2 gap-4">
|
||||
<Input
|
||||
errorMessage={errors.sshHost}
|
||||
isInvalid={!!errors.sshHost}
|
||||
label="SSH 主机"
|
||||
placeholder="例如: 192.0.2.10"
|
||||
value={form.sshHost}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshHost: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
|
||||
<Input
|
||||
errorMessage={errors.sshPort}
|
||||
isInvalid={!!errors.sshPort}
|
||||
label="SSH 端口"
|
||||
placeholder="22"
|
||||
value={form.sshPort}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshPort: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div className="grid grid-cols-1 md:grid-cols-2 gap-4">
|
||||
<Input
|
||||
errorMessage={errors.sshUsername}
|
||||
isInvalid={!!errors.sshUsername}
|
||||
label="SSH 用户名"
|
||||
placeholder="例如: root"
|
||||
value={form.sshUsername}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshUsername: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
|
||||
<Select
|
||||
label="SSH 认证方式"
|
||||
selectedKeys={[form.sshAuthType]}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selected = Array.from(keys)[0] as
|
||||
| "password"
|
||||
| "private_key"
|
||||
| undefined;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshAuthType: selected || "private_key",
|
||||
}));
|
||||
}}
|
||||
>
|
||||
<SelectItem
|
||||
key="private_key"
|
||||
textValue="private_key"
|
||||
>
|
||||
私钥
|
||||
</SelectItem>
|
||||
<SelectItem key="password" textValue="password">
|
||||
密码
|
||||
</SelectItem>
|
||||
</Select>
|
||||
</div>
|
||||
|
||||
{form.sshAuthType === "password" ? (
|
||||
<Input
|
||||
errorMessage={errors.sshPassword}
|
||||
isInvalid={!!errors.sshPassword}
|
||||
label="SSH 密码"
|
||||
placeholder="请输入 SSH 密码"
|
||||
type="password"
|
||||
value={form.sshPassword}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshPassword: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
) : (
|
||||
<Textarea
|
||||
errorMessage={errors.sshPrivateKey}
|
||||
isInvalid={!!errors.sshPrivateKey}
|
||||
label="SSH 私钥"
|
||||
minRows={6}
|
||||
placeholder="粘贴 PEM 格式私钥内容"
|
||||
value={form.sshPrivateKey}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshPrivateKey: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
)}
|
||||
|
||||
<Input
|
||||
label="SSH 私钥密码短语"
|
||||
placeholder="可选"
|
||||
type="password"
|
||||
value={form.sshPassphrase}
|
||||
variant="bordered"
|
||||
onChange={(e) =>
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshPassphrase: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
|
||||
<Select
|
||||
label="sudo 提权方式"
|
||||
selectedKeys={[form.sshSudoMode]}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selected = Array.from(keys)[0] as
|
||||
| "none"
|
||||
| "sudo"
|
||||
| "sudo_su"
|
||||
| undefined;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
sshSudoMode: selected || "none",
|
||||
}));
|
||||
}}
|
||||
>
|
||||
<SelectItem key="none" textValue="none">
|
||||
无需提权
|
||||
</SelectItem>
|
||||
<SelectItem key="sudo" textValue="sudo">
|
||||
sudo
|
||||
</SelectItem>
|
||||
<SelectItem key="sudo_su" textValue="sudo_su">
|
||||
sudo su
|
||||
</SelectItem>
|
||||
</Select>
|
||||
</div>
|
||||
) : (
|
||||
<Alert
|
||||
color="primary"
|
||||
description="agent 模式下无需填写 SSH 配置;节点安装 agent 后会自行接管转发。"
|
||||
variant="flat"
|
||||
/>
|
||||
)}
|
||||
|
||||
<Input
|
||||
description="用于多IP服务器指定使用那个IP请求远程地址,不懂的默认为空就行"
|
||||
errorMessage={errors.interfaceName}
|
||||
@@ -2677,18 +2985,16 @@ export default function NodePage() {
|
||||
<div className="text-xs text-default-500 mb-2">
|
||||
开启开关以屏蔽对应协议
|
||||
</div>
|
||||
{protocolDisabled && (
|
||||
{protocolControlsDisabled && (
|
||||
<Alert
|
||||
className="mb-2"
|
||||
color="warning"
|
||||
description={
|
||||
protocolDisabledReason || "等待节点上线后再设置"
|
||||
}
|
||||
description={protocolControlsReason}
|
||||
variant="flat"
|
||||
/>
|
||||
)}
|
||||
<div
|
||||
className={`grid grid-cols-1 sm:grid-cols-3 gap-3 bg-content1/30 dark:bg-content1/20 p-3 rounded-md border border-divider ${protocolDisabled ? "opacity-70" : ""}`}
|
||||
className={`grid grid-cols-1 sm:grid-cols-3 gap-3 bg-content1/30 dark:bg-content1/20 p-3 rounded-md border border-divider ${protocolControlsDisabled ? "opacity-70" : ""}`}
|
||||
>
|
||||
{/* HTTP tile */}
|
||||
<div className="px-3 py-3 rounded-lg bg-content1/55 dark:bg-content1/35 border border-divider hover:border-primary-200 dark:hover:border-primary-500/30 transition-colors">
|
||||
@@ -2715,7 +3021,7 @@ export default function NodePage() {
|
||||
禁用/启用
|
||||
</div>
|
||||
<Switch
|
||||
isDisabled={protocolDisabled}
|
||||
isDisabled={protocolControlsDisabled}
|
||||
isSelected={form.http === 1}
|
||||
size="sm"
|
||||
onValueChange={(v) =>
|
||||
@@ -2762,7 +3068,7 @@ export default function NodePage() {
|
||||
禁用/启用
|
||||
</div>
|
||||
<Switch
|
||||
isDisabled={protocolDisabled}
|
||||
isDisabled={protocolControlsDisabled}
|
||||
isSelected={form.tls === 1}
|
||||
size="sm"
|
||||
onValueChange={(v) =>
|
||||
@@ -2801,7 +3107,7 @@ export default function NodePage() {
|
||||
禁用/启用
|
||||
</div>
|
||||
<Switch
|
||||
isDisabled={protocolDisabled}
|
||||
isDisabled={protocolControlsDisabled}
|
||||
isSelected={form.socks === 1}
|
||||
size="sm"
|
||||
onValueChange={(v) =>
|
||||
|
||||
@@ -122,6 +122,7 @@ interface Tunnel {
|
||||
inx?: number;
|
||||
name: string;
|
||||
type: number; // 1: 端口转发, 2: 隧道转发
|
||||
forwardMode?: "agent" | "nftables";
|
||||
inNodeId: ChainTunnel[]; // 入口节点列表
|
||||
outNodeId?: ChainTunnel[]; // 出口节点列表
|
||||
chainNodes?: ChainTunnel[][]; // 转发链节点列表,二维数组
|
||||
@@ -150,6 +151,7 @@ interface Node {
|
||||
id: number;
|
||||
name: string;
|
||||
status: number; // 1: 在线, 0: 离线
|
||||
forwardMode?: "agent" | "nftables";
|
||||
serverIp?: string;
|
||||
serverIpV4?: string;
|
||||
serverIpV6?: string;
|
||||
@@ -160,6 +162,7 @@ interface TunnelForm {
|
||||
id?: number;
|
||||
name: string;
|
||||
type: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
inNodeId: ChainTunnel[];
|
||||
outNodeId?: ChainTunnel[];
|
||||
chainNodes?: ChainTunnel[][]; // 转发链节点列表,二维数组,外层是跳数,内层是该跳的节点
|
||||
@@ -172,6 +175,23 @@ interface TunnelForm {
|
||||
status: number;
|
||||
}
|
||||
|
||||
type TunnelForwardMode = NonNullable<TunnelForm["forwardMode"]>;
|
||||
|
||||
const getTunnelForwardMode = (
|
||||
inNodeId: ChainTunnel[],
|
||||
nodes: Node[],
|
||||
): TunnelForwardMode =>
|
||||
inNodeId.some((item) => {
|
||||
const node = nodes.find((candidate) => candidate.id === item.nodeId);
|
||||
|
||||
return node?.forwardMode === "nftables";
|
||||
})
|
||||
? "nftables"
|
||||
: "agent";
|
||||
|
||||
const createTypedTunnelFormDefaults = (): TunnelForm =>
|
||||
createTunnelFormDefaults() as TunnelForm;
|
||||
|
||||
interface BatchProgressState {
|
||||
active: boolean;
|
||||
label: string;
|
||||
@@ -416,7 +436,7 @@ export default function TunnelPage() {
|
||||
};
|
||||
|
||||
// 表单状态
|
||||
const [form, setForm] = useState<TunnelForm>(createTunnelFormDefaults());
|
||||
const [form, setForm] = useState<TunnelForm>(createTypedTunnelFormDefaults());
|
||||
|
||||
// 表单验证错误
|
||||
const [errors, setErrors] = useState<{ [key: string]: string }>({});
|
||||
@@ -564,7 +584,14 @@ export default function TunnelPage() {
|
||||
|
||||
// 表单验证
|
||||
const validateForm = (): boolean => {
|
||||
const newErrors = validateTunnelForm(form, nodes, isEdit);
|
||||
const newErrors = validateTunnelForm(
|
||||
{
|
||||
...form,
|
||||
forwardMode: getTunnelForwardMode(form.inNodeId, nodes),
|
||||
},
|
||||
nodes,
|
||||
isEdit,
|
||||
);
|
||||
|
||||
setErrors(newErrors);
|
||||
|
||||
@@ -574,7 +601,7 @@ export default function TunnelPage() {
|
||||
// 新增隧道
|
||||
const handleAdd = () => {
|
||||
setIsEdit(false);
|
||||
setForm(createTunnelFormDefaults());
|
||||
setForm(createTypedTunnelFormDefaults());
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
};
|
||||
@@ -588,6 +615,7 @@ export default function TunnelPage() {
|
||||
id: tunnel.id,
|
||||
name: tunnel.name,
|
||||
type: tunnel.type,
|
||||
forwardMode: tunnel.forwardMode === "nftables" ? "nftables" : "agent",
|
||||
inNodeId: tunnel.inNodeId || [],
|
||||
outNodeId: tunnel.outNodeId || [],
|
||||
chainNodes: tunnel.chainNodes || [],
|
||||
@@ -885,6 +913,7 @@ export default function TunnelPage() {
|
||||
|
||||
const data = {
|
||||
...form,
|
||||
forwardMode: getTunnelForwardMode(form.inNodeId, nodes),
|
||||
inIp: inIpString,
|
||||
outNodeId: cleanedOutNodeId,
|
||||
chainNodes: cleanedChainNodes,
|
||||
|
||||
@@ -8,6 +8,7 @@ interface TunnelFormInput {
|
||||
inNodeId: TunnelChainNode[];
|
||||
outNodeId?: TunnelChainNode[];
|
||||
trafficRatio: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
probeTargetHost?: string;
|
||||
probeTargetPort?: number;
|
||||
}
|
||||
@@ -15,6 +16,7 @@ interface TunnelFormInput {
|
||||
interface TunnelNodeInput {
|
||||
id: number;
|
||||
status: number;
|
||||
forwardMode?: "agent" | "nftables";
|
||||
}
|
||||
|
||||
const isValidProbeIPv4 = (host: string) => {
|
||||
@@ -116,6 +118,7 @@ export const createTunnelFormDefaults = () => {
|
||||
chainNodes: [],
|
||||
flow: 1,
|
||||
trafficRatio: 1.0,
|
||||
forwardMode: "agent",
|
||||
inIp: "",
|
||||
ipPreference: "",
|
||||
probeTargetHost: "",
|
||||
@@ -183,6 +186,10 @@ export const validateTunnelForm = (
|
||||
}
|
||||
|
||||
if (form.type === 2) {
|
||||
if (form.forwardMode === "nftables") {
|
||||
errors.type = "nftables 节点仅支持端口转发";
|
||||
}
|
||||
|
||||
if (!form.outNodeId || form.outNodeId.length === 0) {
|
||||
errors.outNodeId = "请至少选择一个出口节点";
|
||||
} else {
|
||||
@@ -208,6 +215,23 @@ export const validateTunnelForm = (
|
||||
}
|
||||
}
|
||||
|
||||
if (form.forwardMode === "nftables") {
|
||||
const nftEntryNodes = (form.inNodeId || []).filter((item) => {
|
||||
const node = nodes.find((n) => n.id === item.nodeId);
|
||||
|
||||
return node?.forwardMode === "nftables";
|
||||
});
|
||||
|
||||
if (nftEntryNodes.length > 0) {
|
||||
if (form.inNodeId.length !== 1) {
|
||||
errors.inNodeId = "nftables 节点仅支持单入口隧道";
|
||||
}
|
||||
if ((form.outNodeId || []).length > 0) {
|
||||
errors.outNodeId = "nftables 节点不支持出口节点配置";
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return errors;
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user