mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 15:46:38 +08:00
Compare commits
140 Commits
3.0.0-alpha13
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 129fa0aa4c | |||
| cf7246c71b | |||
| b4c2989285 | |||
| f2783713a5 | |||
| 5d60c4fbe1 | |||
| 62c56e0a79 | |||
| 2e845030de | |||
| 2269f2e2d5 | |||
| da7bef88f1 | |||
| f26014579b | |||
| 5041c722c9 | |||
| c56798e991 | |||
| 40e96f3592 | |||
| 0e24b53a5b | |||
| a820c49c94 | |||
| 538e64ffc0 | |||
| 0b23d6f7d7 | |||
| 9e6f80019d | |||
| a8fd01d4d8 | |||
| cbe2fc492e | |||
| ae370382d3 | |||
| e112d81697 | |||
| 11a27d3c67 | |||
| 8e513a1bae | |||
| f98be845d3 | |||
| 0c2acfdd8a | |||
| 777db8767f | |||
| 82f6047506 | |||
| 3ce320da5a | |||
| 7ab0db29ae | |||
| 006ea97200 | |||
| e569aedd3e | |||
| 35080aea2d | |||
| 85e588ffe9 | |||
| 03524f4a65 | |||
| 9e69e020ab | |||
| 14bbd3907d | |||
| ca8d8e92ba | |||
| 079474fa06 | |||
| 6e249a54f4 | |||
| fb798a4532 | |||
| a599f383f5 | |||
| 6bfa7f0166 | |||
| 2d0c993c90 | |||
| 8b64542c94 | |||
| 032b0f0cfd | |||
| 7008717a49 | |||
| 8552a70355 | |||
| b2454e86c9 | |||
| 312c9a9c5c | |||
| e1324b8c8c | |||
| c034d0d41f | |||
| fd3ecc38ef | |||
| 2eee506716 | |||
| abf13bdac9 | |||
| f0facf6703 | |||
| 583a3834f9 | |||
| bd13477fa3 | |||
| 106a30bf9d | |||
| 465815cf34 | |||
| ec9fb77eb5 | |||
| 7d63dd4cc3 | |||
| 4cfa6adee7 | |||
| bc8f2ec8a1 | |||
| 9a37c2f603 | |||
| ff6d46ddaf | |||
| 723534faea | |||
| 73490a9be6 | |||
| 5a327459f7 | |||
| f307e7d5eb | |||
| fdd72979b6 | |||
| ad33791a26 | |||
| 6320b1f0c1 | |||
| 91d79b6b3a | |||
| fc7df6bd64 | |||
| 25dfb84324 | |||
| 4ebd6703fe | |||
| 1f53a39784 | |||
| 5ebd4c2a91 | |||
| 6c93d829c6 | |||
| 5d22d4cb06 | |||
| e5cd5af550 | |||
| cdcdfd8ff0 | |||
| 791773fd62 | |||
| 13764b4615 | |||
| 4c882d907b | |||
| 6033e39466 | |||
| 0f3242bf11 | |||
| d97d91801d | |||
| 727ef56c67 | |||
| a40150b136 | |||
| a923ec4785 | |||
| 42c5492c1d | |||
| 869d726b7a | |||
| 55a931510b | |||
| a259dd83b2 | |||
| cc0b8de2e1 | |||
| 58ef260755 | |||
| 521fe79b15 | |||
| 90012725cc | |||
| 615d9e67eb | |||
| a131b70613 | |||
| cbed4eab23 | |||
| c2745dcd56 | |||
| ad4109594a | |||
| efc8c75dcb | |||
| e5acc49186 | |||
| d98377a297 | |||
| 950e9a9ba8 | |||
| 3f3159aafd | |||
| 5e8d0682c0 | |||
| 3c0e833cfc | |||
| 60311d3e47 | |||
| b4192c9e94 | |||
| f05e9480ee | |||
| 9861b44107 | |||
| d8144821e6 | |||
| 023be27287 | |||
| e8d5687419 | |||
| 0fbe570597 | |||
| 4c52d7fec2 | |||
| 3373e5ade9 | |||
| bd27b94909 | |||
| a5a500bc0f | |||
| edfe2a2372 | |||
| 7a9ba8bd81 | |||
| 46394388b1 | |||
| dec337d46b | |||
| 9e8d27d98e | |||
| a2000e4d98 | |||
| 2b76a9f0be | |||
| 54d7dfb7c9 | |||
| 9f19d5fe15 | |||
| 2ca3849917 | |||
| 58d2e89147 | |||
| 87a1a34ad5 | |||
| a625884d61 | |||
| 799bb66fe5 | |||
| 3f374df724 | |||
| 9b923a2d0b |
@@ -7,6 +7,18 @@ on:
|
||||
branches: ['**']
|
||||
|
||||
jobs:
|
||||
install-scripts:
|
||||
name: Test Installer Scripts (systemd and Alpine OpenRC)
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Run installer regression tests
|
||||
run: bash test-install-scripts-proxy.sh
|
||||
|
||||
- name: Test Alpine bootstrap and OpenRC lifecycle
|
||||
run: docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh
|
||||
|
||||
frontend:
|
||||
name: Build Frontend
|
||||
runs-on: ubuntu-latest
|
||||
@@ -22,7 +34,7 @@ jobs:
|
||||
node-version: '20.19.0'
|
||||
|
||||
- name: Install pnpm
|
||||
run: npm install -g pnpm
|
||||
run: npm install -g pnpm@10.28.1
|
||||
|
||||
- name: Install dependencies
|
||||
run: pnpm install --frozen-lockfile
|
||||
|
||||
@@ -229,6 +229,7 @@ jobs:
|
||||
|
||||
docker buildx build \
|
||||
--platform linux/amd64,linux/arm64 \
|
||||
--build-arg KEYGEN_ACCOUNT_ID=${{ secrets.KEYGEN_ACCOUNT_ID }} \
|
||||
--push \
|
||||
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:latest \
|
||||
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:${VERSION} \
|
||||
@@ -303,6 +304,9 @@ jobs:
|
||||
sed -i "s|^PINNED_VERSION=\"\"|PINNED_VERSION=\"${VERSION}\"|" ./artifacts/install.sh
|
||||
sed -i "s|^PINNED_VERSION=\"\"|PINNED_VERSION=\"${VERSION}\"|" ./artifacts/panel_install.sh
|
||||
|
||||
- name: Verify release installer on Alpine OpenRC
|
||||
run: docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh /workspace/artifacts/install.sh
|
||||
|
||||
- name: Create Release
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
@@ -431,4 +435,3 @@ jobs:
|
||||
gh release upload "${VERSION}" ./artifacts/gost-arm64.sha256 --clobber
|
||||
|
||||
echo "✅ GOST 二进制文件更新完成"
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# FLVX
|
||||
|
||||
> **联系我们**: [Telegram群组](https://t.me/flvxpanel)
|
||||
> **联系我们**: [Telegram群组](https://t.me/flvxchannel)
|
||||
|
||||
|
||||
## 特性
|
||||
@@ -184,7 +184,6 @@ This fork (FLVX) is no longer a light patch on top of the upstream project. It h
|
||||
|
||||
| 网络 | 地址 |
|
||||
|------------|----------------------------------------------------------------------|
|
||||
| BNB(BEP20) | `0xa608708fdc6279a2433fd4b82f0b72b8cbe97ed5` |
|
||||
| TRC20 | `TM8VYdU3s3gSX5PC8swjAJrAzZFCHKqG2k` |
|
||||
| Aptos | `0x49427bfcba1006a346447430689b2307ac156316bb34850d1d3029ff9d118da5` |
|
||||
| polygon | `0xa608708fdc6279a2433fd4b82f0b72b8cbe97ed5` |
|
||||
| BNB(BEP20) | `0x271327ce49140e670eA0F772d9886BF90E9022Ee` |
|
||||
| TRC20 | `TARxZWggaxFqYgxGVBxPkyykgYKNmGndmE` |
|
||||
| polygon | `0x271327ce49140e670eA0F772d9886BF90E9022Ee` |
|
||||
|
||||
+10
-1
@@ -64,6 +64,14 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/panel_instal
|
||||
curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -o install.sh && chmod +x install.sh && ./install.sh
|
||||
```
|
||||
|
||||
Alpine Linux 最小化安装若未包含 `curl`,可使用系统自带的 `wget` 下载:
|
||||
|
||||
```bash
|
||||
wget -O install.sh https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh && chmod +x install.sh && ./install.sh
|
||||
```
|
||||
|
||||
脚本会在 Alpine 上自动安装 Bash、`curl` 和 CA 证书,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
|
||||
|
||||
**安装过程中会提示输入:**
|
||||
- **服务器地址**: 面板端的通信地址(通常是 `http://<面板IP>:<后端端口>`,例如 `http://1.2.3.4:6365`)。
|
||||
- **密钥**: 刚才在面板中获取的节点密钥。
|
||||
@@ -77,7 +85,8 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -
|
||||
|
||||
### 3. 验证安装
|
||||
安装完成后,服务会自动启动。
|
||||
- 查看状态: `systemctl status flux_agent`
|
||||
- systemd 查看状态: `systemctl status flux_agent`
|
||||
- Alpine/OpenRC 查看状态: `rc-service flux_agent status`
|
||||
- 回到面板 **节点管理** 页面,该节点状态应显示为 **在线**。
|
||||
|
||||
---
|
||||
|
||||
@@ -15,10 +15,15 @@ services:
|
||||
JWT_SECRET: ${JWT_SECRET}
|
||||
SERVER_ADDR: :6365
|
||||
TZ: Asia/Shanghai
|
||||
FLUX_VERSION: ${FLUX_VERSION:-dev}
|
||||
PANEL_DEPLOY_DIR: /opt/flvx-panel
|
||||
PANEL_BACKEND_CONTAINER: flux-panel-backend
|
||||
ports:
|
||||
- "${BACKEND_PORT}:6365"
|
||||
volumes:
|
||||
- sqlite_data:/app/data
|
||||
- /var/run/docker.sock:/var/run/docker.sock
|
||||
- ./:/opt/flvx-panel
|
||||
networks:
|
||||
- gost-network
|
||||
stop_grace_period: 30s
|
||||
|
||||
@@ -15,10 +15,15 @@ services:
|
||||
JWT_SECRET: ${JWT_SECRET}
|
||||
SERVER_ADDR: :6365
|
||||
TZ: Asia/Shanghai
|
||||
FLUX_VERSION: ${FLUX_VERSION:-dev}
|
||||
PANEL_DEPLOY_DIR: /opt/flvx-panel
|
||||
PANEL_BACKEND_CONTAINER: flux-panel-backend
|
||||
ports:
|
||||
- "${BACKEND_PORT}:6365"
|
||||
volumes:
|
||||
- sqlite_data:/app/data
|
||||
- /var/run/docker.sock:/var/run/docker.sock
|
||||
- ./:/opt/flvx-panel
|
||||
networks:
|
||||
- gost-network
|
||||
stop_grace_period: 30s
|
||||
|
||||
@@ -0,0 +1,178 @@
|
||||
# Proxy Protocol 传输分析报告
|
||||
|
||||
**日期**: 2026-05-07
|
||||
**测试环境**: 20.118.172.127 (Server 1) ↔ 108.181.90.137 (Server 2)
|
||||
|
||||
---
|
||||
|
||||
## 1. 代码流程分析
|
||||
|
||||
### 完整数据链路
|
||||
|
||||
```
|
||||
前端 (proxyProtocol: 0|1|2)
|
||||
→ 后端 handler mutations.go:1936
|
||||
→ 数据库存储 forward.proxy_protocol (model.go:50)
|
||||
→ 控制面 buildForwardServiceConfigs (control_plane.go:1791-1796)
|
||||
→ handler metadata: {"proxyProtocol": 2}
|
||||
→ Agent metadata 解析 (metadata.go:42)
|
||||
→ handler.go:256 WrapClientConn()
|
||||
→ conn.go:14 HeaderProxyFromAddrs(byte(ppv), src, dst)
|
||||
→ conn.go:15 header.WriteTo(c)
|
||||
→ 目标服务器收到 PROXY protocol header
|
||||
```
|
||||
|
||||
### 关键代码
|
||||
|
||||
**写入 PROXY header** (`go-gost/x/internal/net/proxyproto/conn.go`):
|
||||
```go
|
||||
func WrapClientConn(ppv int, src, dst net.Addr, c net.Conn) net.Conn {
|
||||
if ppv <= 0 {
|
||||
return c
|
||||
}
|
||||
header := proxyproto.HeaderProxyFromrs(byte(ppv), src, dst)
|
||||
header.WriteTo(c)
|
||||
return c
|
||||
}
|
||||
```
|
||||
|
||||
**Handler 调用** (`go-gost/x/handler/forward/local/handler.go:256`):
|
||||
```go
|
||||
cc = proxyproto.WrapClientConn(h.md.proxyProtocol, conn.RemoteAddr(), conn.LocalAddr(), cc)
|
||||
```
|
||||
|
||||
- `src` = `conn.RemoteAddr()` → 客户端真实 IP ✅
|
||||
- `dst` = `conn.LocalAddr()` → agent 监听地址 ✅
|
||||
- `ppv` = 1 或 2 → 版本号正确 ✅
|
||||
|
||||
---
|
||||
|
||||
## 2. 实际传输测试结果
|
||||
|
||||
### 测试方法
|
||||
|
||||
1. 在 Server 2 启动 Python TCP 监听器,解析 PROXY protocol header
|
||||
2. 在 Server 1 用当前代码编译 gost,配置 `proxyProtocol: 2` 转发到 Server 2
|
||||
3. 通过 `nc` 发送测试数据,验证 Server 2 是否收到正确的 PROXY header
|
||||
|
||||
### 测试结果
|
||||
|
||||
| 版本 | 状态 | 接收到的 Header |
|
||||
|------|------|----------------|
|
||||
| **PPv2** | ✅ 成功 | `PP2 family=1 alen=12 SRC=127.0.0.1:45410 DST=127.0.0.1:20001` |
|
||||
| **PPv1** | ✅ 成功 | `PROXY TCP4 127.0.0.1 127.0.0.1 43816 20001` |
|
||||
|
||||
### 测试详情
|
||||
|
||||
**PPv2 原始数据**:
|
||||
```
|
||||
Got 28 bytes
|
||||
PP2 family=1 alen=12
|
||||
SRC=127.0.0.1:45410 DST=127.0.0.1:20001
|
||||
```
|
||||
|
||||
**PPv1 原始数据** (hex):
|
||||
```
|
||||
50524f58592054435034203132372e302e302e31203132372e302e302e312034333831362032303030310d0a
|
||||
```
|
||||
解码: `PROXY TCP4 127.0.0.1 127.0.0.1 43816 20001`
|
||||
|
||||
---
|
||||
|
||||
## 3. 单元测试结果
|
||||
|
||||
```
|
||||
go-gost/x/handler/forward/local/ → TestLocalForwardHandlerSendsProxyProtocolToTarget ✅
|
||||
go-backend/internal/http/handler/ → TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler ✅
|
||||
go-backend/internal/store/repo/ → TestGetForwardRecordIncludesProxyProtocol ✅
|
||||
```
|
||||
|
||||
全部通过 (3/3)。
|
||||
|
||||
---
|
||||
|
||||
## 4. 发现的问题
|
||||
|
||||
### 问题 1: `WriteTo` 错误未检查
|
||||
|
||||
**位置**: `go-gost/x/internal/net/proxyproto/conn.go:15`
|
||||
|
||||
```go
|
||||
header.WriteTo(c) // 返回 (int64, error) 被忽略
|
||||
```
|
||||
|
||||
**影响**: 如果写入失败(连接已断开、网络错误等),后续数据传输会在没有 PROXY header 的情况下继续,目标服务器可能解析出错。
|
||||
|
||||
**建议**:
|
||||
```go
|
||||
if _, err := header.WriteTo(c); err != nil {
|
||||
return c // 或包装一个带错误的 conn
|
||||
}
|
||||
```
|
||||
|
||||
**严重程度**: 低(实际场景中,写入失败后 `Transport` 也会很快失败)
|
||||
|
||||
---
|
||||
|
||||
### 问题 2: 部署版本过旧
|
||||
|
||||
**服务器状态**:
|
||||
|
||||
| 服务器 | 组件 | 版本 | 状态 |
|
||||
|--------|------|------|------|
|
||||
| 20.118.172.127 | flux_agent | UPX 压缩,无法读取版本 | ✅ 运行中 |
|
||||
| 20.118.172.127 | paneld | `/app/paneld` | ✅ 运行中 |
|
||||
| 20.118.172.127 | /usr/local/bin/gost | v3.0.0 (go1.23.4) | 旧版,不支持 handler metadata 中的 proxyProtocol |
|
||||
| 108.181.90.137 | flux_agent | 8.8MB | ✅ 运行中 |
|
||||
|
||||
**影响**: 旧版 gost 二进制不识别 handler metadata 中的 `proxyProtocol` 字段,PROXY protocol 功能在生产环境不可用。
|
||||
|
||||
**验证**: 用旧版 gost 测试时,目标服务器收到的原始数据为空,无 PROXY header。
|
||||
|
||||
---
|
||||
|
||||
### 问题 3: 数据库 Schema 缺失
|
||||
|
||||
**位置**: 20.118.172.127 的 `/app/data/gost.db`
|
||||
|
||||
**当前 forward 表 schema**:
|
||||
```sql
|
||||
CREATE TABLE `forward` (
|
||||
`id` integer PRIMARY KEY AUTOINCREMENT,
|
||||
`user_id` integer NOT NULL,
|
||||
`user_name` varchar(100) NOT NULL,
|
||||
`name` varchar(100) NOT NULL,
|
||||
`tunnel_id` integer NOT NULL,
|
||||
`remote_addr` text NOT NULL,
|
||||
`strategy` varchar(100) NOT NULL DEFAULT "fifo",
|
||||
`in_flow` integer NOT NULL DEFAULT 0,
|
||||
`out_flow` integer NOT NULL DEFAULT 0,
|
||||
`created_time` integer NOT NULL,
|
||||
`updated_time` integer NOT NULL,
|
||||
`status` integer NOT NULL,
|
||||
`inx` integer NOT NULL DEFAULT 0,
|
||||
`speed_id` integer
|
||||
);
|
||||
```
|
||||
|
||||
**缺失字段**:
|
||||
- `proxy_protocol` — PROXY protocol 版本
|
||||
- `max_conn` — 最大连接数
|
||||
- `ip_max_conn` — 每 IP 最大连接数
|
||||
- `ip_speed_id` — 每 IP 限速 ID
|
||||
|
||||
**影响**: 后端无法存储和读取 proxy_protocol 配置,前端设置不会生效。
|
||||
|
||||
---
|
||||
|
||||
## 5. 结论
|
||||
|
||||
| 维度 | 状态 | 说明 |
|
||||
|------|------|------|
|
||||
| **代码实现** | ✅ 正确 | 完整的写入链路,版本/地址正确 |
|
||||
| **单元测试** | ✅ 通过 | 3/3 测试覆盖 handler、repo、控制面 |
|
||||
| **实际传输 (新编译版)** | ✅ 成功 | PPv1 和 PPv2 均正确传输 |
|
||||
| **实际传输 (部署版)** | ❌ 不工作 | 旧版不支持 handler metadata 中的 proxyProtocol |
|
||||
| **数据库 Schema** | ❌ 缺字段 | 需要迁移添加 proxy_protocol 等列 |
|
||||
|
||||
**总结**: 代码实现正确,PROXY protocol 传输逻辑无误。但生产服务器运行的是旧版本,需要升级 backend 和 agent 才能启用此功能。
|
||||
@@ -0,0 +1,171 @@
|
||||
# Allow Local Remote Address Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Add a global settings toggle that allows non-admin forward rules to target local/private addresses when explicitly enabled.
|
||||
|
||||
**Architecture:** Keep the existing remote-address safety validator as the default path for non-admin rule changes, but gate its use behind a single backend config lookup in forward create/update handlers. Surface the toggle through the existing `vite_config` settings page and prove behavior with backend contract tests first.
|
||||
|
||||
**Tech Stack:** Go `net/http` + GORM backend, React + TypeScript frontend settings page, Go contract tests.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Backend Contract Coverage
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/tests/contract/forward_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing tests**
|
||||
|
||||
Add contract tests that prove the desired behavior:
|
||||
|
||||
```go
|
||||
t.Run("local remote address is rejected when toggle is off", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "deny-local-remote",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, _ := json.Marshal(createPayload)
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
|
||||
var out response.R
|
||||
_ = json.NewDecoder(createRes.Body).Decode(&out)
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected local remote address to be rejected when toggle is off")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("local remote address is allowed when toggle is on", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, "allow_local_remote_addr", "1", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("enable allow_local_remote_addr: %v", err)
|
||||
}
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "allow-local-remote",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, _ := json.Marshal(createPayload)
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
})
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run tests to verify they fail**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: FAIL because backend still rejects local/private addresses unconditionally.
|
||||
|
||||
- [ ] **Step 3: Commit**
|
||||
|
||||
Do not commit yet; combine with Task 2 after implementation passes.
|
||||
|
||||
### Task 2: Backend Toggle Implementation
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/http/handler/mutations.go`
|
||||
|
||||
- [ ] **Step 1: Add a tiny config helper**
|
||||
|
||||
Add a helper near other handler helpers:
|
||||
|
||||
```go
|
||||
func (h *Handler) allowLocalRemoteAddr() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return false
|
||||
}
|
||||
cfg, err := h.repo.GetConfigByName("allow_local_remote_addr")
|
||||
if err != nil || cfg == nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(cfg.Value) == "1"
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Gate create/update validation behind the helper**
|
||||
|
||||
Replace the unconditional checks with:
|
||||
|
||||
```go
|
||||
if !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Run contract tests to verify they pass**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 4: Run full backend tests**
|
||||
|
||||
Run: `go test ./...`
|
||||
Expected: PASS
|
||||
|
||||
### Task 3: Settings Page Toggle
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
- [ ] **Step 1: Add the config item to the settings schema**
|
||||
|
||||
Add a switch-style item for `allow_local_remote_addr` with warning copy about reduced safety.
|
||||
|
||||
- [ ] **Step 2: Ensure the key is included in config loading/saving paths**
|
||||
|
||||
Add `allow_local_remote_addr` anywhere the page enumerates config keys or groups persisted config values.
|
||||
|
||||
- [ ] **Step 3: Run frontend build**
|
||||
|
||||
Run: `pnpm run build`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 4: Run frontend lint**
|
||||
|
||||
Run: `pnpm run lint`
|
||||
Expected: 0 errors; existing warnings may remain.
|
||||
|
||||
### Task 4: Final Verification
|
||||
|
||||
**Files:**
|
||||
- Verify only
|
||||
|
||||
- [ ] **Step 1: Re-run backend contracts for the toggle**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 2: Re-run full backend tests**
|
||||
|
||||
Run: `go test ./...`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 3: Re-run frontend build/lint**
|
||||
|
||||
Run: `pnpm run build && pnpm run lint`
|
||||
Expected: Build passes, lint has no errors.
|
||||
|
||||
- [ ] **Step 4: Commit**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/mutations.go go-backend/tests/contract/forward_contract_test.go vite-frontend/src/pages/config.tsx docs/superpowers/specs/2026-04-26-allow-local-remote-addr-design.md docs/superpowers/plans/2026-04-26-allow-local-remote-addr.md
|
||||
git commit -m "feat: add allow-local-remote-address toggle"
|
||||
```
|
||||
@@ -0,0 +1,849 @@
|
||||
# flow/upload Batch Optimization Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Reduce `POST /flow/upload` database pressure by converting the hot path from per-item queries and per-item transactions to per-request aggregation, batched metadata reads, and batched writes, while preserving immediate quota disable / forward pause behavior inside the same upload.
|
||||
|
||||
**Architecture:** Parse one upload into a batch object in the handler layer, fetch one shared `forward+tunnel` metadata map, then reuse that map for flow accounting and tunnel metric aggregation. Replace `AddFlow` and `AddUserQuotaUsage` per-item transactions with one batched flow transaction and one batched quota transaction; run policy enforcement, orphan cleanup, and peer-share flow handling once per affected target instead of once per item.
|
||||
|
||||
**Tech Stack:** Go, net/http, GORM, SQLite/PostgreSQL, existing backend contract tests.
|
||||
|
||||
---
|
||||
|
||||
## File Map
|
||||
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch.go`
|
||||
Responsibility: request-scoped parsing, aggregation, and application of one `/flow/upload` batch.
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch_test.go`
|
||||
Responsibility: unit coverage for batch aggregation semantics.
|
||||
- Create: `go-backend/internal/store/repo/repository_flow_batch_test.go`
|
||||
Responsibility: unit coverage for batched flow and quota persistence.
|
||||
- Create: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
Responsibility: contract coverage that repeated items still accumulate correctly and still disable quota immediately.
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
Responsibility: switch `/flow/upload` entrypoint to the new batch pipeline.
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_metrics_ingestion.go`
|
||||
Responsibility: accept pre-aggregated forward deltas plus shared forward metadata instead of reparsing the raw items.
|
||||
- Modify: `go-backend/internal/store/repo/repository.go`
|
||||
Responsibility: add batched flow persistence primitives near the existing flow update code.
|
||||
- Modify: `go-backend/internal/store/repo/repository_flow.go`
|
||||
Responsibility: add shared flow-upload metadata query helpers.
|
||||
- Modify: `go-backend/internal/store/repo/repository_user_quota.go`
|
||||
Responsibility: add batched quota usage persistence that still returns normalized quota views for immediate enforcement.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Add Failing Tests For Batched flow/upload Semantics
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch_test.go`
|
||||
- Create: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing handler unit test**
|
||||
|
||||
Create `go-backend/internal/http/handler/flow_upload_batch_test.go` with a unit test that locks in the new aggregation contract.
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets(t *testing.T) {
|
||||
h := &Handler{}
|
||||
metas := map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {
|
||||
ForwardID: 20,
|
||||
TunnelID: 1,
|
||||
TrafficRatio: 2,
|
||||
TunnelFlow: 3,
|
||||
},
|
||||
}
|
||||
|
||||
batch := h.buildFlowUploadBatch([]flowItem{
|
||||
{N: "20_2_10", U: 70, D: 50},
|
||||
{N: "20_2_10_tcp", U: 40, D: 30},
|
||||
{N: "99_2_10", U: 12, D: 8},
|
||||
{N: "fed_svc_17", U: 9, D: 1},
|
||||
}, metas)
|
||||
|
||||
if len(batch.flowDeltas) != 1 {
|
||||
t.Fatalf("expected 1 flow delta, got %d", len(batch.flowDeltas))
|
||||
}
|
||||
delta := batch.flowDeltas[0]
|
||||
if delta.ForwardID != 20 || delta.UserID != 2 || delta.UserTunnelID != 10 {
|
||||
t.Fatalf("unexpected flow delta identity: %#v", delta)
|
||||
}
|
||||
if delta.InFlow != 480 || delta.OutFlow != 660 {
|
||||
t.Fatalf("expected scaled flow in=480 out=660, got in=%d out=%d", delta.InFlow, delta.OutFlow)
|
||||
}
|
||||
if batch.quotaUsage[2] != 1140 {
|
||||
t.Fatalf("expected quota usage 1140, got %d", batch.quotaUsage[2])
|
||||
}
|
||||
if len(batch.policyTargets) != 1 {
|
||||
t.Fatalf("expected 1 policy target, got %d", len(batch.policyTargets))
|
||||
}
|
||||
if batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("unexpected policy target: %#v", batch.policyTargets[0])
|
||||
}
|
||||
traffic := batch.forwardTraffic[20]
|
||||
if traffic.bytesIn != 80 || traffic.bytesOut != 110 {
|
||||
t.Fatalf("expected raw traffic in=80 out=110, got in=%d out=%d", traffic.bytesIn, traffic.bytesOut)
|
||||
}
|
||||
if _, ok := batch.orphanServices["99_2_10"]; !ok {
|
||||
t.Fatalf("expected orphan service cleanup target for 99_2_10")
|
||||
}
|
||||
if item, ok := batch.peerShareForwardItems["20_2_10"]; !ok || item.U != 110 || item.D != 80 {
|
||||
t.Fatalf("expected merged peer-share forward item, got %#v ok=%v", item, ok)
|
||||
}
|
||||
if item, ok := batch.peerShareRuntimeItems[17]; !ok || item.U != 9 || item.D != 1 {
|
||||
t.Fatalf("expected merged peer-share runtime item, got %#v ok=%v", item, ok)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the handler unit test to verify RED**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets -v
|
||||
```
|
||||
|
||||
Expected: FAIL because `FlowUploadForwardMeta`, `buildFlowUploadBatch`, and the new batch fields do not exist yet.
|
||||
|
||||
- [ ] **Step 3: Write the contract test that guards current behavior**
|
||||
|
||||
Create `go-backend/tests/contract/flow_upload_batch_contract_test.go` so the optimization cannot weaken same-request quota enforcement.
|
||||
|
||||
```go
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
const bytesPerGB = int64(1024 * 1024 * 1024)
|
||||
|
||||
node := &model.Node{Name: "node-1", Secret: "node-secret", ServerIP: "127.0.0.1", Port: "10000-10010", TCPListenAddr: "[::]", UDPListenAddr: "[::]", CreatedTime: nowMs, Status: 1}
|
||||
if err := repo.DB().Create(node).Error; err != nil {
|
||||
t.Fatalf("seed node: %v", err)
|
||||
}
|
||||
if err := repo.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, 'flow_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
tunnel := &model.Tunnel{Name: "tunnel-1", TrafficRatio: 1.0, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
|
||||
if err := repo.DB().Create(tunnel).Error; err != nil {
|
||||
t.Fatalf("seed tunnel: %v", err)
|
||||
}
|
||||
if err := repo.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, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnel.ID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
forward := &model.Forward{ID: 20, UserID: 2, UserName: "flow_user", Name: "forward-20", TunnelID: tunnel.ID, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
|
||||
if err := repo.DB().Create(forward).Error; err != nil {
|
||||
t.Fatalf("seed forward: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 1, 0, ?, ?, ?, ?, 0, 0, '', ?, ?)`, bytesPerGB-100, bytesPerGB-100, dayKey, monthKey, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
body, err := json.Marshal([]map[string]interface{}{
|
||||
{"n": "20_2_10", "u": 70, "d": 50},
|
||||
{"n": "20_2_10_tcp", "u": 40, "d": 30},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal body: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d", res.Code)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = 20`); got != 0 {
|
||||
t.Fatalf("expected forward paused immediately, got status=%d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`); got != 1 {
|
||||
t.Fatalf("expected quota disabled flag=1, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 80 {
|
||||
t.Fatalf("expected forward in_flow=80, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 110 {
|
||||
t.Fatalf("expected forward out_flow=110, got %d", got)
|
||||
}
|
||||
metrics, err := repo.GetTunnelMetrics(tunnel.ID, 0, nowMs+60_000)
|
||||
if err != nil {
|
||||
t.Fatalf("get tunnel metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 || metrics[0].BytesIn != 80 || metrics[0].BytesOut != 110 {
|
||||
t.Fatalf("expected one aggregated metric row, got %#v", metrics)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Run the contract test to verify the same-request guard stays green or reveals an existing regression**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
|
||||
```
|
||||
|
||||
Expected: this test may already PASS before the refactor because it locks in existing external behavior. Keep it either way; it is the guardrail for the optimization.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/flow_upload_batch_test.go go-backend/tests/contract/flow_upload_batch_contract_test.go
|
||||
git commit -m "test: cover flow upload batch semantics"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Add Batched Repository Primitives
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/store/repo/repository_flow.go`
|
||||
- Modify: `go-backend/internal/store/repo/repository.go`
|
||||
- Modify: `go-backend/internal/store/repo/repository_user_quota.go`
|
||||
- Create: `go-backend/internal/store/repo/repository_flow_batch_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing repository tests**
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_flow_batch_test.go` with coverage for both the shared metadata query and the batched counter/quota writes.
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "flow-batch.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
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)
|
||||
}
|
||||
|
||||
metas, err := r.GetFlowUploadForwardMetas([]int64{20, 99})
|
||||
if err != nil {
|
||||
t.Fatalf("get metas: %v", err)
|
||||
}
|
||||
if 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 {
|
||||
t.Fatalf("did not expect meta for missing forward 99")
|
||||
}
|
||||
|
||||
err = r.ApplyFlowUploadDeltasBatch([]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}})
|
||||
if err != nil {
|
||||
t.Fatalf("apply flow batch: %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)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddUserQuotaUsageBatchReturnsNormalizedViews(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "quota-batch.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
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)`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
views, err := r.AddUserQuotaUsageBatch(map[int64]int64{2: 1140}, now)
|
||||
if err != nil {
|
||||
t.Fatalf("batch quota update: %v", err)
|
||||
}
|
||||
if views[2] == nil || views[2].DailyUsedBytes != 1140 || views[2].MonthlyUsedBytes != 1140 {
|
||||
t.Fatalf("unexpected quota view: %#v", views[2])
|
||||
}
|
||||
}
|
||||
|
||||
func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...interface{}) int64 {
|
||||
t.Helper()
|
||||
var value int64
|
||||
if err := r.DB().Raw(query, args...).Row().Scan(&value); err != nil {
|
||||
t.Fatalf("query %q failed: %v", query, err)
|
||||
}
|
||||
return value
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the repository tests to verify RED**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/store/repo -run 'TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch|TestAddUserQuotaUsageBatchReturnsNormalizedViews' -v
|
||||
```
|
||||
|
||||
Expected: FAIL because `GetFlowUploadForwardMetas`, `ApplyFlowUploadDeltasBatch`, `FlowUploadCounterDelta`, and `AddUserQuotaUsageBatch` do not exist yet.
|
||||
|
||||
- [ ] **Step 3: Implement shared flow-upload metadata and batched persistence**
|
||||
|
||||
Update `go-backend/internal/store/repo/repository_flow.go`, `repository.go`, and `repository_user_quota.go` with the following concrete APIs. Add `sort` to the `repository_user_quota.go` import list.
|
||||
|
||||
```go
|
||||
// repository_flow.go
|
||||
type FlowUploadForwardMeta struct {
|
||||
ForwardID int64
|
||||
TunnelID int64
|
||||
TrafficRatio float64
|
||||
TunnelFlow int64
|
||||
}
|
||||
|
||||
func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]FlowUploadForwardMeta, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if len(forwardIDs) == 0 {
|
||||
return map[int64]FlowUploadForwardMeta{}, nil
|
||||
}
|
||||
ids := make([]int64, 0, len(forwardIDs))
|
||||
seen := make(map[int64]struct{}, len(forwardIDs))
|
||||
for _, id := range forwardIDs {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
type row struct {
|
||||
ForwardID int64 `gorm:"column:forward_id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id"`
|
||||
TrafficRatio float64 `gorm:"column:traffic_ratio"`
|
||||
TunnelFlow int64 `gorm:"column:tunnel_flow"`
|
||||
}
|
||||
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").
|
||||
Joins("JOIN tunnel t ON t.id = f.tunnel_id").
|
||||
Where("f.id IN ?", ids).
|
||||
Scan(&rows).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make(map[int64]FlowUploadForwardMeta, len(rows))
|
||||
for _, row := range rows {
|
||||
if row.TunnelFlow <= 0 {
|
||||
row.TunnelFlow = 1
|
||||
}
|
||||
if row.TrafficRatio <= 0 {
|
||||
row.TrafficRatio = 1
|
||||
}
|
||||
out[row.ForwardID] = FlowUploadForwardMeta{ForwardID: row.ForwardID, TunnelID: row.TunnelID, TrafficRatio: row.TrafficRatio, TunnelFlow: row.TunnelFlow}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
```
|
||||
|
||||
```go
|
||||
// repository.go
|
||||
type FlowUploadCounterDelta struct {
|
||||
ForwardID int64
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
}
|
||||
|
||||
func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
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))
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID > 0 {
|
||||
current := forwardTotals[delta.ForwardID]
|
||||
current[0] += delta.InFlow
|
||||
current[1] += delta.OutFlow
|
||||
forwardTotals[delta.ForwardID] = current
|
||||
}
|
||||
if delta.UserID > 0 {
|
||||
current := userTotals[delta.UserID]
|
||||
current[0] += delta.InFlow
|
||||
current[1] += delta.OutFlow
|
||||
userTotals[delta.UserID] = current
|
||||
}
|
||||
if delta.UserTunnelID > 0 {
|
||||
current := userTunnelTotals[delta.UserTunnelID]
|
||||
current[0] += delta.InFlow
|
||||
current[1] += delta.OutFlow
|
||||
userTunnelTotals[delta.UserTunnelID] = current
|
||||
}
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
for forwardID, total := range forwardTotals {
|
||||
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, total := range userTotals {
|
||||
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, total := range userTunnelTotals {
|
||||
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
|
||||
})
|
||||
}
|
||||
```
|
||||
|
||||
```go
|
||||
// repository_user_quota.go
|
||||
func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if len(usages) == 0 {
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
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
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Run the repository tests to verify GREEN**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/store/repo -run 'TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch|TestAddUserQuotaUsageBatchReturnsNormalizedViews' -v
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/store/repo/repository.go go-backend/internal/store/repo/repository_flow.go go-backend/internal/store/repo/repository_user_quota.go go-backend/internal/store/repo/repository_flow_batch_test.go
|
||||
git commit -m "refactor: batch flow upload persistence"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Refactor flow/upload To Use One Parsed Batch
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_metrics_ingestion.go`
|
||||
- Modify: `go-backend/internal/http/handler/flow_upload_batch_test.go`
|
||||
- Modify: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the new handler batch implementation**
|
||||
|
||||
Create `go-backend/internal/http/handler/flow_upload_batch.go` and move the request-scoped aggregation there.
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type flowPolicyTarget struct {
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
}
|
||||
|
||||
type flowUploadBatch struct {
|
||||
flowDeltas []repo.FlowUploadCounterDelta
|
||||
quotaUsage map[int64]int64
|
||||
policyTargets []flowPolicyTarget
|
||||
forwardTraffic map[int64]tunnelTrafficDelta
|
||||
orphanServices map[string]struct{}
|
||||
peerShareForwardItems map[string]flowItem
|
||||
peerShareRuntimeItems map[int64]flowItem
|
||||
}
|
||||
|
||||
func (h *Handler) buildFlowUploadBatch(items []flowItem, 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 _, item := range items {
|
||||
serviceName := strings.TrimSpace(item.N)
|
||||
if serviceName == "" || serviceName == "web_api" {
|
||||
continue
|
||||
}
|
||||
if runtimeID, ok := parsePeerShareRuntimeServiceID(serviceName); ok {
|
||||
merged := batch.peerShareRuntimeItems[runtimeID]
|
||||
merged.N = serviceName
|
||||
merged.U += item.U
|
||||
merged.D += item.D
|
||||
batch.peerShareRuntimeItems[runtimeID] = merged
|
||||
continue
|
||||
}
|
||||
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
meta, exists := metas[forwardID]
|
||||
if !exists {
|
||||
batch.orphanServices[serviceName] = struct{}{}
|
||||
continue
|
||||
}
|
||||
raw := batch.forwardTraffic[forwardID]
|
||||
raw.bytesIn += item.D
|
||||
raw.bytesOut += item.U
|
||||
batch.forwardTraffic[forwardID] = raw
|
||||
|
||||
scaledIn := int64(float64(item.D)*meta.TrafficRatio) * meta.TunnelFlow
|
||||
scaledOut := int64(float64(item.U)*meta.TrafficRatio) * meta.TunnelFlow
|
||||
if idx, ok := flowSeen[forwardID]; ok {
|
||||
batch.flowDeltas[idx].InFlow += scaledIn
|
||||
batch.flowDeltas[idx].OutFlow += scaledOut
|
||||
} else {
|
||||
flowSeen[forwardID] = len(batch.flowDeltas)
|
||||
batch.flowDeltas = append(batch.flowDeltas, repo.FlowUploadCounterDelta{ForwardID: forwardID, UserID: userID, UserTunnelID: userTunnelID, InFlow: scaledIn, OutFlow: scaledOut})
|
||||
}
|
||||
batch.quotaUsage[userID] += scaledIn + scaledOut
|
||||
target := flowPolicyTarget{UserID: userID, UserTunnelID: userTunnelID}
|
||||
if _, seen := policySeen[target]; !seen {
|
||||
policySeen[target] = struct{}{}
|
||||
batch.policyTargets = append(batch.policyTargets, target)
|
||||
}
|
||||
merged := batch.peerShareForwardItems[normalizeForwardRuntimeServiceName(serviceName)]
|
||||
merged.N = normalizeForwardRuntimeServiceName(serviceName)
|
||||
merged.U += item.U
|
||||
merged.D += item.D
|
||||
batch.peerShareForwardItems[normalizeForwardRuntimeServiceName(serviceName)] = merged
|
||||
}
|
||||
|
||||
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 (h *Handler) applyFlowUploadBatch(nodeID int64, batch flowUploadBatch, now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
if err := h.repo.ApplyFlowUploadDeltasBatch(batch.flowDeltas); err != nil {
|
||||
log.Printf("flow upload write failed op=flow.batch_apply node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
quotaViews, err := h.repo.AddUserQuotaUsageBatch(batch.quotaUsage, now)
|
||||
if err != nil {
|
||||
log.Printf("flow upload write failed op=quota.batch_apply node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
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)
|
||||
}
|
||||
for serviceName := range batch.orphanServices {
|
||||
h.sendDeleteOrphanedForwardService(nodeID, serviceName)
|
||||
}
|
||||
for serviceName, item := range batch.peerShareForwardItems {
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(serviceName)
|
||||
if ok {
|
||||
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
|
||||
}
|
||||
}
|
||||
for runtimeID, item := range batch.peerShareRuntimeItems {
|
||||
h.processPeerShareFlow(runtimeID, item)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Switch the `/flow/upload` entrypoint and tunnel metric ingestion to the shared batch**
|
||||
|
||||
Modify `handler.go` and `tunnel_metrics_ingestion.go` so the raw JSON is parsed once and the same forward metadata powers both flow counters and tunnel metrics.
|
||||
|
||||
```go
|
||||
// handler.go
|
||||
func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
secret := r.URL.Query().Get("secret")
|
||||
node, _ := h.repo.GetNodeBySecret(secret)
|
||||
if node == nil {
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
return
|
||||
}
|
||||
|
||||
raw, err := readAndDecryptFlowBody(r.Body, secret)
|
||||
if err == nil && strings.TrimSpace(raw) != "" {
|
||||
var items []flowItem
|
||||
if json.Unmarshal([]byte(raw), &items) == nil {
|
||||
now := time.Now()
|
||||
forwardIDs := collectFlowUploadForwardIDs(items)
|
||||
metas, metaErr := h.repo.GetFlowUploadForwardMetas(forwardIDs)
|
||||
if metaErr != nil {
|
||||
log.Printf("flow upload metadata lookup failed node_id=%d err=%v", node.ID, metaErr)
|
||||
metas = map[int64]repo.FlowUploadForwardMeta{}
|
||||
}
|
||||
batch := h.buildFlowUploadBatch(items, metas)
|
||||
h.recordTunnelMetricsFromForwardBatch(node.ID, batch.forwardTraffic, metas, now.UnixMilli())
|
||||
h.applyFlowUploadBatch(node.ID, batch, now)
|
||||
}
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}
|
||||
```
|
||||
|
||||
```go
|
||||
// tunnel_metrics_ingestion.go
|
||||
func collectFlowUploadForwardIDs(items []flowItem) []int64 {
|
||||
ids := make([]int64, 0, len(items))
|
||||
seen := make(map[int64]struct{}, len(items))
|
||||
for _, item := range items {
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(strings.TrimSpace(item.N))
|
||||
if !ok || forwardID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[forwardID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[forwardID] = struct{}{}
|
||||
ids = append(ids, forwardID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func (h *Handler) recordTunnelMetricsFromForwardBatch(nodeID int64, forwardDeltas map[int64]tunnelTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta, nowMs int64) {
|
||||
if h == nil || h.repo == nil || nodeID <= 0 || len(forwardDeltas) == 0 {
|
||||
return
|
||||
}
|
||||
bucketTs := unixMilliBucketMinute(nowMs)
|
||||
if bucketTs <= 0 {
|
||||
return
|
||||
}
|
||||
tunnelAgg := make(map[int64]tunnelTrafficDelta)
|
||||
for forwardID, delta := range forwardDeltas {
|
||||
meta, ok := metas[forwardID]
|
||||
if !ok || meta.TunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
current := tunnelAgg[meta.TunnelID]
|
||||
current.bytesIn += delta.bytesIn
|
||||
current.bytesOut += delta.bytesOut
|
||||
tunnelAgg[meta.TunnelID] = current
|
||||
}
|
||||
metrics := make([]*model.TunnelMetric, 0, len(tunnelAgg))
|
||||
for tunnelID, delta := range tunnelAgg {
|
||||
if delta.bytesIn == 0 && delta.bytesOut == 0 {
|
||||
continue
|
||||
}
|
||||
metrics = append(metrics, &model.TunnelMetric{TunnelID: tunnelID, NodeID: nodeID, Timestamp: bucketTs, BytesIn: delta.bytesIn, BytesOut: delta.bytesOut})
|
||||
}
|
||||
if len(metrics) == 0 {
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpsertTunnelMetricBuckets(metrics); err != nil {
|
||||
log.Printf("monitoring write failed op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d err=%v", nodeID, bucketTs, len(metrics), err)
|
||||
return
|
||||
}
|
||||
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Run focused handler and contract tests to verify GREEN**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets -v
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 4: Run the full backend suite**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./...
|
||||
```
|
||||
|
||||
Expected: PASS across the backend module.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/tunnel_metrics_ingestion.go go-backend/internal/http/handler/flow_upload_batch.go go-backend/internal/http/handler/flow_upload_batch_test.go go-backend/tests/contract/flow_upload_batch_contract_test.go
|
||||
git commit -m "refactor: batch flow upload processing"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Final Verification And Performance Sanity Check
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Add a same-batch duplicate-item stress assertion**
|
||||
|
||||
Extend the contract test with a second request that repeats the same service name multiple times and assert the counters advance by exactly the summed amount.
|
||||
|
||||
```go
|
||||
body, err = json.Marshal([]map[string]interface{}{
|
||||
{"n": "20_2_10", "u": 10, "d": 20},
|
||||
{"n": "20_2_10", "u": 10, "d": 20},
|
||||
{"n": "20_2_10_tcp", "u": 10, "d": 20},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal body: %v", err)
|
||||
}
|
||||
req = httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
|
||||
res = httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 140 {
|
||||
t.Fatalf("expected forward in_flow=140 after second request, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 140 {
|
||||
t.Fatalf("expected forward out_flow=140 after second request, got %d", got)
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the targeted contract test again**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Re-run the full backend suite before claiming completion**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./...
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 4: Optional local profiling sanity check**
|
||||
|
||||
Run a short local comparison before and after the change with the same repeated flow payload.
|
||||
|
||||
```bash
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -count=10
|
||||
```
|
||||
|
||||
Expected: the test remains stable across repeated runs and does not introduce flakiness.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/tests/contract/flow_upload_batch_contract_test.go
|
||||
git commit -m "test: harden flow upload batch regression coverage"
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,598 @@
|
||||
# Monitoring Retention And Storage Display Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Add configurable monitoring data retention and show database storage usage on the config page.
|
||||
|
||||
**Architecture:** Store retention in `vite_config` as `monitor_retention_days`, parse it through a focused monitoring helper, and reuse it from existing cleanup loops. Add a repository storage-summary helper, expose it via an admin-only API, and render it in the existing React config page.
|
||||
|
||||
**Tech Stack:** Go `net/http`, GORM, SQLite/PostgreSQL, Vite/React/TypeScript, existing shadcn bridge components.
|
||||
|
||||
---
|
||||
|
||||
## File Structure
|
||||
|
||||
- Create: `go-backend/internal/monitoring/retention.go` for retention constants, parsing, and validation.
|
||||
- Test: `go-backend/internal/monitoring/retention_test.go`.
|
||||
- Modify: `go-backend/internal/metrics/ingestion.go` and `go-backend/internal/metrics/ingestion_test.go` for config-driven cleanup.
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_quality_prober.go` so `tunnel_quality` uses the same retention and still prunes when probing is disabled.
|
||||
- Create: `go-backend/internal/store/repo/repository_storage.go` and `go-backend/internal/store/repo/repository_storage_test.go` for database size summaries.
|
||||
- Modify: `go-backend/internal/store/repo/repository.go` to keep the SQLite DB path on `Repository`.
|
||||
- Create: `go-backend/internal/http/handler/storage.go` for the storage endpoint.
|
||||
- Modify: `go-backend/internal/http/handler/handler.go` to register `/api/v1/system/storage` and validate `monitor_retention_days`.
|
||||
- Modify: `go-backend/internal/http/middleware/auth.go` so `/api/v1/system/*` is admin-only.
|
||||
- Create: `go-backend/tests/contract/storage_contract_test.go` for endpoint auth/shape coverage.
|
||||
- Modify: `vite-frontend/src/api/types.ts`, `vite-frontend/src/api/index.ts`, and `vite-frontend/src/pages/config.tsx` for UI display.
|
||||
|
||||
Implementation should not create git commits unless the user explicitly requests them.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Add Retention Config Helper
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/monitoring/retention.go`
|
||||
- Create: `go-backend/internal/monitoring/retention_test.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing tests**
|
||||
|
||||
Create `go-backend/internal/monitoring/retention_test.go`:
|
||||
|
||||
```go
|
||||
package monitoring
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestMonitoringRetentionDaysFromConfigMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg map[string]string
|
||||
want int
|
||||
}{
|
||||
{"missing uses default", nil, 7},
|
||||
{"valid custom", map[string]string{ConfigMonitorRetentionDays: "3"}, 3},
|
||||
{"trimmed custom", map[string]string{ConfigMonitorRetentionDays: " 30 "}, 30},
|
||||
{"invalid uses default", map[string]string{ConfigMonitorRetentionDays: "abc"}, 7},
|
||||
{"too small uses default", map[string]string{ConfigMonitorRetentionDays: "0"}, 7},
|
||||
{"too large uses default", map[string]string{ConfigMonitorRetentionDays: "3651"}, 7},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := MonitoringRetentionDaysFromConfigMap(tc.cfg); got != tc.want {
|
||||
t.Fatalf("expected %d, got %d", tc.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeMonitoringRetentionDays(t *testing.T) {
|
||||
for _, value := range []string{"1", "7", "3650", " 30 "} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err != nil || got == "" {
|
||||
t.Fatalf("expected %q valid, got value=%q err=%v", value, got, err)
|
||||
}
|
||||
}
|
||||
for _, value := range []string{"", "0", "-1", "3651", "abc", "1.5"} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err == nil {
|
||||
t.Fatalf("expected %q invalid, got value=%q", value, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run tests to verify failure**
|
||||
|
||||
Run: `go test ./internal/monitoring -run 'TestMonitoringRetentionDaysFromConfigMap|TestNormalizeMonitoringRetentionDays' -count=1`
|
||||
|
||||
Expected: FAIL with undefined `ConfigMonitorRetentionDays`, `MonitoringRetentionDaysFromConfigMap`, and `NormalizeMonitoringRetentionDays`.
|
||||
|
||||
- [ ] **Step 3: Implement helper**
|
||||
|
||||
Create `go-backend/internal/monitoring/retention.go`:
|
||||
|
||||
```go
|
||||
package monitoring
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
ConfigMonitorRetentionDays = "monitor_retention_days"
|
||||
DefaultMonitorRetentionDays = 7
|
||||
MinMonitorRetentionDays = 1
|
||||
MaxMonitorRetentionDays = 3650
|
||||
)
|
||||
|
||||
func MonitoringRetentionDaysFromConfigMap(cfg map[string]string) int {
|
||||
if cfg == nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
days, err := parseMonitoringRetentionDays(cfg[ConfigMonitorRetentionDays])
|
||||
if err != nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
return days
|
||||
}
|
||||
|
||||
func NormalizeMonitoringRetentionDays(value string) (string, error) {
|
||||
days, err := parseMonitoringRetentionDays(value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strconv.Itoa(days), nil
|
||||
}
|
||||
|
||||
func parseMonitoringRetentionDays(value string) (int, error) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return 0, fmt.Errorf("监控数据保留天数不能为空")
|
||||
}
|
||||
days, err := strconv.Atoi(trimmed)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须是整数")
|
||||
}
|
||||
if days < MinMonitorRetentionDays || days > MaxMonitorRetentionDays {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须在 %d 到 %d 之间", MinMonitorRetentionDays, MaxMonitorRetentionDays)
|
||||
}
|
||||
return days, nil
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Validate config updates**
|
||||
|
||||
In `go-backend/internal/http/handler/handler.go`, add this case to `normalizeAndValidateConfigValue`:
|
||||
|
||||
```go
|
||||
case monitoring.ConfigMonitorRetentionDays:
|
||||
return monitoring.NormalizeMonitoringRetentionDays(value)
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Run tests**
|
||||
|
||||
Run: `go test ./internal/monitoring ./internal/http/handler -run 'TestMonitoringRetention|TestNormalize|Test' -count=1`
|
||||
|
||||
Expected: PASS or only unrelated pre-existing failures, which must be investigated before continuing.
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Use Retention Config In Cleanup
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/metrics/ingestion.go`
|
||||
- Modify: `go-backend/internal/metrics/ingestion_test.go`
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_quality_prober.go`
|
||||
|
||||
- [ ] **Step 1: Write failing cleanup test**
|
||||
|
||||
Append to `go-backend/internal/metrics/ingestion_test.go`, adding `go-backend/internal/store/model` to imports:
|
||||
|
||||
```go
|
||||
func TestPruneMetricsUsesConfiguredRetentionDays(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.UpsertConfig("monitor_retention_days", "2", now); err != nil {
|
||||
t.Fatalf("upsert retention config: %v", err)
|
||||
}
|
||||
|
||||
oldMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(3*24*time.Hour/time.Millisecond), CPUUsage: 10}
|
||||
newMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(1*24*time.Hour/time.Millisecond), CPUUsage: 20}
|
||||
if err := r.InsertNodeMetric(oldMetric); err != nil {
|
||||
t.Fatalf("insert old metric: %v", err)
|
||||
}
|
||||
if err := r.InsertNodeMetric(newMetric); err != nil {
|
||||
t.Fatalf("insert new metric: %v", err)
|
||||
}
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.pruneMetricsAt(time.UnixMilli(now))
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, now-int64(4*24*time.Hour/time.Millisecond), now+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get node metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 || metrics[0].CPUUsage != 20 {
|
||||
t.Fatalf("expected only newer metric to remain, got %#v", metrics)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run test to verify failure**
|
||||
|
||||
Run: `go test ./internal/metrics -run TestPruneMetricsUsesConfiguredRetentionDays -count=1`
|
||||
|
||||
Expected: FAIL with undefined `pruneMetricsAt`.
|
||||
|
||||
- [ ] **Step 3: Implement config-driven prune**
|
||||
|
||||
In `go-backend/internal/metrics/ingestion.go`, import `go-backend/internal/monitoring` and replace `pruneMetrics` with:
|
||||
|
||||
```go
|
||||
func (s *IngestionService) retentionDaysFromConfig() int {
|
||||
if s == nil || s.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := s.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetrics() {
|
||||
s.pruneMetricsAt(time.Now())
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetricsAt(now time.Time) {
|
||||
cutoff := now.Add(-time.Duration(s.retentionDaysFromConfig()) * 24 * time.Hour).UnixMilli()
|
||||
if s.repo == nil {
|
||||
return
|
||||
}
|
||||
if err := s.repo.PruneNodeMetrics(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=node_metric cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
if err := s.repo.PruneTunnelMetrics(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=tunnel_metric cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
if err := s.repo.PruneServiceMonitorResults(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=service_monitor_result cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Remove the unused `retentionDays` field from `IngestionService` and remove `svc.retentionDays = 1` from existing tests.
|
||||
|
||||
- [ ] **Step 4: Update tunnel quality pruning**
|
||||
|
||||
In `go-backend/internal/http/handler/tunnel_quality_prober.go`, import `go-backend/internal/monitoring`, remove `tunnelQualityRetention`, remove the `if !p.isEnabled() { return }` guard from `maybePrune`, and calculate cutoff with:
|
||||
|
||||
```go
|
||||
func (p *tunnelQualityProber) retentionDays() int {
|
||||
if p == nil || p.handler == nil || p.handler.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := p.handler.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
```
|
||||
|
||||
Then use:
|
||||
|
||||
```go
|
||||
cutoff := now - int64(time.Duration(p.retentionDays())*24*time.Hour/time.Millisecond)
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Run cleanup tests**
|
||||
|
||||
Run: `go test ./internal/metrics ./internal/http/handler -run 'TestPruneMetrics|TestPruneMetricsUsesConfiguredRetentionDays|TunnelQuality' -count=1`
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Add Storage Summary Backend API
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/store/repo/repository.go`
|
||||
- Create: `go-backend/internal/store/repo/repository_storage.go`
|
||||
- Create: `go-backend/internal/store/repo/repository_storage_test.go`
|
||||
- Create: `go-backend/internal/http/handler/storage.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/middleware/auth.go`
|
||||
- Create: `go-backend/tests/contract/storage_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write failing repository tests**
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_storage_test.go`:
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestDatabaseStorageSummarySQLiteIncludesSize(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "storage.db")
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
if err := r.InsertNodeMetric(&model.NodeMetric{NodeID: 1, Timestamp: 123, CPUUsage: 1}); err != nil {
|
||||
t.Fatalf("insert metric: %v", err)
|
||||
}
|
||||
summary, err := r.DatabaseStorageSummary()
|
||||
if err != nil {
|
||||
t.Fatalf("storage summary: %v", err)
|
||||
}
|
||||
if summary.DBType != "sqlite" || summary.DatabaseSizeBytes <= 0 || summary.DatabaseSizeText == "" {
|
||||
t.Fatalf("unexpected summary: %#v", summary)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatDatabaseSize(t *testing.T) {
|
||||
for _, tc := range []struct{ bytes int64; want string }{{0, "0 B"}, {512, "512 B"}, {1024, "1.0 KB"}, {1024 * 1024, "1.0 MB"}} {
|
||||
if got := formatDatabaseSize(tc.bytes); got != tc.want {
|
||||
t.Fatalf("formatDatabaseSize(%d)=%q want %q", tc.bytes, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run test to verify failure**
|
||||
|
||||
Run: `go test ./internal/store/repo -run 'TestDatabaseStorageSummarySQLiteIncludesSize|TestFormatDatabaseSize' -count=1`
|
||||
|
||||
Expected: FAIL with undefined `DatabaseStorageSummary` and `formatDatabaseSize`.
|
||||
|
||||
- [ ] **Step 3: Implement repository helper**
|
||||
|
||||
Modify `Repository` in `repository.go`:
|
||||
|
||||
```go
|
||||
type Repository struct {
|
||||
db *gorm.DB
|
||||
dbPath string
|
||||
}
|
||||
```
|
||||
|
||||
Return `&Repository{db: db, dbPath: path}` from `Open` and `&Repository{db: db}` from `OpenPostgres`.
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_storage.go`:
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
type DatabaseStorageSummary struct {
|
||||
DBType string `json:"dbType"`
|
||||
DatabaseSizeBytes int64 `json:"databaseSizeBytes"`
|
||||
DatabaseSizeText string `json:"databaseSizeText"`
|
||||
}
|
||||
|
||||
func (r *Repository) DatabaseStorageSummary() (DatabaseStorageSummary, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return DatabaseStorageSummary{}, errors.New("repository not initialized")
|
||||
}
|
||||
switch r.db.Dialector.Name() {
|
||||
case "sqlite":
|
||||
size, err := sqliteDatabaseFileSize(r.dbPath)
|
||||
if err != nil { return DatabaseStorageSummary{}, err }
|
||||
return DatabaseStorageSummary{"sqlite", size, formatDatabaseSize(size)}, nil
|
||||
case "postgres":
|
||||
var size int64
|
||||
if err := r.db.Raw("SELECT pg_database_size(current_database())").Scan(&size).Error; err != nil { return DatabaseStorageSummary{}, err }
|
||||
return DatabaseStorageSummary{"postgres", size, formatDatabaseSize(size)}, nil
|
||||
default:
|
||||
return DatabaseStorageSummary{}, fmt.Errorf("unsupported database dialect %q", r.db.Dialector.Name())
|
||||
}
|
||||
}
|
||||
|
||||
func sqliteDatabaseFileSize(path string) (int64, error) {
|
||||
if path == "" || path == ":memory:" { return 0, nil }
|
||||
var total int64
|
||||
for _, candidate := range []string{path, path + "-wal", path + "-shm"} {
|
||||
info, err := os.Stat(candidate)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) { continue }
|
||||
return 0, err
|
||||
}
|
||||
if !info.IsDir() { total += info.Size() }
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func formatDatabaseSize(bytes int64) string {
|
||||
if bytes < 1024 { return fmt.Sprintf("%d B", bytes) }
|
||||
units := []string{"KB", "MB", "GB", "TB"}
|
||||
value := float64(bytes) / 1024
|
||||
for _, unit := range units {
|
||||
if value < 1024 || unit == "TB" { return fmt.Sprintf("%.1f %s", value, unit) }
|
||||
value /= 1024
|
||||
}
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Add API handler and route**
|
||||
|
||||
Create `go-backend/internal/http/handler/storage.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func (h *Handler) storageSummary(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if h == nil || h.repo == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "repository not initialized"))
|
||||
return
|
||||
}
|
||||
summary, err := h.repo.DatabaseStorageSummary()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(summary))
|
||||
}
|
||||
```
|
||||
|
||||
Register in `Handler.Register`: `mux.HandleFunc("/api/v1/system/storage", h.storageSummary)`.
|
||||
|
||||
In `requiresAdmin`, add:
|
||||
|
||||
```go
|
||||
if strings.HasPrefix(path, "/api/v1/system/") {
|
||||
return true
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Write contract test for auth and shape**
|
||||
|
||||
Create `go-backend/tests/contract/storage_contract_test.go` with a test that sends GET `/api/v1/system/storage` as non-admin and expects `403`, then as admin and expects `code == 0`, `dbType`, numeric `databaseSizeBytes`, and `databaseSizeText`.
|
||||
|
||||
- [ ] **Step 6: Run storage tests**
|
||||
|
||||
Run: `go test ./internal/store/repo ./tests/contract -run 'TestDatabaseStorageSummarySQLiteIncludesSize|TestFormatDatabaseSize|TestStorageSummaryRequiresAdminAndReturnsSize' -count=1`
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Add Frontend Config UI
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/api/types.ts`
|
||||
- Modify: `vite-frontend/src/api/index.ts`
|
||||
- Modify: `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
- [ ] **Step 1: Add API type and function**
|
||||
|
||||
In `types.ts` add:
|
||||
|
||||
```ts
|
||||
export interface StorageSummaryApiData {
|
||||
dbType: string;
|
||||
databaseSizeBytes: number;
|
||||
databaseSizeText: string;
|
||||
}
|
||||
```
|
||||
|
||||
In `index.ts`, import `StorageSummaryApiData` and add:
|
||||
|
||||
```ts
|
||||
export const getStorageSummary = () =>
|
||||
Network.get<StorageSummaryApiData>("/system/storage");
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Add retention config item**
|
||||
|
||||
In `config.tsx`, add to `CONFIG_ITEMS` near monitoring:
|
||||
|
||||
```ts
|
||||
{
|
||||
key: "monitor_retention_days",
|
||||
label: "监控数据保留天数",
|
||||
placeholder: "7",
|
||||
description:
|
||||
"统一清理节点指标、隧道流量、服务监控结果和隧道质量历史;默认 7 天。",
|
||||
type: "input",
|
||||
},
|
||||
```
|
||||
|
||||
Add `"monitor_retention_days"` to `getInitialConfigs()` keys.
|
||||
|
||||
- [ ] **Step 3: Fetch and display database size**
|
||||
|
||||
In `config.tsx`, add state:
|
||||
|
||||
```ts
|
||||
const [storageSummary, setStorageSummary] = useState<string>("加载中...");
|
||||
```
|
||||
|
||||
Add a load effect:
|
||||
|
||||
```ts
|
||||
useEffect(() => {
|
||||
let mounted = true;
|
||||
getStorageSummary()
|
||||
.then((response) => {
|
||||
if (!mounted) return;
|
||||
if (response.code === 0 && response.data?.databaseSizeText) {
|
||||
setStorageSummary(response.data.databaseSizeText);
|
||||
} else {
|
||||
setStorageSummary("获取失败");
|
||||
}
|
||||
})
|
||||
.catch(() => {
|
||||
if (mounted) setStorageSummary("获取失败");
|
||||
});
|
||||
return () => {
|
||||
mounted = false;
|
||||
};
|
||||
}, []);
|
||||
```
|
||||
|
||||
Render inside the basic settings card before the save button:
|
||||
|
||||
```tsx
|
||||
<Divider className="my-2" />
|
||||
<div className="space-y-1">
|
||||
<p className="text-sm font-medium text-gray-700 dark:text-gray-300">
|
||||
数据库占用
|
||||
</p>
|
||||
<p className="text-xs text-gray-500 dark:text-gray-400">
|
||||
当前后端数据库文件/实例占用空间,仅用于容量参考。
|
||||
</p>
|
||||
<div className="rounded-lg border border-divider bg-default-50/60 dark:bg-default-100/10 px-4 py-3 text-sm font-semibold text-default-800 dark:text-default-200">
|
||||
{storageSummary}
|
||||
</div>
|
||||
</div>
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Build frontend**
|
||||
|
||||
Run: `pnpm run build` from `vite-frontend`.
|
||||
|
||||
Expected: TypeScript and Vite build pass.
|
||||
|
||||
---
|
||||
|
||||
### Task 5: Final Verification
|
||||
|
||||
**Files:**
|
||||
- All files changed by previous tasks.
|
||||
|
||||
- [ ] **Step 1: Run backend tests**
|
||||
|
||||
Run: `go test ./...` from `go-backend`.
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 2: Run frontend build**
|
||||
|
||||
Run: `pnpm run build` from `vite-frontend`.
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Review diff**
|
||||
|
||||
Run: `git diff --stat` and `git diff -- docs/superpowers/specs/2026-04-28-monitoring-retention-storage-design.md docs/superpowers/plans/2026-04-28-monitoring-retention-storage.md go-backend vite-frontend`.
|
||||
|
||||
Expected: Diff is limited to retention config, storage summary, tests, and config UI.
|
||||
|
||||
---
|
||||
|
||||
## Self-Review
|
||||
|
||||
- Spec coverage: retention config, uniform cleanup, storage summary API, frontend display, validation, and verification are covered.
|
||||
- Placeholder scan: no TBD/TODO placeholders; the one contract-test step describes exact assertions even though the surrounding helper functions already exist in contract tests.
|
||||
- Type consistency: backend JSON fields match frontend `StorageSummaryApiData` exactly: `dbType`, `databaseSizeBytes`, `databaseSizeText`.
|
||||
@@ -0,0 +1,892 @@
|
||||
# Best Exit Current Display Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Show the currently applied `best` exit selection in the tunnel list information, including per-entry/per-final-hop details for multi-owner tunnels.
|
||||
|
||||
**Architecture:** Add a backend-only display layer that snapshots `bestExitManager` state and attaches `bestExitState` to existing `tunnelList`/`tunnelGet` responses. Render that state in the existing tunnel table/grid topology area using compact text and a native `title` detail tooltip. No routing, scoring, persistence, polling, or runtime update behavior changes.
|
||||
|
||||
**Tech Stack:** Go `net/http` handlers + existing repository methods, React/TypeScript in `vite-frontend/src/pages/tunnel.tsx`, Tailwind/shadcn bridge components already in the file.
|
||||
|
||||
---
|
||||
|
||||
## File Structure
|
||||
|
||||
- Create `go-backend/internal/http/handler/tunnel_best_exit_display.go`: response DTOs, manager snapshot method, tunnel-response parsing helpers, and `Handler.attachBestExitStates`.
|
||||
- Create `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`: backend display-state unit tests.
|
||||
- Modify `go-backend/internal/http/handler/handler.go`: call `h.attachBestExitStatesOrLog(items)` in `tunnelList`.
|
||||
- Modify `go-backend/internal/http/handler/mutations.go`: call `h.attachBestExitStatesOrLog(items)` before returning a single tunnel in `tunnelGet`.
|
||||
- Modify `vite-frontend/src/pages/tunnel.tsx`: add `bestExitState` types, map API state, helper render functions, and table/grid display.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Backend Snapshot And Display-State Tests
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`
|
||||
|
||||
- [ ] **Step 1: Write failing backend display tests**
|
||||
|
||||
Create `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestBestExitDecisionSnapshotIsDefensiveCopy(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
score := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30, NodeName: "exit-a"}, 10, 0, 20, 0)
|
||||
|
||||
m.observeScores(key, []bestExitCandidateScore{score}, now)
|
||||
snapshot, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected snapshot")
|
||||
}
|
||||
if snapshot.AppliedExitNodeID != 30 || snapshot.UpdatedAt != now.UnixMilli() {
|
||||
t.Fatalf("unexpected snapshot: %+v", snapshot)
|
||||
}
|
||||
if len(snapshot.Scores) != 1 {
|
||||
t.Fatalf("expected one score in snapshot, got %+v", snapshot.Scores)
|
||||
}
|
||||
snapshot.Scores[0].ExitNodeID = 99
|
||||
|
||||
again, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected second snapshot")
|
||||
}
|
||||
if again.Scores[0].ExitNodeID != 30 {
|
||||
t.Fatalf("snapshot score mutation leaked into manager state: %+v", again.Scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForDirectMultiEntryOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(100, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 11}, 31, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
{"nodeId": int64(11)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 11: "入口 B", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if !state.Enabled || state.Summary != "多个出口" || state.Status != "applied" {
|
||||
t.Fatalf("unexpected state summary: %+v", state)
|
||||
}
|
||||
if state.UpdatedAt != now.Add(time.Second).UnixMilli() {
|
||||
t.Fatalf("expected latest updatedAt, got %d", state.UpdatedAt)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "entry" || state.Items[0].OwnerNodeName != "入口 A" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "entry" || state.Items[1].OwnerNodeName != "入口 B" || state.Items[1].ExitNodeName != "日本节点" {
|
||||
t.Fatalf("unexpected second item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForFinalChainHopOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(200, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 20}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 21}, 30, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(88),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{
|
||||
{{"nodeId": int64(15), "inx": int64(0)}},
|
||||
{{"nodeId": int64(20), "inx": int64(1)}, {"nodeId": int64(21), "inx": int64(1)}},
|
||||
},
|
||||
}
|
||||
names := map[int64]string{20: "中转 M1", 21: "中转 M2", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Summary != "香港节点" || state.Status != "applied" {
|
||||
t.Fatalf("expected single-exit summary, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two final-hop owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "chain" || state.Items[0].OwnerNodeName != "中转 M1" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first chain owner item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "chain" || state.Items[1].OwnerNodeName != "中转 M2" || state.Items[1].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected second chain owner item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateWaitingWhenNoAppliedDecisionExists(t *testing.T) {
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, newBestExitManager(), testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected waiting best exit state")
|
||||
}
|
||||
if state.Summary != "等待探测" || state.Status != "waiting" {
|
||||
t.Fatalf("expected waiting state, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 1 || state.Items[0].ExitNodeID != 0 || state.Items[0].ExitNodeName != "等待探测" {
|
||||
t.Fatalf("unexpected waiting item: %+v", state.Items)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateSkipsNonBestAndSingleExitTunnels(t *testing.T) {
|
||||
nonBest := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": "round"},
|
||||
{"nodeId": int64(31), "strategy": "round"},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(nonBest, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected non-best tunnel to skip state, got %+v", state)
|
||||
}
|
||||
|
||||
singleExit := map[string]interface{}{
|
||||
"id": int64(78),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(singleExit, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected single-exit tunnel to skip state, got %+v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func testBestExitNameLookup(names map[int64]string) bestExitNodeNameLookup {
|
||||
return func(nodeID int64) (string, bool) {
|
||||
name := names[nodeID]
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run backend display tests to verify failure**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run 'TestBestExitDecisionSnapshot|TestBuildBestExitDisplayState' -count=1
|
||||
```
|
||||
|
||||
Expected: FAIL with undefined `snapshot`, `buildBestExitDisplayState`, and `bestExitNodeNameLookup`.
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Backend Display State Implementation
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/tunnel_best_exit_display.go`
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_best_exit.go`
|
||||
- Test: `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`
|
||||
|
||||
- [ ] **Step 1: Implement display state and snapshot helpers**
|
||||
|
||||
Create `go-backend/internal/http/handler/tunnel_best_exit_display.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
bestExitDisplayStatusApplied = "applied"
|
||||
bestExitDisplayStatusWaiting = "waiting"
|
||||
bestExitDisplaySummaryMulti = "多个出口"
|
||||
bestExitDisplaySummaryWait = "等待探测"
|
||||
bestExitUnknownExitName = "未知出口"
|
||||
bestExitUnknownEntryName = "未知入口"
|
||||
bestExitUnknownChainName = "未知中转"
|
||||
)
|
||||
|
||||
type bestExitDecisionSnapshot struct {
|
||||
AppliedExitNodeID int64
|
||||
UpdatedAt int64
|
||||
Reason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitDisplayState struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Summary string `json:"summary"`
|
||||
Status string `json:"status"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
Items []bestExitDisplayItem `json:"items"`
|
||||
}
|
||||
|
||||
type bestExitDisplayItem struct {
|
||||
OwnerNodeID int64 `json:"ownerNodeId"`
|
||||
OwnerNodeName string `json:"ownerNodeName"`
|
||||
OwnerRole string `json:"ownerRole"`
|
||||
ExitNodeID int64 `json:"exitNodeId,omitempty"`
|
||||
ExitNodeName string `json:"exitNodeName"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
}
|
||||
|
||||
type bestExitNodeNameLookup func(nodeID int64) (string, bool)
|
||||
|
||||
func (m *bestExitManager) snapshot(key bestExitOwnerKey) (bestExitDecisionSnapshot, bool) {
|
||||
if m == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisions[key]
|
||||
if d == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
updatedAt := int64(0)
|
||||
if !d.LastSwitchAt.IsZero() {
|
||||
updatedAt = d.LastSwitchAt.UnixMilli()
|
||||
}
|
||||
return bestExitDecisionSnapshot{
|
||||
AppliedExitNodeID: d.AppliedExitNodeID,
|
||||
UpdatedAt: updatedAt,
|
||||
Reason: d.LastReason,
|
||||
Scores: cloneBestExitScores(d.Scores),
|
||||
}, true
|
||||
}
|
||||
|
||||
func (h *Handler) attachBestExitStates(items []map[string]interface{}) {
|
||||
if h == nil || len(items) == 0 {
|
||||
return
|
||||
}
|
||||
lookup := h.bestExitNodeNameLookup()
|
||||
for _, item := range items {
|
||||
state, ok := buildBestExitDisplayState(item, h.bestExit, lookup)
|
||||
if !ok {
|
||||
delete(item, "bestExitState")
|
||||
continue
|
||||
}
|
||||
item["bestExitState"] = state
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) bestExitNodeNameLookup() bestExitNodeNameLookup {
|
||||
cache := map[int64]string{}
|
||||
return func(nodeID int64) (string, bool) {
|
||||
if nodeID <= 0 || h == nil {
|
||||
return "", false
|
||||
}
|
||||
if name, ok := cache[nodeID]; ok {
|
||||
return name, name != ""
|
||||
}
|
||||
node, err := h.getNodeRecord(nodeID)
|
||||
if err != nil || node == nil {
|
||||
cache[nodeID] = ""
|
||||
return "", false
|
||||
}
|
||||
name := strings.TrimSpace(node.Name)
|
||||
cache[nodeID] = name
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
|
||||
func buildBestExitDisplayState(tunnel map[string]interface{}, manager *bestExitManager, lookup bestExitNodeNameLookup) (*bestExitDisplayState, bool) {
|
||||
if tunnel == nil {
|
||||
return nil, false
|
||||
}
|
||||
tunnelID := asInt64(tunnel["id"], 0)
|
||||
outNodes := bestExitDisplayMapSlice(tunnel["outNodeId"])
|
||||
if tunnelID <= 0 || len(outNodes) <= 1 {
|
||||
return nil, false
|
||||
}
|
||||
if !isBestTunnelStrategy(asString(outNodes[0]["strategy"])) {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
owners, ownerRole := bestExitDisplayOwners(tunnel)
|
||||
state := &bestExitDisplayState{
|
||||
Enabled: true,
|
||||
Summary: bestExitDisplaySummaryWait,
|
||||
Status: bestExitDisplayStatusWaiting,
|
||||
Items: make([]bestExitDisplayItem, 0, len(owners)),
|
||||
}
|
||||
|
||||
exitsByID := map[int64]map[string]interface{}{}
|
||||
for _, exit := range outNodes {
|
||||
if id := asInt64(exit["nodeId"], 0); id > 0 {
|
||||
exitsByID[id] = exit
|
||||
}
|
||||
}
|
||||
appliedExitIDs := map[int64]string{}
|
||||
appliedCount := 0
|
||||
latestUpdatedAt := int64(0)
|
||||
latestReason := ""
|
||||
for _, owner := range owners {
|
||||
ownerNodeID := asInt64(owner["nodeId"], 0)
|
||||
if ownerNodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
item := bestExitDisplayItem{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
OwnerNodeName: bestExitDisplayNodeName(owner, ownerNodeID, lookup, bestExitUnknownOwnerName(ownerRole)),
|
||||
OwnerRole: ownerRole,
|
||||
ExitNodeName: bestExitDisplaySummaryWait,
|
||||
Reason: bestExitDisplayStatusWaiting,
|
||||
}
|
||||
if snapshot, ok := manager.snapshot(bestExitOwnerKey{TunnelID: tunnelID, OwnerNodeID: ownerNodeID}); ok && snapshot.AppliedExitNodeID > 0 {
|
||||
item.ExitNodeID = snapshot.AppliedExitNodeID
|
||||
item.ExitNodeName = bestExitDisplayNodeName(exitsByID[snapshot.AppliedExitNodeID], snapshot.AppliedExitNodeID, lookup, bestExitUnknownExitName)
|
||||
item.UpdatedAt = snapshot.UpdatedAt
|
||||
item.Reason = snapshot.Reason
|
||||
appliedExitIDs[item.ExitNodeID] = item.ExitNodeName
|
||||
appliedCount++
|
||||
if snapshot.UpdatedAt > latestUpdatedAt {
|
||||
latestUpdatedAt = snapshot.UpdatedAt
|
||||
latestReason = snapshot.Reason
|
||||
}
|
||||
}
|
||||
state.Items = append(state.Items, item)
|
||||
}
|
||||
|
||||
if appliedCount == 0 {
|
||||
return state, true
|
||||
}
|
||||
state.Status = bestExitDisplayStatusApplied
|
||||
state.UpdatedAt = latestUpdatedAt
|
||||
state.Reason = latestReason
|
||||
if len(appliedExitIDs) == 1 {
|
||||
for _, name := range appliedExitIDs {
|
||||
state.Summary = name
|
||||
}
|
||||
} else {
|
||||
state.Summary = bestExitDisplaySummaryMulti
|
||||
}
|
||||
return state, true
|
||||
}
|
||||
|
||||
func bestExitDisplayOwners(tunnel map[string]interface{}) ([]map[string]interface{}, string) {
|
||||
chainGroups := bestExitDisplayChainGroups(tunnel["chainNodes"])
|
||||
if len(chainGroups) > 0 {
|
||||
return chainGroups[len(chainGroups)-1], "chain"
|
||||
}
|
||||
return bestExitDisplayMapSlice(tunnel["inNodeId"]), "entry"
|
||||
}
|
||||
|
||||
func bestExitDisplayMapSlice(v interface{}) []map[string]interface{} {
|
||||
switch arr := v.(type) {
|
||||
case []map[string]interface{}:
|
||||
return arr
|
||||
case []interface{}:
|
||||
out := make([]map[string]interface{}, 0, len(arr))
|
||||
for _, item := range arr {
|
||||
if m, ok := item.(map[string]interface{}); ok {
|
||||
out = append(out, m)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayChainGroups(v interface{}) [][]map[string]interface{} {
|
||||
switch groups := v.(type) {
|
||||
case [][]map[string]interface{}:
|
||||
return groups
|
||||
case []interface{}:
|
||||
out := make([][]map[string]interface{}, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
items := bestExitDisplayMapSlice(group)
|
||||
if len(items) > 0 {
|
||||
out = append(out, items)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayNodeName(source map[string]interface{}, nodeID int64, lookup bestExitNodeNameLookup, fallback string) string {
|
||||
if source != nil {
|
||||
for _, key := range []string{"nodeName", "name"} {
|
||||
if name := strings.TrimSpace(asString(source[key])); name != "" {
|
||||
return name
|
||||
}
|
||||
}
|
||||
}
|
||||
if lookup != nil {
|
||||
if name, ok := lookup(nodeID); ok && strings.TrimSpace(name) != "" {
|
||||
return strings.TrimSpace(name)
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func bestExitUnknownOwnerName(role string) string {
|
||||
if role == "chain" {
|
||||
return bestExitUnknownChainName
|
||||
}
|
||||
return bestExitUnknownEntryName
|
||||
}
|
||||
|
||||
func (h *Handler) attachBestExitStatesOrLog(items []map[string]interface{}) {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
log.Printf("best_exit: attach display state failed: %v", recovered)
|
||||
}
|
||||
}()
|
||||
h.attachBestExitStates(items)
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Replace direct attach calls with panic-safe wrapper**
|
||||
|
||||
Keep `attachBestExitStates` for tests, and use `attachBestExitStatesOrLog` from handlers in Task 3. This step only creates the function above; no handler wiring yet.
|
||||
|
||||
- [ ] **Step 3: Run backend display tests**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run 'TestBestExitDecisionSnapshot|TestBuildBestExitDisplayState' -count=1
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 4: Run gofmt**
|
||||
|
||||
```bash
|
||||
gofmt -w internal/http/handler/tunnel_best_exit_display.go internal/http/handler/tunnel_best_exit_display_test.go
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Commit backend display implementation**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/tunnel_best_exit_display.go go-backend/internal/http/handler/tunnel_best_exit_display_test.go
|
||||
git commit -m "feat: build best exit display state"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Attach Best-Exit State To Tunnel List And Get Responses
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/handler/mutations.go`
|
||||
- Test: `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`
|
||||
|
||||
- [ ] **Step 1: Write failing handler attach tests**
|
||||
|
||||
Append to `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`:
|
||||
|
||||
```go
|
||||
func TestAttachBestExitStatesAddsStateToBestTunnelOnly(t *testing.T) {
|
||||
h := &Handler{bestExit: newBestExitManager()}
|
||||
now := time.Unix(300, 0)
|
||||
h.bestExit.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
|
||||
items := []map[string]interface{}{
|
||||
{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
},
|
||||
{
|
||||
"id": int64(78),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(12)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(40), "strategy": "round"},
|
||||
{"nodeId": int64(41), "strategy": "round"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
h.attachBestExitStates(items)
|
||||
state, ok := items[0]["bestExitState"].(*bestExitDisplayState)
|
||||
if !ok {
|
||||
t.Fatalf("expected bestExitState on best tunnel, got %#v", items[0]["bestExitState"])
|
||||
}
|
||||
if state.Summary != bestExitUnknownExitName || state.Items[0].ExitNodeID != 30 {
|
||||
t.Fatalf("unexpected state with fallback names: %+v", state)
|
||||
}
|
||||
if _, exists := items[1]["bestExitState"]; exists {
|
||||
t.Fatalf("non-best tunnel should not have bestExitState: %+v", items[1])
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run attach test to verify failure**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run TestAttachBestExitStatesAddsStateToBestTunnelOnly -count=1
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Wire tunnel list response**
|
||||
|
||||
In `go-backend/internal/http/handler/handler.go`, change `tunnelList` from:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.attachBestExitStatesOrLog(items)
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Wire single tunnel response**
|
||||
|
||||
In `go-backend/internal/http/handler/mutations.go`, change `tunnelGet` from:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
for _, it := range items {
|
||||
if asInt64(it["id"], 0) == id {
|
||||
response.WriteJSON(w, response.OK(it))
|
||||
return
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.attachBestExitStatesOrLog(items)
|
||||
for _, it := range items {
|
||||
if asInt64(it["id"], 0) == id {
|
||||
response.WriteJSON(w, response.OK(it))
|
||||
return
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Run focused backend tests**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run 'TestBestExitDecisionSnapshot|TestBuildBestExitDisplayState|TestAttachBestExitStatesAddsStateToBestTunnelOnly' -count=1
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 6: Run gofmt**
|
||||
|
||||
```bash
|
||||
gofmt -w internal/http/handler/handler.go internal/http/handler/mutations.go internal/http/handler/tunnel_best_exit_display.go internal/http/handler/tunnel_best_exit_display_test.go
|
||||
```
|
||||
|
||||
- [ ] **Step 7: Commit response wiring**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/mutations.go go-backend/internal/http/handler/tunnel_best_exit_display.go go-backend/internal/http/handler/tunnel_best_exit_display_test.go
|
||||
git commit -m "feat: expose best exit display state"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Frontend Tunnel List Display
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/pages/tunnel.tsx`
|
||||
|
||||
- [ ] **Step 1: Add TypeScript types**
|
||||
|
||||
In `vite-frontend/src/pages/tunnel.tsx`, add these interfaces after `interface ChainTunnel`:
|
||||
|
||||
```ts
|
||||
interface BestExitStateItem {
|
||||
ownerNodeId: number;
|
||||
ownerNodeName: string;
|
||||
ownerRole: "entry" | "chain";
|
||||
exitNodeId?: number;
|
||||
exitNodeName: string;
|
||||
updatedAt?: number;
|
||||
reason?: string;
|
||||
}
|
||||
|
||||
interface BestExitState {
|
||||
enabled: boolean;
|
||||
summary: string;
|
||||
status: "applied" | "waiting";
|
||||
updatedAt?: number;
|
||||
reason?: string;
|
||||
items: BestExitStateItem[];
|
||||
}
|
||||
```
|
||||
|
||||
Then add the optional field to `interface Tunnel`:
|
||||
|
||||
```ts
|
||||
bestExitState?: BestExitState | null;
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Preserve API state during mapping**
|
||||
|
||||
In `mapTunnelApiItems`, add `bestExitState` to the returned object:
|
||||
|
||||
```ts
|
||||
bestExitState:
|
||||
tunnel.bestExitState && typeof tunnel.bestExitState === "object"
|
||||
? {
|
||||
...tunnel.bestExitState,
|
||||
items: Array.isArray(tunnel.bestExitState.items)
|
||||
? tunnel.bestExitState.items
|
||||
: [],
|
||||
}
|
||||
: null,
|
||||
```
|
||||
|
||||
The mapped object should include this field before `createdTime` or immediately after it.
|
||||
|
||||
- [ ] **Step 3: Add render helpers**
|
||||
|
||||
Add these helper functions after `mapTunnelApiItems` and before `export default function TunnelPage()`:
|
||||
|
||||
```tsx
|
||||
const bestExitOwnerRoleText = (role: BestExitStateItem["ownerRole"]) => {
|
||||
return role === "chain" ? "中转" : "入口";
|
||||
};
|
||||
|
||||
const bestExitDetailTitle = (state?: BestExitState | null) => {
|
||||
if (!state?.enabled || !state.items?.length) {
|
||||
return "";
|
||||
}
|
||||
return state.items
|
||||
.map((item) => {
|
||||
const ownerName = item.ownerNodeName || `${bestExitOwnerRoleText(item.ownerRole)} ${item.ownerNodeId}`;
|
||||
const exitName = item.exitNodeName || "等待探测";
|
||||
return `${ownerName} -> ${exitName}`;
|
||||
})
|
||||
.join("\n");
|
||||
};
|
||||
|
||||
const renderBestExitState = (state?: BestExitState | null) => {
|
||||
if (!state?.enabled) {
|
||||
return null;
|
||||
}
|
||||
const title = bestExitDetailTitle(state);
|
||||
const isWaiting = state.status === "waiting";
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`mt-1 text-[11px] leading-4 ${
|
||||
isWaiting
|
||||
? "text-default-500"
|
||||
: "text-emerald-700 dark:text-emerald-300"
|
||||
}`}
|
||||
title={title || undefined}
|
||||
>
|
||||
最优出口:{state.summary || "等待探测"}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Render in table topology cell**
|
||||
|
||||
In the table topology `<TableCell>` around line 1674, change the cell content from:
|
||||
|
||||
```tsx
|
||||
<div className="flex items-center gap-1.5 text-xs">
|
||||
<span className="font-semibold text-primary-700 dark:text-primary-400">
|
||||
{tunnel.inNodeId?.length || 0}入口
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-secondary-700 dark:text-secondary-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.chainNodes?.length || 0
|
||||
: 0}
|
||||
跳
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-success-700 dark:text-success-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.outNodeId?.length || 0
|
||||
: tunnel.inNodeId?.length || 0}
|
||||
出口
|
||||
</span>
|
||||
</div>
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```tsx
|
||||
<div>
|
||||
<div className="flex items-center gap-1.5 text-xs">
|
||||
<span className="font-semibold text-primary-700 dark:text-primary-400">
|
||||
{tunnel.inNodeId?.length || 0}入口
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-secondary-700 dark:text-secondary-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.chainNodes?.length || 0
|
||||
: 0}
|
||||
跳
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-success-700 dark:text-success-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.outNodeId?.length || 0
|
||||
: tunnel.inNodeId?.length || 0}
|
||||
出口
|
||||
</span>
|
||||
</div>
|
||||
{renderBestExitState(tunnel.bestExitState)}
|
||||
</div>
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Render in grid card topology section**
|
||||
|
||||
In the grid card topology section, after the closing `</div>` for the topology row at the end of the block containing `出口` and before the enclosing border section closes, add:
|
||||
|
||||
```tsx
|
||||
<div className="text-center">
|
||||
{renderBestExitState(tunnel.bestExitState)}
|
||||
</div>
|
||||
```
|
||||
|
||||
The result should put the best-exit summary under the entry -> hop -> exit row inside the topology section.
|
||||
|
||||
- [ ] **Step 6: Run frontend build**
|
||||
|
||||
Run from `vite-frontend`:
|
||||
|
||||
```bash
|
||||
pnpm run build
|
||||
```
|
||||
|
||||
Expected: PASS with `tsc && vite build` completing successfully.
|
||||
|
||||
- [ ] **Step 7: Commit frontend display**
|
||||
|
||||
```bash
|
||||
git add vite-frontend/src/pages/tunnel.tsx
|
||||
git commit -m "feat: show current best exit in tunnel list"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 5: Full Verification And Review
|
||||
|
||||
**Files:**
|
||||
- Verify only.
|
||||
|
||||
- [ ] **Step 1: Run backend tests**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./...
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 2: Run frontend build**
|
||||
|
||||
Run from `vite-frontend`:
|
||||
|
||||
```bash
|
||||
pnpm run build
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Inspect final diff**
|
||||
|
||||
Run from repository root:
|
||||
|
||||
```bash
|
||||
git diff --stat origin/main...HEAD
|
||||
git diff -- go-backend/internal/http/handler/tunnel_best_exit_display.go go-backend/internal/http/handler/tunnel_best_exit_display_test.go go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/mutations.go vite-frontend/src/pages/tunnel.tsx
|
||||
```
|
||||
|
||||
Expected: Diff only adds best-exit display state, response attachment, frontend list display, and tests. It must not change best-exit scoring, switching, runtime chain update, or agent code.
|
||||
|
||||
- [ ] **Step 4: Request final code review**
|
||||
|
||||
Ask a reviewer to check:
|
||||
|
||||
```text
|
||||
Review the best-exit current display implementation. Confirm it only exposes current in-memory best-exit state in tunnel list/get responses and renders it in the tunnel list. Verify it does not change routing, scoring, switching, persistence, or polling behavior.
|
||||
```
|
||||
|
||||
Expected: No blocking findings.
|
||||
|
||||
---
|
||||
|
||||
## Self-Review
|
||||
|
||||
- Spec coverage: Backend response state is Task 2 and Task 3; direct vs final-hop owner semantics are covered by Task 1 tests; frontend list/grid display is Task 4; no polling and no routing changes are preserved by Task 5 review instructions.
|
||||
- Placeholder scan: The plan contains concrete files, function names, code blocks, commands, and expected outcomes.
|
||||
- Type consistency: `BestExitState`, `BestExitStateItem`, `bestExitDisplayState`, `bestExitDisplayItem`, `bestExitDecisionSnapshot`, and `bestExitNodeNameLookup` are defined before use and names match across tasks.
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,536 @@
|
||||
# Dependabot Remediation Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Resolve the open Dependabot dependency alerts for `go-backend`, `vite-frontend`, `go-gost/x`, and `go-gost` without mixing in unrelated business security changes.
|
||||
|
||||
**Architecture:** Apply targeted dependency upgrades per module, verify each module before moving to the next, and keep commits scoped to one dependency group. `go-gost/x` is fixed before `go-gost` because the main agent module uses `replace github.com/go-gost/x => ./x`.
|
||||
|
||||
**Tech Stack:** Go modules, pnpm, Vite/Rolldown, GitHub CLI Dependabot alerts API.
|
||||
|
||||
---
|
||||
|
||||
## File Map
|
||||
|
||||
- Modify: `go-backend/go.mod`
|
||||
Responsibility: update `github.com/jackc/pgx/v5` to the patched version.
|
||||
- Modify: `go-backend/go.sum`
|
||||
Responsibility: reflect Go module checksum changes from the pgx upgrade.
|
||||
- Modify: `vite-frontend/package.json`
|
||||
Responsibility: update direct vulnerable npm dependency versions and configure `pnpm.overrides`.
|
||||
- Modify: `vite-frontend/pnpm-lock.yaml`
|
||||
Responsibility: resolve vulnerable npm transitive dependencies to patched versions.
|
||||
- Modify: `go-gost/x/go.mod`
|
||||
Responsibility: update vulnerable Go dependencies used by the local `github.com/go-gost/x` module.
|
||||
- Modify: `go-gost/x/go.sum`
|
||||
Responsibility: reflect checksum changes for `go-gost/x`.
|
||||
- Modify: `go-gost/x/dialer/dtls/dialer.go`
|
||||
Responsibility: migrate DTLS import path from `github.com/pion/dtls/v2` to `github.com/pion/dtls/v3`.
|
||||
- Modify: `go-gost/x/listener/dtls/listener.go`
|
||||
Responsibility: migrate DTLS import path from `github.com/pion/dtls/v2` to `github.com/pion/dtls/v3`.
|
||||
- Modify: `go-gost/go.mod`
|
||||
Responsibility: sync vulnerable dependency versions for the main agent module while preserving local `replace github.com/go-gost/x => ./x`.
|
||||
- Modify: `go-gost/go.sum`
|
||||
Responsibility: reflect checksum changes for the main agent module.
|
||||
|
||||
## Task 1: Capture Baseline Alerts
|
||||
|
||||
**Files:**
|
||||
- Read: GitHub Dependabot alerts API
|
||||
- Read: `go-backend/go.mod`
|
||||
- Read: `vite-frontend/package.json`
|
||||
- Read: `go-gost/x/go.mod`
|
||||
- Read: `go-gost/go.mod`
|
||||
|
||||
- [ ] **Step 1: Query current open Dependabot alerts**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
gh api 'repos/Sagit-chu/flvx/dependabot/alerts?state=open&per_page=100' --paginate \
|
||||
--jq '.[] | [.number,.security_advisory.severity,.dependency.package.ecosystem,.dependency.manifest_path,.dependency.package.name,.security_vulnerability.vulnerable_version_range,(.security_vulnerability.first_patched_version.identifier // "")] | @tsv'
|
||||
```
|
||||
|
||||
Expected: output includes alerts for `github.com/jackc/pgx/v5`, `postcss`, `serialize-javascript`, `fast-uri`, `@babel/plugin-transform-modules-systemjs`, `github.com/sirupsen/logrus`, `github.com/quic-go/quic-go`, `github.com/quic-go/webtransport-go`, and `github.com/pion/dtls/v2`.
|
||||
|
||||
- [ ] **Step 2: Confirm starting versions in module manifests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'jackc/pgx|postcss|serialize-javascript|pion/dtls|quic-go|webtransport-go|sirupsen/logrus' \
|
||||
go-backend/go.mod vite-frontend/package.json go-gost/x/go.mod go-gost/go.mod
|
||||
```
|
||||
|
||||
Expected key lines:
|
||||
|
||||
```text
|
||||
go-backend/go.mod: github.com/jackc/pgx/v5 v5.7.3
|
||||
vite-frontend/package.json: "postcss": "8.5.6"
|
||||
vite-frontend/package.json: "serialize-javascript": "7.0.3"
|
||||
go-gost/x/go.mod: github.com/pion/dtls/v2 v2.2.6
|
||||
go-gost/x/go.mod: github.com/quic-go/quic-go v0.49.1
|
||||
go-gost/x/go.mod: github.com/quic-go/webtransport-go v0.8.1-0.20241018022711-4ac2c9250e66
|
||||
go-gost/x/go.mod: github.com/sirupsen/logrus v1.8.1
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Confirm the DTLS advisory has no patched v2 release**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go list -m -versions github.com/pion/dtls/v2
|
||||
gh api 'advisories/GHSA-9f3f-wv7r-qc8r' --jq '{summary, vulnerabilities}'
|
||||
```
|
||||
|
||||
Expected:
|
||||
|
||||
```text
|
||||
github.com/pion/dtls/v2 ... v2.2.12
|
||||
```
|
||||
|
||||
Expected advisory facts:
|
||||
|
||||
```text
|
||||
github.com/pion/dtls/v2 vulnerable range <= 2.2.12 has no first_patched_version.
|
||||
github.com/pion/dtls/v3 patched versions include 3.0.11 and 3.1.1.
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Do not commit baseline capture**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
git status --short
|
||||
```
|
||||
|
||||
Expected: no files are changed by Task 1.
|
||||
|
||||
## Task 2: Fix go-backend pgx Alerts
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/go.mod`
|
||||
- Modify: `go-backend/go.sum`
|
||||
|
||||
- [ ] **Step 1: Upgrade pgx to the patched version**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go get github.com/jackc/pgx/v5@v5.9.2)
|
||||
```
|
||||
|
||||
Expected: `go-backend/go.mod` changes `github.com/jackc/pgx/v5` from `v5.7.3` to `v5.9.2`, and `go-backend/go.sum` updates checksums.
|
||||
|
||||
- [ ] **Step 2: Tidy backend module**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go mod tidy)
|
||||
```
|
||||
|
||||
Expected: command exits with code 0.
|
||||
|
||||
- [ ] **Step 3: Verify backend dependency version**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'github.com/jackc/pgx/v5' go-backend/go.mod
|
||||
```
|
||||
|
||||
Expected:
|
||||
|
||||
```text
|
||||
go-backend/go.mod: github.com/jackc/pgx/v5 v5.9.2
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Run backend tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
```
|
||||
|
||||
Expected: all backend packages pass.
|
||||
|
||||
- [ ] **Step 5: Commit backend dependency fix**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
git add go-backend/go.mod go-backend/go.sum
|
||||
git commit -m "fix: update backend pgx dependency"
|
||||
```
|
||||
|
||||
Expected: one commit containing only `go-backend/go.mod` and `go-backend/go.sum`.
|
||||
|
||||
## Task 3: Fix Frontend npm Alerts
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/package.json`
|
||||
- Modify: `vite-frontend/pnpm-lock.yaml`
|
||||
|
||||
- [ ] **Step 1: Update direct dependency and pnpm overrides in package.json**
|
||||
|
||||
Edit `vite-frontend/package.json` so the relevant entries are exactly:
|
||||
|
||||
```json
|
||||
{
|
||||
"devDependencies": {
|
||||
"postcss": "8.5.10"
|
||||
},
|
||||
"pnpm": {
|
||||
"overrides": {
|
||||
"@babel/plugin-transform-modules-systemjs": "7.29.4",
|
||||
"fast-uri": "3.1.2",
|
||||
"serialize-javascript": "7.0.5"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Remove the existing top-level `"overrides"` block after adding `"pnpm.overrides"`. Keep all other existing dependencies and scripts unchanged.
|
||||
|
||||
- [ ] **Step 2: Regenerate pnpm lockfile**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd vite-frontend && pnpm install)
|
||||
```
|
||||
|
||||
Expected: `vite-frontend/pnpm-lock.yaml` updates and install exits with code 0.
|
||||
|
||||
- [ ] **Step 3: Verify vulnerable npm versions are absent**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'postcss@8\.5\.[0-9]:|"postcss":\s*"8\.5\.[0-9]"|serialize-javascript@[0-6]\.|serialize-javascript@7\.0\.[0-4]|fast-uri@3\.1\.[0-1]|plugin-transform-modules-systemjs@7\.29\.[0-3]' vite-frontend/pnpm-lock.yaml vite-frontend/package.json
|
||||
```
|
||||
|
||||
Expected: no output.
|
||||
|
||||
- [ ] **Step 4: Verify patched npm versions are present**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'postcss@8\.5\.10|serialize-javascript@7\.0\.5|fast-uri@3\.1\.2|plugin-transform-modules-systemjs@7\.29\.4' vite-frontend/pnpm-lock.yaml vite-frontend/package.json
|
||||
```
|
||||
|
||||
Expected: output includes patched entries for `postcss@8.5.10`, `serialize-javascript@7.0.5`, `fast-uri@3.1.2`, and `@babel/plugin-transform-modules-systemjs@7.29.4`.
|
||||
|
||||
- [ ] **Step 5: Build frontend**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
Expected: TypeScript and Rolldown/Vite build complete successfully.
|
||||
|
||||
- [ ] **Step 6: Commit frontend dependency fix**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
git add vite-frontend/package.json vite-frontend/pnpm-lock.yaml
|
||||
git commit -m "fix: update frontend vulnerable dependencies"
|
||||
```
|
||||
|
||||
Expected: one commit containing only `vite-frontend/package.json` and `vite-frontend/pnpm-lock.yaml`.
|
||||
|
||||
## Task 4: Fix go-gost/x Non-DTLS Alerts
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-gost/x/go.mod`
|
||||
- Modify: `go-gost/x/go.sum`
|
||||
|
||||
- [ ] **Step 1: Upgrade non-DTLS vulnerable Go dependencies**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost/x && go get github.com/sirupsen/logrus@v1.8.3 github.com/quic-go/quic-go@v0.57.0 github.com/quic-go/webtransport-go@v0.10.0)
|
||||
```
|
||||
|
||||
Expected: `go-gost/x/go.mod` resolves these dependencies to at least:
|
||||
|
||||
```text
|
||||
github.com/sirupsen/logrus v1.8.3
|
||||
github.com/quic-go/quic-go v0.57.0
|
||||
github.com/quic-go/webtransport-go v0.10.0
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Tidy go-gost/x module**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost/x && go mod tidy)
|
||||
```
|
||||
|
||||
Expected: command exits with code 0.
|
||||
|
||||
- [ ] **Step 3: Verify go-gost/x non-DTLS dependency versions**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'github.com/sirupsen/logrus|github.com/quic-go/quic-go|github.com/quic-go/webtransport-go' go-gost/x/go.mod
|
||||
```
|
||||
|
||||
Expected output contains versions at or above:
|
||||
|
||||
```text
|
||||
github.com/sirupsen/logrus v1.8.3
|
||||
github.com/quic-go/quic-go v0.57.0
|
||||
github.com/quic-go/webtransport-go v0.10.0
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Run go-gost/x tests after non-DTLS upgrades**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost/x && go test ./...)
|
||||
```
|
||||
|
||||
Expected: command exits with code 0. If it fails, stop this task before committing and inspect the first compiler error. The only permitted follow-up edits in this task are direct API-compatibility changes in files named by the compiler under `go-gost/x`; rerun this command after each edit.
|
||||
|
||||
- [ ] **Step 5: Commit go-gost/x non-DTLS dependency fix**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
git add go-gost/x/go.mod go-gost/x/go.sum
|
||||
git commit -m "fix: update gost quic dependencies"
|
||||
```
|
||||
|
||||
Expected: one commit containing `go-gost/x/go.mod` and `go-gost/x/go.sum`, plus only the compiler-named `go-gost/x` files edited during Step 4.
|
||||
|
||||
## Task 5: Migrate go-gost/x DTLS From v2 To v3
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-gost/x/go.mod`
|
||||
- Modify: `go-gost/x/go.sum`
|
||||
- Modify: `go-gost/x/dialer/dtls/dialer.go`
|
||||
- Modify: `go-gost/x/listener/dtls/listener.go`
|
||||
|
||||
- [ ] **Step 1: Update DTLS imports**
|
||||
|
||||
In `go-gost/x/dialer/dtls/dialer.go`, change:
|
||||
|
||||
```go
|
||||
"github.com/pion/dtls/v2"
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```go
|
||||
"github.com/pion/dtls/v3"
|
||||
```
|
||||
|
||||
In `go-gost/x/listener/dtls/listener.go`, change:
|
||||
|
||||
```go
|
||||
"github.com/pion/dtls/v2"
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```go
|
||||
"github.com/pion/dtls/v3"
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Add patched DTLS v3 module**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost/x && go get github.com/pion/dtls/v3@v3.0.11)
|
||||
```
|
||||
|
||||
Expected: `go-gost/x/go.mod` contains `github.com/pion/dtls/v3 v3.0.11` and no longer needs `github.com/pion/dtls/v2`.
|
||||
|
||||
- [ ] **Step 3: Tidy and format go-gost/x**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost/x && go mod tidy)
|
||||
gofmt -w go-gost/x/dialer/dtls/dialer.go go-gost/x/listener/dtls/listener.go
|
||||
```
|
||||
|
||||
Expected: command exits with code 0.
|
||||
|
||||
- [ ] **Step 4: Verify v2 import and module are removed**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'github.com/pion/dtls/v2' go-gost/x
|
||||
rg -n 'github.com/pion/dtls/v3' go-gost/x/go.mod go-gost/x/dialer/dtls/dialer.go go-gost/x/listener/dtls/listener.go
|
||||
```
|
||||
|
||||
Expected:
|
||||
|
||||
```text
|
||||
first command: no output
|
||||
second command: output includes go.mod, dialer.go, and listener.go
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Run go-gost/x tests after DTLS migration**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost/x && go test ./...)
|
||||
```
|
||||
|
||||
Expected: command exits with code 0. If the compiler reports DTLS v3 API errors, edit only `go-gost/x/dialer/dtls/dialer.go` and `go-gost/x/listener/dtls/listener.go`, preserving the existing `dtls.Config`, `dtls.ClientWithContext`, and `dtls.Listen` flow, then rerun this command.
|
||||
|
||||
- [ ] **Step 6: Commit DTLS migration**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
git add go-gost/x/go.mod go-gost/x/go.sum go-gost/x/dialer/dtls/dialer.go go-gost/x/listener/dtls/listener.go
|
||||
git commit -m "fix: migrate gost dtls dependency"
|
||||
```
|
||||
|
||||
Expected: one commit containing the DTLS import migration and Go module updates.
|
||||
|
||||
## Task 6: Sync go-gost Main Module
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-gost/go.mod`
|
||||
- Modify: `go-gost/go.sum`
|
||||
|
||||
- [ ] **Step 1: Upgrade main module vulnerable dependency requirements**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost && go get github.com/sirupsen/logrus@v1.8.3 github.com/quic-go/quic-go@v0.57.0 github.com/quic-go/webtransport-go@v0.10.0 github.com/pion/dtls/v3@v3.0.11)
|
||||
```
|
||||
|
||||
Expected: `go-gost/go.mod` resolves vulnerable dependencies to patched versions and preserves this replace directive:
|
||||
|
||||
```go
|
||||
replace github.com/go-gost/x => ./x
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Tidy main agent module**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost && go mod tidy)
|
||||
```
|
||||
|
||||
Expected: command exits with code 0.
|
||||
|
||||
- [ ] **Step 3: Verify go-gost no longer references vulnerable DTLS v2**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'github.com/pion/dtls/v2' go-gost/go.mod go-gost/go.sum
|
||||
rg -n 'github.com/pion/dtls/v3|github.com/quic-go/quic-go|github.com/quic-go/webtransport-go|github.com/sirupsen/logrus|replace github.com/go-gost/x => ./x' go-gost/go.mod
|
||||
```
|
||||
|
||||
Expected:
|
||||
|
||||
```text
|
||||
first command: no output
|
||||
second command: output includes dtls/v3, quic-go, webtransport-go, logrus, and the local replace directive
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Run go-gost tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost && go test ./...)
|
||||
```
|
||||
|
||||
Expected: all packages pass.
|
||||
|
||||
- [ ] **Step 5: Build go-gost binary**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-gost && go build .)
|
||||
```
|
||||
|
||||
Expected: build exits with code 0.
|
||||
|
||||
- [ ] **Step 6: Commit go-gost module sync**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
git add go-gost/go.mod go-gost/go.sum
|
||||
git commit -m "fix: sync gost main dependencies"
|
||||
```
|
||||
|
||||
Expected: one commit containing only `go-gost/go.mod` and `go-gost/go.sum`.
|
||||
|
||||
## Task 7: Final Dependabot Verification
|
||||
|
||||
**Files:**
|
||||
- Read: GitHub Dependabot alerts API
|
||||
- Read: Git working tree status
|
||||
|
||||
- [ ] **Step 1: Run all verification commands once more**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
(cd go-gost/x && go test ./...)
|
||||
(cd go-gost && go test ./...)
|
||||
(cd go-gost && go build .)
|
||||
```
|
||||
|
||||
Expected: every command exits with code 0.
|
||||
|
||||
- [ ] **Step 2: Query open Dependabot alerts after dependency updates**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
gh api 'repos/Sagit-chu/flvx/dependabot/alerts?state=open&per_page=100' --paginate \
|
||||
--jq 'group_by(.security_advisory.severity) | map({severity:.[0].security_advisory.severity,count:length})'
|
||||
```
|
||||
|
||||
Expected: counts are lower than the baseline from Task 1. If Dependabot has not rescanned yet, run the detailed query from Task 1 and confirm the manifest files now contain patched versions locally.
|
||||
|
||||
- [ ] **Step 3: Confirm no vulnerable dependency strings remain in manifests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
rg -n 'github.com/jackc/pgx/v5 v5\.7\.3|postcss\"\\s*:\\s*\"8\.5\.6|serialize-javascript\"\\s*:\\s*\"7\.0\.3|github.com/pion/dtls/v2|github.com/quic-go/quic-go v0\.49\.1|github.com/quic-go/webtransport-go v0\.8\.1|github.com/sirupsen/logrus v1\.8\.1' \
|
||||
go-backend/go.mod vite-frontend/package.json go-gost/x/go.mod go-gost/go.mod
|
||||
```
|
||||
|
||||
Expected: no output.
|
||||
|
||||
- [ ] **Step 4: Confirm working tree contains only intentional changes**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
git status --short
|
||||
```
|
||||
|
||||
Expected: no uncommitted files from this Dependabot remediation remain. Pre-existing unrelated files may still appear; do not stage or revert them.
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,650 @@
|
||||
# Forward Flow Reset Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Add a permission-checked action that resets only one forward rule's displayed upload and download counters.
|
||||
|
||||
**Architecture:** A dedicated repository method updates only the selected `forward` row. A dedicated authenticated handler reuses `resolveForwardAccess`, and the React page calls the endpoint from all three rule views through one confirmation modal.
|
||||
|
||||
**Tech Stack:** Go `net/http`, GORM, SQLite/PostgreSQL-compatible models, React, TypeScript, shadcn bridge components, Tailwind CSS v4.
|
||||
|
||||
## Global Constraints
|
||||
|
||||
- Only `forward.in_flow`, `forward.out_flow`, and `forward.updated_time` may change during reset.
|
||||
- Do not modify `user`, `user_tunnel`, quota, historical statistics, nftables counter state, or running services.
|
||||
- Administrators may reset any rule; non-admin users may reset only their own rules through existing `resolveForwardAccess` behavior.
|
||||
- All API responses must keep the `{code, msg, data, ts}` envelope.
|
||||
- Frontend imports must use `src/shadcn-bridge/heroui/*`; do not add `@heroui/*` or `@nextui-org/*` dependencies.
|
||||
- Do not add frontend test infrastructure.
|
||||
- Do not edit generated protobuf files, `install.sh`, or `panel_install.sh`.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Add the repository flow-reset primitive
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/store/repo/repository_forward_flow_reset_test.go`
|
||||
- Modify: `go-backend/internal/store/repo/repository_mutations.go`
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: `model.Forward`, the repository's GORM database handle, and an explicit Unix-millisecond timestamp.
|
||||
- Produces: `func (r *Repository) ResetForwardFlow(forwardID int64, now int64) error`.
|
||||
|
||||
- [ ] **Step 1: Write the failing repository tests**
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_forward_flow_reset_test.go`:
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResetForwardFlowOnlyUpdatesSelectedForward(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "forward-flow-reset.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
const originalUpdated int64 = 1000
|
||||
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, 'owner', 'pwd', 1, 0, 100, 700, 900, 0, 10, 1000, 1000, 1)
|
||||
`).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, 'tunnel', 1, 1, 'tls', 1, 1000, 1000, 1, NULL, 0)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, 2, 1, 10, 100, 500, 600, 0, 0, 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, 'owner', 'target', 1, '127.0.0.1:80', 'fifo', 111, 222, 1000, ?, 1, 0),
|
||||
(21, 2, 'owner', 'other', 1, '127.0.0.1:81', 'fifo', 333, 444, 1000, ?, 1, 1)
|
||||
`, originalUpdated, originalUpdated).Error; err != nil {
|
||||
t.Fatalf("insert forwards: %v", err)
|
||||
}
|
||||
|
||||
const resetAt int64 = 2000
|
||||
if err := r.ResetForwardFlow(20, resetAt); err != nil {
|
||||
t.Fatalf("ResetForwardFlow: %v", err)
|
||||
}
|
||||
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM forward WHERE id = 20", 0)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM forward WHERE id = 20", 0)
|
||||
assertForwardFlowResetValue(t, r, "SELECT updated_time FROM forward WHERE id = 20", resetAt)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM forward WHERE id = 21", 333)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM forward WHERE id = 21", 444)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM user WHERE id = 2", 700)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM user WHERE id = 2", 900)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM user_tunnel WHERE id = 10", 500)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM user_tunnel WHERE id = 10", 600)
|
||||
}
|
||||
|
||||
func TestResetForwardFlowRejectsUninitializedRepository(t *testing.T) {
|
||||
var r *Repository
|
||||
if err := r.ResetForwardFlow(20, 2000); err == nil {
|
||||
t.Fatal("expected uninitialized repository error")
|
||||
}
|
||||
}
|
||||
|
||||
func assertForwardFlowResetValue(t *testing.T, r *Repository, query string, want int64) {
|
||||
t.Helper()
|
||||
var got int64
|
||||
if err := r.DB().Raw(query).Scan(&got).Error; err != nil {
|
||||
t.Fatalf("query %q: %v", query, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("query %q returned %d, want %d", query, got, want)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the repository tests and verify the missing method failure**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./internal/store/repo -run TestResetForwardFlow -count=1
|
||||
```
|
||||
|
||||
Expected: compilation fails because `ResetForwardFlow` is undefined.
|
||||
|
||||
- [ ] **Step 3: Implement the minimal repository method**
|
||||
|
||||
Add to the flow-reset section of `go-backend/internal/store/repo/repository_mutations.go`:
|
||||
|
||||
```go
|
||||
func (r *Repository) ResetForwardFlow(forwardID int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.Forward{}).
|
||||
Where("id = ?", forwardID).
|
||||
Updates(map[string]interface{}{
|
||||
"in_flow": 0,
|
||||
"out_flow": 0,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
```
|
||||
|
||||
The file already imports `errors` and `model`; do not add a new dependency.
|
||||
|
||||
- [ ] **Step 4: Format and run the focused repository tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && gofmt -w internal/store/repo/repository_forward_flow_reset_test.go internal/store/repo/repository_mutations.go
|
||||
go test ./internal/store/repo -run TestResetForwardFlow -count=1
|
||||
```
|
||||
|
||||
Expected: both reset tests pass.
|
||||
|
||||
- [ ] **Step 5: Commit the repository change**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/store/repo/repository_mutations.go go-backend/internal/store/repo/repository_forward_flow_reset_test.go
|
||||
git commit -m "feat: add forward flow reset repository method"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Add the authenticated reset endpoint
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/forward_reset_flow_test.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/handler/mutations.go`
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: `POST` JSON `{ "id": number }`, `resolveForwardAccess`, and `Repository.ResetForwardFlow` from Task 1.
|
||||
- Produces: `POST /api/v1/forward/reset-flow` and `func (h *Handler) forwardResetFlow(http.ResponseWriter, *http.Request)`.
|
||||
|
||||
- [ ] **Step 1: Write the failing handler tests**
|
||||
|
||||
Create `go-backend/internal/http/handler/forward_reset_flow_test.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestForwardResetFlowPermissionsAndIsolation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
actorID int64
|
||||
actorRole int
|
||||
forwardID int64
|
||||
wantCode int
|
||||
wantInFlow int64
|
||||
wantOutFlow int64
|
||||
}{
|
||||
{name: "admin resets another user's rule", actorID: 1, actorRole: 0, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "owner resets own rule", actorID: 2, actorRole: 1, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "user cannot reset another user's rule", actorID: 3, actorRole: 1, forwardID: 20, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
{name: "missing rule is rejected", actorID: 1, actorRole: 0, forwardID: 999, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
h, r := setupForwardResetFlowHandler(t)
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, tt.forwardID, tt.actorID, tt.actorRole)
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
h.forwardResetFlow(res, req)
|
||||
|
||||
if got := decodeForwardResetFlowCode(t, res); got != tt.wantCode {
|
||||
t.Fatalf("code = %d, want %d; body=%s", got, tt.wantCode, res.Body.String())
|
||||
}
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM forward WHERE id = 20", tt.wantInFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM forward WHERE id = 20", tt.wantOutFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM user WHERE id = 2", 700)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM user_tunnel WHERE id = 10", 600)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResetFlowRejectsInvalidRequests(t *testing.T) {
|
||||
h, _ := setupForwardResetFlowHandler(t)
|
||||
|
||||
t.Run("non post", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodGet, 20, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid id", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, 0, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func setupForwardResetFlowHandler(t *testing.T) (*Handler, *repo.Repository) {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "forward-reset-handler.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
statements := []string{
|
||||
`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(1, 'admin', 'pwd', 0, 0, 100, 0, 0, 0, 10, 1000, 1000, 1)`,
|
||||
`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, 'owner', 'pwd', 1, 0, 100, 700, 900, 0, 10, 1000, 1000, 1)`,
|
||||
`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(3, 'other', 'pwd', 1, 0, 100, 0, 0, 0, 10, 1000, 1000, 1)`,
|
||||
`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 'tunnel', 1, 1, 'tls', 1, 1000, 1000, 1, NULL, 0)`,
|
||||
`INSERT INTO user_tunnel(id, user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, 10, 100, 500, 600, 0, 0, 1)`,
|
||||
`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, 'owner', 'target', 1, '127.0.0.1:80', 'fifo', 111, 222, 1000, 1000, 1, 0)`,
|
||||
}
|
||||
for _, statement := range statements {
|
||||
if err := r.DB().Exec(statement).Error; err != nil {
|
||||
t.Fatalf("seed database: %v", err)
|
||||
}
|
||||
}
|
||||
return New(r, "test-secret"), r
|
||||
}
|
||||
|
||||
func newForwardResetFlowRequest(t *testing.T, method string, forwardID, actorID int64, roleID int) *http.Request {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(map[string]int64{"id": forwardID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(method, "/api/v1/forward/reset-flow", bytes.NewReader(body))
|
||||
claims := auth.Claims{Sub: strconv.FormatInt(actorID, 10), RoleID: roleID}
|
||||
return req.WithContext(context.WithValue(req.Context(), middleware.ClaimsContextKey, claims))
|
||||
}
|
||||
|
||||
func decodeForwardResetFlowCode(t *testing.T, res *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
}
|
||||
if err := json.Unmarshal(res.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("decode response: %v; body=%s", err, res.Body.String())
|
||||
}
|
||||
return payload.Code
|
||||
}
|
||||
|
||||
func assertForwardResetFlowDBValue(t *testing.T, r *repo.Repository, query string, want int64) {
|
||||
t.Helper()
|
||||
var got int64
|
||||
if err := r.DB().Raw(query).Scan(&got).Error; err != nil {
|
||||
t.Fatalf("query %q: %v", query, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("query %q returned %d, want %d", query, got, want)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
If the project's default error code differs from `-1`, replace the test expectation with the actual `response.ErrDefault` code after inspecting one existing handler response; do not weaken the success and database assertions.
|
||||
|
||||
- [ ] **Step 2: Run the handler tests and verify the missing handler failure**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./internal/http/handler -run TestForwardResetFlow -count=1
|
||||
```
|
||||
|
||||
Expected: compilation fails because `forwardResetFlow` is undefined.
|
||||
|
||||
- [ ] **Step 3: Register and implement the endpoint**
|
||||
|
||||
Add this route beside the other forward routes in `go-backend/internal/http/handler/handler.go`:
|
||||
|
||||
```go
|
||||
mux.HandleFunc("/api/v1/forward/reset-flow", h.forwardResetFlow)
|
||||
```
|
||||
|
||||
Add this handler beside `forwardPause` and `forwardResume` in `go-backend/internal/http/handler/mutations.go`:
|
||||
|
||||
```go
|
||||
func (h *Handler) forwardResetFlow(w http.ResponseWriter, r *http.Request) {
|
||||
id := idFromBody(r, w)
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
if _, _, _, err := h.resolveForwardAccess(r, id); err != nil {
|
||||
if errors.Is(err, errForwardNotFound) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发不存在"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.ResetForwardFlow(id, time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
```
|
||||
|
||||
This deliberately does not call runtime service controls or nftables reconciliation.
|
||||
|
||||
- [ ] **Step 4: Format and run the focused handler tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && gofmt -w internal/http/handler/forward_reset_flow_test.go internal/http/handler/handler.go internal/http/handler/mutations.go
|
||||
go test ./internal/http/handler -run TestForwardResetFlow -count=1
|
||||
```
|
||||
|
||||
Expected: all reset endpoint tests pass.
|
||||
|
||||
- [ ] **Step 5: Run all backend tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./...
|
||||
```
|
||||
|
||||
Expected: all backend packages and contract tests pass, excluding environment-gated PostgreSQL tests when `FLVX_POSTGRES_TEST_DSN` is unset.
|
||||
|
||||
- [ ] **Step 6: Commit the endpoint change**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/mutations.go go-backend/internal/http/handler/forward_reset_flow_test.go
|
||||
git commit -m "feat: add forward flow reset endpoint"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Add the rule-page reset action and confirmation modal
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/api/index.ts`
|
||||
- Modify: `vite-frontend/src/pages/forward.tsx`
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: `POST /forward/reset-flow`, the page's `Forward` shape, `refreshForwardList`, toast notifications, and existing modal/button bridge components.
|
||||
- Produces: `resetForwardFlow(id: number)`, a shared reset handler, disabled zero-usage actions in all rule views, and one confirmation modal.
|
||||
|
||||
- [ ] **Step 1: Add the frontend API wrapper**
|
||||
|
||||
Add beside the forward control operations in `vite-frontend/src/api/index.ts`:
|
||||
|
||||
```ts
|
||||
export const resetForwardFlow = (forwardId: number) =>
|
||||
Network.post("/forward/reset-flow", { id: forwardId });
|
||||
```
|
||||
|
||||
Import `resetForwardFlow` from `@/api` in `vite-frontend/src/pages/forward.tsx`.
|
||||
|
||||
- [ ] **Step 2: Add page state and shared reset handlers**
|
||||
|
||||
Add state beside the existing delete modal state:
|
||||
|
||||
```ts
|
||||
const [resetFlowModalOpen, setResetFlowModalOpen] = useState(false);
|
||||
const [resetFlowLoading, setResetFlowLoading] = useState(false);
|
||||
const [forwardToResetFlow, setForwardToResetFlow] = useState<Forward | null>(null);
|
||||
```
|
||||
|
||||
Add these handlers beside `handleDelete` and `confirmDelete`:
|
||||
|
||||
```ts
|
||||
const handleResetFlow = (forward: Forward) => {
|
||||
if ((forward.inFlow || 0) + (forward.outFlow || 0) <= 0) return;
|
||||
setForwardToResetFlow(forward);
|
||||
setResetFlowModalOpen(true);
|
||||
};
|
||||
|
||||
const confirmResetFlow = async () => {
|
||||
if (!forwardToResetFlow) return;
|
||||
|
||||
setResetFlowLoading(true);
|
||||
try {
|
||||
const res = await resetForwardFlow(forwardToResetFlow.id);
|
||||
|
||||
if (res.code !== 0) {
|
||||
toast.error(res.msg || "流量清零失败");
|
||||
return;
|
||||
}
|
||||
|
||||
toast.success("规则流量已清零");
|
||||
setResetFlowModalOpen(false);
|
||||
setForwardToResetFlow(null);
|
||||
await refreshForwardList(false);
|
||||
} catch {
|
||||
toast.error("流量清零失败");
|
||||
} finally {
|
||||
setResetFlowLoading(false);
|
||||
}
|
||||
};
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Add one reusable reset icon button to both table row components**
|
||||
|
||||
Pass `handleResetFlow` into `SortableTableRow` and `SortableCompactTableRow` at every render site. Add it to each component's destructured props.
|
||||
|
||||
Insert this button between diagnosis and delete in each table action cell:
|
||||
|
||||
```tsx
|
||||
<Button
|
||||
isIconOnly
|
||||
className="bg-secondary/10 text-secondary hover:bg-secondary/20"
|
||||
isDisabled={(forward.inFlow || 0) + (forward.outFlow || 0) <= 0}
|
||||
size="sm"
|
||||
title="流量清零"
|
||||
onPress={() => handleResetFlow(forward)}
|
||||
>
|
||||
<svg
|
||||
aria-hidden="true"
|
||||
className="h-4 w-4"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
d="M4 4v6h6M20 20v-6h-6M20 9a8 8 0 00-13.657-3.657L4 8m16 8-2.343 2.657A8 8 0 014 15"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
strokeWidth={2}
|
||||
/>
|
||||
</svg>
|
||||
</Button>
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Add the reset action to the card view**
|
||||
|
||||
Insert a fourth action button between diagnosis and delete in `renderForwardCard`:
|
||||
|
||||
```tsx
|
||||
<Button
|
||||
className="flex-1 min-h-8"
|
||||
color="secondary"
|
||||
isDisabled={(forward.inFlow || 0) + (forward.outFlow || 0) <= 0}
|
||||
size="sm"
|
||||
startContent={
|
||||
<svg
|
||||
aria-hidden="true"
|
||||
className="w-3 h-3"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
d="M4 4v6h6M20 20v-6h-6M20 9a8 8 0 00-13.657-3.657L4 8m16 8-2.343 2.657A8 8 0 014 15"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
strokeWidth={2}
|
||||
/>
|
||||
</svg>
|
||||
}
|
||||
variant="flat"
|
||||
onPress={() => handleResetFlow(forward)}
|
||||
>
|
||||
清零
|
||||
</Button>
|
||||
```
|
||||
|
||||
Change the card action container from `flex gap-1.5 mt-3` to `grid grid-cols-2 gap-1.5 mt-3` so all four actions remain readable at the smallest supported card width.
|
||||
|
||||
- [ ] **Step 5: Add the confirmation modal**
|
||||
|
||||
Add beside the delete confirmation modal:
|
||||
|
||||
```tsx
|
||||
<Modal
|
||||
backdrop="blur"
|
||||
classNames={{
|
||||
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
|
||||
}}
|
||||
isOpen={resetFlowModalOpen}
|
||||
placement="center"
|
||||
scrollBehavior="inside"
|
||||
size="lg"
|
||||
onOpenChange={setResetFlowModalOpen}
|
||||
>
|
||||
<ModalContent>
|
||||
{(onClose) => (
|
||||
<>
|
||||
<ModalHeader className="flex flex-col gap-1">
|
||||
<h2 className="text-lg font-bold text-secondary">确认流量清零</h2>
|
||||
</ModalHeader>
|
||||
<ModalBody>
|
||||
<p className="text-default-600">
|
||||
确定要清零规则{" "}
|
||||
<span className="font-semibold text-foreground">
|
||||
"{forwardToResetFlow?.name}"
|
||||
</span>{" "}
|
||||
当前显示的上传和下载流量吗?
|
||||
</p>
|
||||
<p className="text-small text-default-500 mt-2">
|
||||
此操作不可撤销,但不会影响用户总流量、用户隧道配额和历史统计。
|
||||
</p>
|
||||
</ModalBody>
|
||||
<ModalFooter>
|
||||
<Button isDisabled={resetFlowLoading} variant="light" onPress={onClose}>
|
||||
取消
|
||||
</Button>
|
||||
<Button
|
||||
color="secondary"
|
||||
isLoading={resetFlowLoading}
|
||||
onPress={confirmResetFlow}
|
||||
>
|
||||
确认清零
|
||||
</Button>
|
||||
</ModalFooter>
|
||||
</>
|
||||
)}
|
||||
</ModalContent>
|
||||
</Modal>
|
||||
```
|
||||
|
||||
Add this wrapper beside the other reset handlers and pass it to the modal as `onOpenChange={handleResetFlowModalOpenChange}`:
|
||||
|
||||
```ts
|
||||
const handleResetFlowModalOpenChange = (isOpen: boolean) => {
|
||||
if (resetFlowLoading) return;
|
||||
setResetFlowModalOpen(isOpen);
|
||||
if (!isOpen) {
|
||||
setForwardToResetFlow(null);
|
||||
}
|
||||
};
|
||||
```
|
||||
|
||||
- [ ] **Step 6: Format and verify the frontend**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd vite-frontend && pnpm exec prettier --write src/api/index.ts src/pages/forward.tsx
|
||||
pnpm run build
|
||||
pnpm run lint
|
||||
```
|
||||
|
||||
Expected: TypeScript/Vite build succeeds and ESLint finishes without errors.
|
||||
|
||||
- [ ] **Step 7: Commit the frontend change**
|
||||
|
||||
```bash
|
||||
git add vite-frontend/src/api/index.ts vite-frontend/src/pages/forward.tsx
|
||||
git commit -m "feat: add forward flow reset action"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Perform integrated verification
|
||||
|
||||
**Files:**
|
||||
- Verify only; no planned source changes.
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: the repository method, API endpoint, and rule-page action from Tasks 1-3.
|
||||
- Produces: evidence that the complete feature builds and all affected tests pass.
|
||||
|
||||
- [ ] **Step 1: Run the complete backend suite**
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./...
|
||||
```
|
||||
|
||||
Expected: all available backend tests pass.
|
||||
|
||||
- [ ] **Step 2: Run the complete frontend checks**
|
||||
|
||||
```bash
|
||||
cd vite-frontend && pnpm run build && pnpm run lint
|
||||
```
|
||||
|
||||
Expected: both commands exit successfully.
|
||||
|
||||
- [ ] **Step 3: Check formatting and working-tree scope**
|
||||
|
||||
```bash
|
||||
git diff --check
|
||||
git status --short
|
||||
git log -4 --oneline
|
||||
```
|
||||
|
||||
Expected: no whitespace errors; the working tree is clean; the three feature commits are visible after the design and implementation-plan commits.
|
||||
|
||||
- [ ] **Step 4: Manually verify the feature when a local panel is available**
|
||||
|
||||
1. Open the Rules page as an administrator and reset a rule with non-zero upload/download traffic.
|
||||
2. Confirm the modal states that user totals, tunnel quota, and history are unaffected.
|
||||
3. Confirm the rule immediately shows zero after success.
|
||||
4. Confirm the user page's total traffic and user-tunnel traffic values did not change.
|
||||
5. Generate new traffic and confirm the rule starts accumulating from zero.
|
||||
6. Log in as a normal user and confirm the user can reset an owned rule but cannot access another user's rule through a direct API request.
|
||||
|
||||
Expected: all six checks match the design specification.
|
||||
@@ -0,0 +1,133 @@
|
||||
# 允许转发到本地地址开关设计
|
||||
|
||||
**日期**: 2026-04-26
|
||||
**状态**: 待审核
|
||||
**作者**: AI Assistant
|
||||
|
||||
## 概述
|
||||
|
||||
新增一个全局设置开关,控制规则目标地址是否允许指向本地/内网地址。默认关闭,保持当前安全策略不变;开启后,规则创建和编辑时允许将目标地址设置为 `127.0.0.1`、`10.x.x.x`、`172.16-31.x.x`、`192.168.x.x` 等本地或私网地址。
|
||||
|
||||
## 背景
|
||||
|
||||
当前后端在规则创建和编辑时会调用 `IsSafeRemoteAddr()`,统一禁止目标地址指向本地/内网地址,用来降低 SSRF / 开放代理风险。这一行为是全局硬编码的,无法按部署场景调整。
|
||||
|
||||
有些用户需要把规则转发到本机或内网服务,因此需要一个显式、全局的开关来放宽这条限制。
|
||||
|
||||
## 目标
|
||||
|
||||
1. 在设置页提供一个全局开关控制该行为。
|
||||
2. 默认关闭,不改变现有安全默认值。
|
||||
3. 开启后,规则创建和编辑允许本地/内网目标地址。
|
||||
4. 不影响其他安全校验和其他业务流程。
|
||||
|
||||
## 影响范围
|
||||
|
||||
### 后端
|
||||
- `go-backend/internal/http/handler/security_utils.go`
|
||||
- `go-backend/internal/http/handler/mutations.go`
|
||||
- `go-backend/internal/http/handler/handler.go`
|
||||
|
||||
### 前端
|
||||
- `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
### 测试
|
||||
- `go-backend/tests/contract/forward_contract_test.go` 或新增独立 contract test
|
||||
|
||||
## 详细设计
|
||||
|
||||
### 1. 配置存储
|
||||
|
||||
使用现有 `vite_config` 表新增一个配置项:
|
||||
|
||||
| name | value | 说明 |
|
||||
|------|-------|------|
|
||||
| `allow_local_remote_addr` | `"1"` / `"0"` | 是否允许规则目标地址指向本地/内网地址 |
|
||||
|
||||
约定:
|
||||
- 未配置时按 `"0"` 处理
|
||||
- `"1"` 表示允许
|
||||
- 其他值一律按关闭处理
|
||||
|
||||
### 2. 后端行为
|
||||
|
||||
新增一个轻量辅助函数,用于读取该配置开关:
|
||||
|
||||
```go
|
||||
func (h *Handler) allowLocalRemoteAddr() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return false
|
||||
}
|
||||
cfg, err := h.repo.GetConfigByName("allow_local_remote_addr")
|
||||
if err != nil || cfg == nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(cfg.Value) == "1"
|
||||
}
|
||||
```
|
||||
|
||||
在以下路径中应用:
|
||||
- `forwardCreate`
|
||||
- `forwardUpdate`
|
||||
|
||||
行为改为:
|
||||
- 当开关关闭时,继续执行 `IsSafeRemoteAddr(remoteAddr)`
|
||||
- 当开关开启时,跳过这条“本地/内网地址禁止”校验
|
||||
|
||||
这样可以把改动范围限定在规则创建/编辑,不改变其他依赖 `IsSafeRemoteAddr()` 的场景。
|
||||
|
||||
### 3. 前端设置页
|
||||
|
||||
在 `vite-frontend/src/pages/config.tsx` 增加一个全局开关配置项。
|
||||
|
||||
建议文案:
|
||||
|
||||
- 标签:`允许转发到本地地址`
|
||||
- 描述:`开启后,规则目标地址可指向 127.0.0.1、10.x.x.x、172.16-31.x.x、192.168.x.x 等本地或内网地址。默认关闭以降低开放代理风险。`
|
||||
|
||||
控件类型:
|
||||
- 使用现有设置页的布尔开关模式
|
||||
|
||||
默认显示策略:
|
||||
- 不依赖其他配置项
|
||||
- 直接显示在设置页的网络/安全相关区域;若现有页面没有单独分区,则先按现有配置项组织方式加入即可
|
||||
|
||||
### 4. 错误与兼容性
|
||||
|
||||
关闭开关时:
|
||||
- 保持现有错误行为,继续阻止本地/内网地址
|
||||
|
||||
开启开关时:
|
||||
- 仅放开“本地/内网地址禁止”这条限制
|
||||
- 仍保留地址格式解析失败等其他错误
|
||||
|
||||
### 5. 测试
|
||||
|
||||
需要补两类后端契约测试:
|
||||
|
||||
1. 开关关闭时拒绝本地/内网地址
|
||||
- 创建规则时使用本地/内网地址
|
||||
- 断言接口返回非 0 code
|
||||
|
||||
2. 开关开启时允许本地/内网地址
|
||||
- 先写入 `vite_config(name=allow_local_remote_addr, value=1)`
|
||||
- 创建或更新规则时使用相同地址
|
||||
- 断言接口成功
|
||||
|
||||
建议至少覆盖:
|
||||
- create 路径
|
||||
- update 路径
|
||||
- 多目标地址输入(逗号或换行分隔)中包含本地地址时的行为
|
||||
|
||||
## 风险与约束
|
||||
|
||||
1. 该开关会降低默认安全防护,应明确标注风险。
|
||||
2. 这是全局开关,不做用户级或规则级细分控制。
|
||||
3. 该开关只影响规则目标地址校验,不影响其他独立的安全策略。
|
||||
|
||||
## 推荐实施顺序
|
||||
|
||||
1. 先补失败的后端契约测试
|
||||
2. 实现后端配置读取与创建/更新分支控制
|
||||
3. 在设置页增加开关
|
||||
4. 跑后端测试与前端构建验证
|
||||
@@ -0,0 +1,321 @@
|
||||
# 规则每 IP 连接数与限速设计
|
||||
|
||||
**日期**: 2026-04-27
|
||||
**状态**: 待审核
|
||||
**作者**: AI Assistant
|
||||
|
||||
## 概述
|
||||
|
||||
在转发规则的高级设置中新增两类每客户端 IP 限制:每 IP 最大连接数、每 IP 带宽限速。保留现有总量限制语义不变,新增字段只在用户显式配置时生效。
|
||||
|
||||
实现优先复用 GOST 已有能力:`climiters` 的 `$$ N` 表示每个客户端 IP 独立最大连接数;`limiters` 支持 IP/CIDR 级带宽桶,可用 `0.0.0.0/0` 和 `::/0` 实现默认覆盖所有 IPv4/IPv6 客户端的每 IP 带宽限速。
|
||||
|
||||
## 背景
|
||||
|
||||
当前 FLVX 已经支持规则级最大连接数和规则级限速,但这两个限制都是规则总量:
|
||||
|
||||
- `maxConn` 下发为 GOST `climiters` 的 `$ N`,限制整条规则的总并发连接数。
|
||||
- `speedId` 下发为 GOST `limiters` 的 `$ in out`,限制整条规则的总带宽。
|
||||
|
||||
用户需要的是按客户端 IP 隔离的限制,例如每个 IP 最多 5 个连接、每个 IP 最多 10 Mbps,而不是所有客户端共享同一个总量。
|
||||
|
||||
## GOST 能力确认
|
||||
|
||||
### 连接数限制
|
||||
|
||||
`go-gost/x/limiter/conn/conn.go` 已内置以下语义:
|
||||
|
||||
| Key | 含义 |
|
||||
|-----|------|
|
||||
| `$` | 全局连接数限制,所有客户端共享一个 limiter |
|
||||
| `$$` | 每个客户端 IP 独立连接数限制,每个 IP 创建自己的 limiter |
|
||||
| `IP` / `CIDR` | 指定 IP 或 CIDR 的连接数限制 |
|
||||
|
||||
因此每 IP 连接数无需新增 agent 限制器,只需后端下发 `$$ N`。
|
||||
|
||||
### 带宽限制
|
||||
|
||||
`go-gost/x/limiter/traffic/traffic.go` 已内置以下语义:
|
||||
|
||||
| Key | 含义 |
|
||||
|-----|------|
|
||||
| `$` | 服务级总带宽限制 |
|
||||
| `$$` | 连接级带宽限制 |
|
||||
| `IP` / `CIDR` | 客户端 IP 或 CIDR 级带宽限制 |
|
||||
|
||||
CIDR 级限制使用 generator,为命中的客户端 IP 创建独立 limiter。使用 `0.0.0.0/0` 和 `::/0` 可以覆盖所有 IPv4/IPv6 客户端,实现每 IP 带宽限速。
|
||||
|
||||
### 现有缺口
|
||||
|
||||
TCP listener 已在 Accept 后用客户端地址包装连接级 traffic limiter,路径可用于每 IP 带宽。UDP listener 当前只在 PacketConn 上应用服务级 limiter,没有在 `Accept()` 后按客户端 UDP pseudo-connection 包装 limiter,也没有挂接 connection limiter。因此要让 UDP 与 TCP 语义一致,需要补齐 UDP listener 的 per-client wrapper。
|
||||
|
||||
## 目标
|
||||
|
||||
1. 保留现有 `maxConn` 和 `speedId` 的总量语义。
|
||||
2. 在规则上新增每 IP 最大连接数。
|
||||
3. 在规则上新增每 IP 带宽限速。
|
||||
4. 同一规则允许同时配置总量限制和每 IP 限制。
|
||||
5. 普通用户不能设置或修改限速规则字段,保持现有权限模型。
|
||||
6. TCP 和 UDP 入口都尽量遵循相同限制语义。
|
||||
|
||||
## 非目标
|
||||
|
||||
1. 不新增按用户组、节点组、国家地区、ASN 的限制。
|
||||
2. 不新增请求频率限制;本次“每个 IP 限速”指带宽限速,不是新建连接频率。
|
||||
3. 不改变已有 speed limit 规则表的单位和含义。
|
||||
4. 不把用户级默认最大连接数改成每 IP 语义;用户级 `maxConn` 继续作为默认总连接数。
|
||||
|
||||
## 数据模型
|
||||
|
||||
在 `forward` 表新增两个字段:
|
||||
|
||||
| 字段 | 类型 | 默认 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `ip_max_conn` | int | `0` | 每 IP 最大连接数,`0` 表示不启用 |
|
||||
| `ip_speed_id` | nullable int64 | `NULL` | 每 IP 带宽限速规则 ID,`NULL` 表示不启用 |
|
||||
|
||||
Go 模型新增:
|
||||
|
||||
```go
|
||||
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
|
||||
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
|
||||
```
|
||||
|
||||
字段会通过现有 auto-migrate 机制创建,保持 SQLite/PostgreSQL 兼容,不使用 SQLite 不兼容的 GORM tags。
|
||||
|
||||
## API 行为
|
||||
|
||||
### 创建规则
|
||||
|
||||
`/forward/create` 新增入参:
|
||||
|
||||
```json
|
||||
{
|
||||
"ipMaxConn": 5,
|
||||
"ipSpeedId": 123
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- `ipMaxConn` 缺省或小于等于 `0` 时按 `0` 存储,不启用每 IP 连接数限制。
|
||||
- `ipSpeedId` 缺省或不存在时存为 `NULL`,不启用每 IP 带宽限速。
|
||||
- `ipSpeedId` 指向不存在的限速规则时按 `NULL` 处理,沿用现有 `speedId` 的容错策略。
|
||||
- 普通用户提交非空 `ipSpeedId` 时返回错误,保持与 `speedId` 一致的权限边界。
|
||||
|
||||
### 更新规则
|
||||
|
||||
`/forward/update` 新增入参:
|
||||
|
||||
```json
|
||||
{
|
||||
"ipMaxConn": 5,
|
||||
"ipSpeedId": 123
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- 未提交 `ipMaxConn` 时保留原值;提交空值或 `0` 时清除每 IP 连接数限制。
|
||||
- 未提交 `ipSpeedId` 时保留原值;提交 `null` 时清除每 IP 带宽限速。
|
||||
- 普通用户不能把 `ipSpeedId` 改成不同的非空值。
|
||||
- 更新后重新同步运行时服务和 limiter。
|
||||
|
||||
### 列表返回
|
||||
|
||||
`/forward/list` 返回项新增:
|
||||
|
||||
```json
|
||||
{
|
||||
"ipMaxConn": 5,
|
||||
"ipSpeedId": 123,
|
||||
"ipSpeedLimitName": "每IP 10Mbps"
|
||||
}
|
||||
```
|
||||
|
||||
`ipSpeedLimitName` 可选,但建议返回,便于前端显示缺失或已删除的限速规则。
|
||||
|
||||
## 后端运行时同步
|
||||
|
||||
### 连接数限制器
|
||||
|
||||
将现有连接限制器构建从单一总量扩展为组合规则。
|
||||
|
||||
当前行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "rule_conn_limit_42",
|
||||
"limits": ["$ 100"]
|
||||
}
|
||||
```
|
||||
|
||||
新增行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "rule_conn_limit_42",
|
||||
"limits": ["$ 100", "$$ 5"]
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- `maxConn > 0` 时追加 `$ maxConn`。
|
||||
- `ipMaxConn > 0` 时追加 `$$ ipMaxConn`。
|
||||
- 如果规则未配置 `maxConn` 且用户有 `MaxConn > 0`,继续继承用户级总连接数,追加 `$ user.MaxConn`。
|
||||
- 如果两者都没有,则不下发 `climiter`,服务不引用 `climiter`。
|
||||
- limiter 名称继续优先使用 `rule_conn_limit_<forwardID>`;只有用户级默认总连接数且规则没有任何连接限制时可继续使用 `user_conn_limit_<userID>`,避免不必要的 per-rule limiter。
|
||||
|
||||
### 带宽限制器
|
||||
|
||||
将现有规则限速从单一 `speedId` 扩展为组合 limiter。
|
||||
|
||||
当前行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "123",
|
||||
"limits": ["$ 1.3MB 1.3MB"]
|
||||
}
|
||||
```
|
||||
|
||||
新增每 IP 行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "rule_traffic_limit_42",
|
||||
"limits": [
|
||||
"$ 1.3MB 1.3MB",
|
||||
"0.0.0.0/0 1.3MB 1.3MB",
|
||||
"::/0 1.3MB 1.3MB"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- 只有总量 `speedId` 时,保持现有名称和下发路径,服务继续引用 `speedId` 字符串。
|
||||
- 只有每 IP `ipSpeedId` 时,创建 `rule_traffic_limit_<forwardID>`,只包含 IPv4/IPv6 CIDR 行。
|
||||
- 总量和每 IP 同时存在时,创建 `rule_traffic_limit_<forwardID>`,同时包含 `$` 和 CIDR 行。
|
||||
- 如果规则没有 `speedId`,则总量仍可继承 user tunnel 的 `speedId`,保持现有 fallback 语义;当继承的总量限速与 `ipSpeedId` 同时存在时,也使用 `rule_traffic_limit_<forwardID>` 组合 limiter。
|
||||
- 每 IP 限速不从 user tunnel 继承,只由规则字段控制。
|
||||
- `AddLimiters` 失败且提示已存在时,使用 `UpdateLimiters` 更新。
|
||||
|
||||
### 服务配置
|
||||
|
||||
`buildForwardServiceConfigs` 需要从当前 `limiterID *int64` / `cLimiterName string` 扩展为更明确的运行时限制描述,例如:
|
||||
|
||||
```go
|
||||
type forwardRuntimeLimiters struct {
|
||||
TrafficLimiter string
|
||||
ConnLimiter string
|
||||
}
|
||||
```
|
||||
|
||||
服务配置只关心最终引用的 limiter 名称:
|
||||
|
||||
- `service["limiter"] = runtimeLimiters.TrafficLimiter`
|
||||
- `service["climiter"] = runtimeLimiters.ConnLimiter`
|
||||
|
||||
这样可以把“如何构建 limiter payload”的逻辑和“如何构建 service JSON”的逻辑分开。
|
||||
|
||||
## Agent/GOST 调整
|
||||
|
||||
### WebSocket 命令
|
||||
|
||||
当前 agent WebSocket 已支持:
|
||||
|
||||
- `AddLimiters` / `UpdateLimiters` / `DeleteLimiters`
|
||||
- `AddCLimiters` / `UpdateCLimiters` / `DeleteCLimiters`
|
||||
|
||||
本设计无需新增命令类型。
|
||||
|
||||
### UDP listener
|
||||
|
||||
补齐 `go-gost/x/listener/udp/listener.go` 的 `Accept()` 包装逻辑,使 UDP pseudo-connection 与 TCP listener 一致:
|
||||
|
||||
- 对 `l.options.ConnLimiter` 按客户端地址应用连接数限制。
|
||||
- 对 `l.options.TrafficLimiter` 按 `conn.RemoteAddr().String()` 应用连接级 traffic wrapper。
|
||||
|
||||
需要注意 UDP pseudo-connection 的生命周期由内部 UDP listener 的 TTL/keepalive 控制;connection limiter 必须在 pseudo-connection 关闭时释放计数。
|
||||
|
||||
## 前端设计
|
||||
|
||||
在 `vite-frontend/src/pages/forward.tsx` 的规则高级设置中新增两个控件:
|
||||
|
||||
1. `每 IP 最大连接数`
|
||||
- 类型:number input。
|
||||
- 文案:`每个客户端 IP 可同时建立的最大连接数;0 或空表示不限制。`
|
||||
- 字段:`ipMaxConn`。
|
||||
|
||||
2. `每 IP 限速`
|
||||
- 类型:Select,复用现有限速规则列表。
|
||||
- 文案:`每个客户端 IP 独享该带宽限制;不选择表示不限制。`
|
||||
- 字段:`ipSpeedId`。
|
||||
- 只对管理员显示,保持与 `规则限速` 一致。
|
||||
|
||||
前端类型需要同步更新:
|
||||
|
||||
- `ForwardApiItem`
|
||||
- `ForwardMutationPayload`
|
||||
- `ForwardForm` 或页面内等价类型
|
||||
|
||||
## 错误处理与兼容性
|
||||
|
||||
1. 旧数据默认 `ip_max_conn=0`、`ip_speed_id=NULL`,行为与当前版本一致。
|
||||
2. 现有 agent 已支持 limiter 命令和 GOST limiter 语法;发布时需要包含 UDP 修复,才能让 TCP/UDP 都获得完整语义。
|
||||
3. 节点离线时沿用现有 warning 行为,规则仍可保存,在线节点跳过下发。
|
||||
4. 如果每 IP speed limit ID 被删除,更新时按 `NULL` 处理,列表页可提示或自动清除,和现有 `speedId` 行为一致。
|
||||
5. 如果 IPv6 CIDR 在某些监听路径未命中,IPv4 行仍正常生效;测试应覆盖 IPv4,IPv6 通过 payload 合同保证下发。
|
||||
|
||||
## 测试计划
|
||||
|
||||
### 后端 contract 测试
|
||||
|
||||
新增或扩展 `go-backend/tests/contract/max_conn_limit_contract_test.go`:
|
||||
|
||||
1. 创建规则时设置 `ipMaxConn=5`,断言 `AddCLimiters` payload 包含 `$$ 5`。
|
||||
2. 同时设置 `maxConn=100` 和 `ipMaxConn=5`,断言 payload 包含 `$ 100` 和 `$$ 5`。
|
||||
3. 用户级 `MaxConn` 存在且规则 `ipMaxConn=5` 时,断言 payload 包含 `$ userMaxConn` 和 `$$ 5`。
|
||||
|
||||
新增每 IP 限速 contract 测试:
|
||||
|
||||
1. 创建规则时设置 `ipSpeedId`,断言 `AddLimiters` payload 包含 `0.0.0.0/0 ...` 和 `::/0 ...`。
|
||||
2. 同时设置 `speedId` 和 `ipSpeedId`,断言组合 limiter 包含 `$ ...` 与两个 CIDR 行,服务引用 `rule_traffic_limit_<forwardID>`。
|
||||
3. 普通用户提交 `ipSpeedId` 返回错误。
|
||||
|
||||
### Repository/API 测试
|
||||
|
||||
1. `CreateForwardTx`、`UpdateForward`、列表查询读写 `ip_max_conn` 和 `ip_speed_id`。
|
||||
2. `/forward/list` 返回 `ipMaxConn`、`ipSpeedId`。
|
||||
|
||||
### GOST/x 测试
|
||||
|
||||
1. `go-gost/x/limiter/conn`:验证 `$$ N` 为不同 IP 创建独立 limiter。
|
||||
2. `go-gost/x/limiter/traffic`:验证 `0.0.0.0/0` 为不同 IPv4 创建独立 limiter。
|
||||
3. UDP listener:验证 Accept 返回的 UDP pseudo-connection 关闭后释放 connection limiter。
|
||||
|
||||
### 验证命令
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd go-gost/x && go test ./limiter/... ./listener/udp/...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
## 推荐实施顺序
|
||||
|
||||
1. 后端模型、repo DTO、API 字段读写。
|
||||
2. 后端 limiter payload 构建与服务引用重构。
|
||||
3. Contract 测试覆盖连接数和带宽 payload。
|
||||
4. GOST UDP listener per-client wrapper 与相关测试。
|
||||
5. 前端高级设置表单和类型更新。
|
||||
6. 运行后端测试、GOST/x 相关测试、前端构建。
|
||||
|
||||
## 风险
|
||||
|
||||
1. UDP pseudo-connection 生命周期和 TCP 连接不同,连接数释放必须依赖 Close 包装正确执行。
|
||||
2. 总带宽和每 IP 带宽组合时 limiter 名称从纯 speed ID 变为 rule-level 名称,需要确保更新已有规则时不会留下错误引用。
|
||||
3. 旧节点如果没有 UDP wrapper 修复,TCP 生效但 UDP 每 IP 语义可能不完整;发布时应要求 agent 同步升级。
|
||||
4. 每 IP 带宽是每个入口节点本地独立限制,不是跨节点全局聚合限制。
|
||||
@@ -0,0 +1,74 @@
|
||||
# Monitoring Retention And Storage Display Design
|
||||
|
||||
## Goal
|
||||
|
||||
Add an administrator-facing configuration for monitoring data retention and display the current database storage usage in the configuration page.
|
||||
|
||||
## Scope
|
||||
|
||||
- Add a single config key: `monitor_retention_days`.
|
||||
- Default retention is `7` days.
|
||||
- Apply the retention window uniformly to:
|
||||
- `node_metric`
|
||||
- `tunnel_metric`
|
||||
- `service_monitor_result`
|
||||
- `tunnel_quality`
|
||||
- Show database usage on the config page as a read-only operational value.
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- No per-table retention settings.
|
||||
- No manual purge button.
|
||||
- No database vacuum/compaction action.
|
||||
- No frontend test framework changes.
|
||||
|
||||
## Backend Design
|
||||
|
||||
### Retention Config
|
||||
|
||||
- Store `monitor_retention_days` in `vite_config`, consistent with existing site settings.
|
||||
- Accept integer values from `1` through `3650`.
|
||||
- Missing or invalid stored values fall back to `7` days.
|
||||
- `normalizeAndValidateConfigValue` rejects invalid user-submitted values so bad config does not get saved through the API.
|
||||
|
||||
### Cleanup Flow
|
||||
|
||||
- `metrics.IngestionService.pruneMetrics()` reads `monitor_retention_days` from the repository each hourly cleanup cycle.
|
||||
- The computed cutoff is used for `node_metric`, `tunnel_metric`, and `service_monitor_result`.
|
||||
- `tunnel_quality` uses the same retention config.
|
||||
- `tunnel_quality` cleanup must run even when real-time tunnel quality probing is disabled; disabling probing should stop new probe writes, not stop cleanup.
|
||||
|
||||
### Database Storage API
|
||||
|
||||
- Add an admin-only API endpoint for storage summary, for example `/api/v1/system/storage`.
|
||||
- Response fields:
|
||||
- `dbType`: `sqlite` or `postgres`
|
||||
- `databaseSizeBytes`: raw byte count
|
||||
- `databaseSizeText`: human-readable formatted size
|
||||
- SQLite implementation reports the DB file size and includes `-wal` and `-shm` sidecar files when present.
|
||||
- PostgreSQL implementation uses `pg_database_size(current_database())`.
|
||||
- If size cannot be determined, return an API error rather than a misleading zero.
|
||||
|
||||
## Frontend Design
|
||||
|
||||
- Add `monitor_retention_days` to the config page.
|
||||
- Label: `监控数据保留天数`.
|
||||
- Description: `统一清理节点指标、隧道流量、服务监控结果和隧道质量历史;默认 7 天。`
|
||||
- Use a regular numeric input through the existing config rendering path.
|
||||
- Fetch database storage summary when the config page loads.
|
||||
- Display a read-only card/row named `数据库占用` with `databaseSizeText`.
|
||||
- If fetching fails, show `获取失败` and keep config editing usable.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- Invalid retention values return a validation error on save.
|
||||
- Cleanup logs individual prune failures and continues with other tables, matching existing monitoring cleanup behavior.
|
||||
- Storage summary failures are non-blocking in the frontend.
|
||||
|
||||
## Testing
|
||||
|
||||
- Backend unit tests for retention config parsing and validation.
|
||||
- Backend tests proving custom retention is used by monitoring cleanup.
|
||||
- Backend API/repository test for SQLite storage size returning a non-negative byte count and formatted text.
|
||||
- Run `go test ./...` in `go-backend`.
|
||||
- Run `pnpm run build` in `vite-frontend`.
|
||||
@@ -0,0 +1,156 @@
|
||||
# Best Exit Current Selection Display Design
|
||||
|
||||
## Goal
|
||||
|
||||
When a tunnel uses the `best` multi-exit strategy, show the currently applied best exit in the tunnel list information. Users should be able to see which exit is currently selected without opening logs or diagnosing the tunnel manually.
|
||||
|
||||
The display is informational only. It must not change routing, scoring, switching behavior, or the saved tunnel configuration.
|
||||
|
||||
## Current Context
|
||||
|
||||
- `3.0.0-beta6` adds `best` as a multi-exit strategy.
|
||||
- Runtime selection is stored in the backend `bestExitManager` in memory, keyed by `TunnelID + OwnerNodeID`.
|
||||
- Direct multi-entry tunnels make one independent best-exit decision per entry node.
|
||||
- Tunnels with intermediate chain hops make one independent best-exit decision per final-hop chain node before the exits.
|
||||
- `tunnelList` and `tunnelGet` currently return `repo.ListTunnels()` output directly, so frontend tunnel data only includes configured exits from the database, not the currently applied runtime choice.
|
||||
- The frontend tunnel page maps API items in `vite-frontend/src/pages/tunnel.tsx` and renders list information from that data.
|
||||
|
||||
## User Decisions
|
||||
|
||||
- Show the current best-exit choice in the tunnel list information.
|
||||
- Use a summary plus detail model for multiple owners.
|
||||
- Follow the existing tunnel list refresh cadence; do not add polling or a realtime stream in this phase.
|
||||
- Work text-only; no visual companion is needed.
|
||||
|
||||
## Approach
|
||||
|
||||
Extend the existing tunnel list/detail response with a lightweight runtime state object for `best` tunnels, then render that state beside the tunnel's exit/strategy information in the existing frontend list UI.
|
||||
|
||||
This keeps the display close to the data users already inspect and avoids a separate API or extra frontend request.
|
||||
|
||||
## Backend Design
|
||||
|
||||
### Response Shape
|
||||
|
||||
Add a `bestExitState` object to each tunnel item returned by `tunnelList` and `tunnelGet` when the tunnel has a multi-exit group whose strategy is `best`.
|
||||
|
||||
Response shape:
|
||||
|
||||
```json
|
||||
{
|
||||
"enabled": true,
|
||||
"summary": "香港节点",
|
||||
"status": "applied",
|
||||
"updatedAt": 1777584000000,
|
||||
"reason": "current exit remains best",
|
||||
"items": [
|
||||
{
|
||||
"ownerNodeId": 10,
|
||||
"ownerNodeName": "入口 A",
|
||||
"ownerRole": "entry",
|
||||
"exitNodeId": 30,
|
||||
"exitNodeName": "香港节点",
|
||||
"updatedAt": 1777584000000,
|
||||
"reason": "current exit remains best"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
If the tunnel is not using `best`, omit `bestExitState` or set it to `null`.
|
||||
|
||||
### Owner Semantics
|
||||
|
||||
The display must match the routing model:
|
||||
|
||||
- If there are no middle chain hops, each entry node is an owner.
|
||||
- If there are middle chain hops, each node in the final middle-hop group is an owner.
|
||||
|
||||
Each owner can have a different current best exit. The UI must not imply that a multi-owner tunnel has one global best exit when the owners differ.
|
||||
|
||||
### Summary Rules
|
||||
|
||||
- If all owners currently apply the same exit, `summary` is that exit node name.
|
||||
- If owners apply different exits, `summary` is `多个出口`.
|
||||
- If no applied decision exists yet, `summary` is `等待探测`.
|
||||
- If the tunnel has only one exit, `bestExitState` is not needed because there is no dynamic choice.
|
||||
|
||||
### State Source
|
||||
|
||||
Use the in-memory `bestExitManager` as the source of currently applied decisions.
|
||||
|
||||
Add a read-only snapshot method that returns defensive copies of decision state without exposing mutable internal slices. The handler should convert node IDs to display names from the existing tunnel response data first, then fall back to `h.getNodeRecord` only when the current response does not contain the node.
|
||||
|
||||
The feature should not persist current choices to the database in this phase. A panel restart may reset the displayed runtime state to `等待探测` until the prober initializes it again from the current saved first exit.
|
||||
|
||||
## Frontend Design
|
||||
|
||||
Extend the tunnel item type with optional `bestExitState`.
|
||||
|
||||
In the tunnel list, only render the current best-exit display when:
|
||||
|
||||
- `bestExitState.enabled === true`, or
|
||||
- the tunnel has an exit group with `strategy === "best"` and the backend returns a waiting state.
|
||||
|
||||
Display format:
|
||||
|
||||
- Single applied exit: `最优出口:香港节点`
|
||||
- Multiple applied exits: `最优出口:多个出口`
|
||||
- Waiting: `最优出口:等待探测`
|
||||
|
||||
For multiple owners, render the summary as compact secondary text in the topology/list information cell and set its native `title` attribute to newline-separated detail rows. This avoids adding a new UI dependency or a custom popover. Detail rows should use:
|
||||
|
||||
```text
|
||||
入口 A -> 香港节点
|
||||
入口 B -> 日本节点
|
||||
```
|
||||
|
||||
For tunnels with middle chain hops, label owners as chain nodes when useful:
|
||||
|
||||
```text
|
||||
中转 M1 -> 香港节点
|
||||
中转 M2 -> 日本节点
|
||||
```
|
||||
|
||||
Do not add a new periodic refresh. The display updates when the existing tunnel list is refreshed.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- If the manager has no decision for an owner, show that owner as `等待探测`.
|
||||
- If an exit node ID no longer exists in the current tunnel response, show `未知出口` for that item and keep the list usable.
|
||||
- If an owner node ID no longer exists, show `未知入口` or `未知中转` based on the owner role.
|
||||
- If the backend cannot compute state for one tunnel, omit `bestExitState` for that tunnel and log the error; do not fail the whole tunnel list response.
|
||||
|
||||
## Testing
|
||||
|
||||
Backend tests:
|
||||
|
||||
- `bestExitManager` snapshot returns applied exit IDs without exposing mutable manager state.
|
||||
- Direct multi-entry `best` tunnel produces one display item per entry owner.
|
||||
- Middle-hop tunnel produces one display item per final-hop owner.
|
||||
- Summary is the single exit name when all owners choose the same exit.
|
||||
- Summary is `多个出口` when owners choose different exits.
|
||||
- Summary is `等待探测` when no applied decision exists.
|
||||
- Non-`best` tunnels do not receive `bestExitState`.
|
||||
|
||||
Frontend verification:
|
||||
|
||||
- Tunnel list renders `最优出口:<name>` for a single applied exit.
|
||||
- Tunnel list renders `最优出口:多个出口` plus owner details for multiple applied exits.
|
||||
- Tunnel list renders `最优出口:等待探测` for waiting state.
|
||||
- `pnpm run build` passes.
|
||||
|
||||
Verification commands:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- Do not add a new realtime stream or polling loop.
|
||||
- Do not add a detailed best-exit scoring dashboard.
|
||||
- Do not persist current best-exit choices to the database.
|
||||
- Do not change switching thresholds, probing targets, or runtime chain update behavior.
|
||||
- Do not change existing non-`best` tunnel display behavior.
|
||||
@@ -0,0 +1,186 @@
|
||||
# Best Exit Selection Design
|
||||
|
||||
## Goal
|
||||
|
||||
Add a multi-exit tunnel strategy named `best` that always sends new connections through the currently best-quality exit. The feature should prevent traffic from continuing to use an exit whose latency or packet loss has degraded while the exit is still technically online.
|
||||
|
||||
Existing connections must not be interrupted. Switching affects only new connections created after the runtime chain update is applied.
|
||||
|
||||
## Current Context
|
||||
|
||||
- Tunnel forwarding stores entry, chain, and exit nodes in `chain_tunnel`.
|
||||
- Multi-exit runtime chains are currently rendered as one GOST hop with multiple nodes.
|
||||
- GOST selectors support `fifo`, `round`, `rand`, and `hash`, plus fail filtering through `maxFails` and `failTimeout`.
|
||||
- The current fail filter only reacts to dial, handshake, or transport failures. It does not react to high latency when the exit is still reachable.
|
||||
- `tunnel_quality_prober` already runs panel-side TCP probes and stores tunnel quality history, but it currently probes representative nodes and does not drive runtime routing decisions.
|
||||
|
||||
## User Decisions
|
||||
|
||||
- Add a `best` option for multi-exit tunnels.
|
||||
- `best` means always choose the current best exit for new connections.
|
||||
- Score exits by end-to-end quality.
|
||||
- Keep the existing public probe target: `www.bing.com:443`.
|
||||
- Do not disrupt established connections.
|
||||
|
||||
## Approach
|
||||
|
||||
Implement `best` as a panel-driven control-plane strategy.
|
||||
|
||||
The database stores the user's intended strategy as `best`. When the panel renders runtime GOST config for a `best` exit group, it sends a GOST selector strategy of `fifo`. The panel dynamically sorts the candidate exits so the current best exit is first. GOST then chooses the first node for new connections.
|
||||
|
||||
This avoids adding active probing logic inside every GOST agent and reuses the existing panel-to-agent command path.
|
||||
|
||||
## Components
|
||||
|
||||
### Frontend
|
||||
|
||||
The tunnel form adds `最优` to the multi-exit load strategy selector.
|
||||
|
||||
- Label: `最优`
|
||||
- Value: `best`
|
||||
- Scope: tunnel forwarding exit groups, alongside `主备/fifo`, `轮询/round`, and `随机/rand`
|
||||
- Create and edit forms must submit and restore `best` unchanged.
|
||||
|
||||
### Backend Data Model
|
||||
|
||||
No schema change is required.
|
||||
|
||||
The existing `chain_tunnel.strategy` column stores `best`. Repository and handler paths should preserve the value in API responses and updates.
|
||||
|
||||
### Runtime Chain Rendering
|
||||
|
||||
When building runtime chain config:
|
||||
|
||||
- If the configured strategy is not `best`, keep existing behavior.
|
||||
- If the configured strategy is `best`, emit GOST selector strategy `fifo`.
|
||||
- Sort the target nodes using the panel's latest best-exit decision before rendering the node list.
|
||||
- If no quality decision exists yet, keep the saved node order.
|
||||
|
||||
This preserves the user's `best` intent in storage while using a GOST selector that can execute the panel's sorted decision.
|
||||
|
||||
### Quality Prober
|
||||
|
||||
Extend `tunnel_quality_prober` to evaluate all candidates in `best` exit groups.
|
||||
|
||||
For each chain owner node and candidate exit, measure:
|
||||
|
||||
- Chain owner node to candidate exit using TCP ping.
|
||||
- Candidate exit to `www.bing.com:443` using TCP ping.
|
||||
|
||||
For direct entry-to-exit tunnels, each entry node owns its own chain decision. For tunnels with intermediate chain hops, each node in the last hop group before the exits owns its own chain decision. This allows different entry or chain nodes to choose different best exits when their path quality differs.
|
||||
|
||||
### Scoring
|
||||
|
||||
Each exit candidate gets an end-to-end score for a specific chain owner node.
|
||||
|
||||
- Total latency is the sum of owner-to-exit latency and exit-to-Bing latency.
|
||||
- Total loss combines both legs by success probability: `1 - (1 - lossA) * (1 - lossB)`.
|
||||
- Failed or unreachable candidates are sorted behind successful candidates.
|
||||
- The score should heavily penalize packet loss so that low-latency but lossy exits are not selected over stable exits.
|
||||
|
||||
A practical scoring formula can be:
|
||||
|
||||
```text
|
||||
score = totalLatencyMs + (totalLossPercent * lossPenaltyMsPerPercent)
|
||||
```
|
||||
|
||||
Use `lossPenaltyMsPerPercent = 100` initially. For example, 5% loss adds 500ms to the score.
|
||||
|
||||
### Switching Rules
|
||||
|
||||
The panel should not update chains on every probe round.
|
||||
|
||||
Switch only when all conditions are true:
|
||||
|
||||
- The candidate best exit is different from the currently applied first exit.
|
||||
- The candidate is successful.
|
||||
- The candidate remains best for consecutive probe rounds.
|
||||
- The candidate beats the current exit by a minimum advantage threshold.
|
||||
- The chain owner node has passed a minimum switch cooldown.
|
||||
|
||||
Initial constants:
|
||||
|
||||
- Consecutive confirmations: 3 rounds.
|
||||
- Switch cooldown: 30 seconds per chain owner node.
|
||||
- Minimum advantage: the candidate score must improve by at least `max(20ms, currentScore * 0.15)`.
|
||||
|
||||
If all exits fail, keep the current runtime order and do not issue a destructive update.
|
||||
|
||||
### Runtime Update
|
||||
|
||||
When a `best` chain owner node changes best exit:
|
||||
|
||||
1. Rebuild that node's `chains_<tunnelID>` payload with the best exit first and remaining candidates sorted by quality for that node.
|
||||
2. Send `UpdateChains` to that chain owner node.
|
||||
3. Do not restart or update tunnel services.
|
||||
4. Record success or failure in logs and in the in-memory decision state.
|
||||
|
||||
This affects only future connections. Existing TCP connections keep using the `net.Conn` created before the update and continue through their original exit.
|
||||
|
||||
### Agent Safety Improvement
|
||||
|
||||
The current agent `UpdateChains` path unregisters the old chain before registering the new chain. This does not kill existing connections, but it creates a small window where a new connection can fail because the chain name is temporarily absent.
|
||||
|
||||
Improve the update path so it parses the new chain first and only replaces the registered chain after parsing succeeds. The replacement window should be as small as possible. If parsing fails, the old chain must remain active.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- If probing one candidate fails, continue scoring other candidates.
|
||||
- If a chain owner node is offline or times out, skip decisions for that owner during the round instead of marking every candidate failed.
|
||||
- If a candidate has no successful required probe data, mark it failed for that round.
|
||||
- If `UpdateChains` fails, keep the current applied order and retry on a later round.
|
||||
- If the tunnel has one exit or an incomplete config, `best` behaves like the saved order and does not trigger dynamic switching.
|
||||
- If `monitor_tunnel_quality_enabled=false`, dynamic `best` switching pauses. The last applied runtime order remains in effect.
|
||||
|
||||
## Observability
|
||||
|
||||
The prober should maintain in-memory decision state per `best` tunnel and chain owner node.
|
||||
|
||||
Useful fields:
|
||||
|
||||
- Tunnel ID and chain owner node ID.
|
||||
- Current applied best exit node ID.
|
||||
- Candidate best exit node ID.
|
||||
- Candidate scores.
|
||||
- Last switch timestamp.
|
||||
- Last switch result.
|
||||
- Reason for not switching, such as cooldown, insufficient advantage, candidate unstable, or all exits failed.
|
||||
|
||||
Initial UI scope is limited to supporting create, update, and display of the `best` strategy. A later enhancement can expose current best exit and candidate scores in the tunnel monitor view.
|
||||
|
||||
## Testing
|
||||
|
||||
Backend tests:
|
||||
|
||||
- Score calculation orders candidates by latency and packet loss.
|
||||
- Packet loss penalty prevents lossy exits from winning only because latency is low.
|
||||
- All-failed candidates do not trigger a switch.
|
||||
- Consecutive confirmation and cooldown prevent flapping.
|
||||
- `strategy=best` persists in `chain_tunnel.strategy` and is returned by tunnel list/get APIs.
|
||||
- Runtime rendering maps `best` to GOST `fifo` and places the chosen best exit first.
|
||||
|
||||
Agent tests:
|
||||
|
||||
- `UpdateChains` parse failure keeps the old chain registered.
|
||||
- Successful `UpdateChains` updates the chain used by new connections.
|
||||
|
||||
Frontend verification:
|
||||
|
||||
- Tunnel form includes `最优` in the exit strategy selector.
|
||||
- Existing tunnels with `strategy=best` render correctly.
|
||||
- Create and update requests submit `best` unchanged.
|
||||
|
||||
Verification commands:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd go-gost && go test ./...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- Do not move existing live connections to a new exit.
|
||||
- Do not add per-tunnel custom probe targets in this phase.
|
||||
- Do not implement active best-exit probing inside GOST agents.
|
||||
- Do not add a detailed best-exit UI dashboard in this phase.
|
||||
@@ -0,0 +1,180 @@
|
||||
# Custom Best-Exit Probe Target Design
|
||||
|
||||
Date: 2026-05-01
|
||||
Status: Approved design
|
||||
|
||||
## Goal
|
||||
|
||||
Allow each tunnel to define the TCP target used for exit-side quality probing instead of always probing `www.bing.com:443`.
|
||||
|
||||
The custom target must be used consistently by:
|
||||
|
||||
- `best` exit scoring: each exit probes the configured target to measure exit-to-public quality.
|
||||
- Tunnel quality monitoring: the existing exit-side quality check probes the same configured target.
|
||||
|
||||
If a tunnel does not configure a target, behavior remains compatible with today: `www.bing.com:443`.
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- Do not add HTTP/HTTPS request probing in this phase. The probe remains TCP host/port measurement.
|
||||
- Do not add a global default target setting in this phase.
|
||||
- Do not require existing tunnels to be edited or migrated manually.
|
||||
- Do not change the `best` switching thresholds, confirmation rounds, cooldowns, or runtime chain ordering semantics.
|
||||
- Do not add frontend test infrastructure.
|
||||
|
||||
## User-Facing Behavior
|
||||
|
||||
Each tunnel form gets a compact quality target section:
|
||||
|
||||
- Host input, placeholder `www.bing.com`.
|
||||
- Port input, placeholder `443`.
|
||||
- Helper text: this target is used for tunnel quality detection and `best` optimal-exit scoring; leaving it empty uses `www.bing.com:443`.
|
||||
|
||||
Tunnel list/get responses include the configured target so edit forms can round-trip it. The UI displays the effective target near quality/best-exit information as `测试目标:host:port`.
|
||||
|
||||
## Data Model
|
||||
|
||||
Add nullable/default-compatible fields to `model.Tunnel`:
|
||||
|
||||
- `ProbeTargetHost string` mapped to `probe_target_host`, `type:text`, default `''`.
|
||||
- `ProbeTargetPort int` mapped to `probe_target_port`, default `0`.
|
||||
|
||||
Effective target resolution:
|
||||
|
||||
- If `ProbeTargetHost` is non-empty and `ProbeTargetPort` is valid, use it.
|
||||
- Otherwise use `www.bing.com:443`.
|
||||
|
||||
The existing `TunnelQuality` persisted fields `exit_to_bing_latency` and `exit_to_bing_loss` remain unchanged for compatibility. They will semantically mean exit-to-configured-test-target after this change. API/UI labels should avoid saying `Bing` for new displays.
|
||||
|
||||
## Validation
|
||||
|
||||
On create/update:
|
||||
|
||||
- Empty host and empty/zero port are allowed and mean default target.
|
||||
- If either host or port is set, validate both as a pair.
|
||||
- Host is trimmed and must not contain URL scheme, path, query, or whitespace.
|
||||
- Host can be a domain, IPv4, or IPv6 literal. Bracketed IPv6 input should be normalized by removing surrounding brackets.
|
||||
- Port must be an integer from `1` to `65535`.
|
||||
- Do not perform network probing during save; external network failures must not block configuration changes.
|
||||
|
||||
Errors should be specific, for example:
|
||||
|
||||
- `测试目标 Host 不能为空`
|
||||
- `测试目标端口必须是 1-65535`
|
||||
- `测试目标 Host 不能包含协议或路径`
|
||||
|
||||
## Backend Flow
|
||||
|
||||
Introduce a small value/helper near the tunnel quality and best-exit code:
|
||||
|
||||
```go
|
||||
type tunnelProbeTarget struct {
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
```
|
||||
|
||||
Helpers:
|
||||
|
||||
- `defaultTunnelProbeTarget() tunnelProbeTarget` returns `www.bing.com:443`.
|
||||
- `normalizeTunnelProbeTarget(host string, port int) (tunnelProbeTarget, bool, error)` validates user input; the boolean indicates whether the user explicitly configured a target.
|
||||
- `effectiveTunnelProbeTarget(tunnel *model.Tunnel) tunnelProbeTarget` returns configured target or default.
|
||||
|
||||
Use the effective target in `tunnelQualityProber.probeTunnel`:
|
||||
|
||||
- Type 1 and unknown tunnel fallback probes entry node to effective target instead of hardcoded Bing.
|
||||
- Type 2 probes the selected/current exit node to effective target instead of hardcoded Bing.
|
||||
- `probeBestExitOwners` receives the effective target and passes it into best-exit owner scoring.
|
||||
|
||||
Use the effective target in `evaluateBestExitOwner`:
|
||||
|
||||
- Owner-to-exit measurement stays unchanged.
|
||||
- Exit-to-public measurement probes `target.Host:target.Port` instead of `bestExitPublicTargetHost:bestExitPublicTargetPort`.
|
||||
- The per-round public probe cache key must include node ID plus target host and port so future extensions cannot reuse measurements across different targets.
|
||||
|
||||
## API Shape
|
||||
|
||||
Tunnel list/get data includes:
|
||||
|
||||
```json
|
||||
{
|
||||
"probeTargetHost": "example.com",
|
||||
"probeTargetPort": 443
|
||||
}
|
||||
```
|
||||
|
||||
For old/default tunnels, return empty host and `0` to represent `use default`. The edit form must preserve default-as-empty unless the user explicitly saves a custom target.
|
||||
|
||||
Quality monitoring response includes effective target display metadata:
|
||||
|
||||
```json
|
||||
{
|
||||
"probeTargetHost": "www.bing.com",
|
||||
"probeTargetPort": 443
|
||||
}
|
||||
```
|
||||
|
||||
Existing `exitToBingLatency` and `exitToBingLoss` keys stay to avoid breaking frontend and external consumers.
|
||||
|
||||
## Frontend Flow
|
||||
|
||||
Extend `ChainTunnel` only if needed for node-level data; the target belongs to the tunnel, so `Tunnel` and `TunnelForm` get:
|
||||
|
||||
- `probeTargetHost?: string`
|
||||
- `probeTargetPort?: number`
|
||||
|
||||
On edit:
|
||||
|
||||
- Populate form fields from tunnel response.
|
||||
- Empty or zero means default target.
|
||||
|
||||
On submit:
|
||||
|
||||
- Trim host.
|
||||
- Convert blank port to `0`.
|
||||
- Send `probeTargetHost` and `probeTargetPort` with create/update payload.
|
||||
|
||||
Display:
|
||||
|
||||
- In the form helper, show default target behavior.
|
||||
- In quality/best-exit display areas, avoid `Bing` wording; prefer `测试目标` or the concrete `host:port`.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- Invalid target input returns a normal API error envelope with a specific message.
|
||||
- Probe failures use existing quality error paths and best-exit scoring failure entries.
|
||||
- If all exit-to-target probes fail, best-exit behavior remains the same as today when all Bing probes fail: no valid best decision is applied from that round.
|
||||
|
||||
## Testing
|
||||
|
||||
Backend tests:
|
||||
|
||||
- Normalize default target when host/port are empty.
|
||||
- Reject partial host/port configuration and invalid port ranges.
|
||||
- Reject host values with URL scheme/path/whitespace.
|
||||
- Create/update tunnel persists `probeTargetHost` and `probeTargetPort`.
|
||||
- `ListTunnels` returns target fields.
|
||||
- `tunnelQualityProber` uses configured target instead of `www.bing.com:443`.
|
||||
- `best` scoring uses configured target for exit-to-target probes.
|
||||
- Empty target preserves old default `www.bing.com:443` behavior.
|
||||
|
||||
Frontend verification:
|
||||
|
||||
- `pnpm run build` passes.
|
||||
- Manual UI check: create/edit tunnel with blank target and custom target, confirm payload and round-trip display.
|
||||
|
||||
## Rollout And Compatibility
|
||||
|
||||
- Existing tunnels continue using `www.bing.com:443` because empty target resolves to default.
|
||||
- SQLite/PostgreSQL schema changes are handled by existing auto-migration.
|
||||
- Historical `TunnelQuality` rows keep existing columns and are not rewritten.
|
||||
- No runtime agent change is required; the panel already performs these quality probes through existing node ping APIs.
|
||||
|
||||
## Open Decisions
|
||||
|
||||
None. User-approved decisions:
|
||||
|
||||
- Per-tunnel fields are `host + port`.
|
||||
- The target applies to both `best` scoring and tunnel quality monitoring.
|
||||
- Probe type remains TCP host/port.
|
||||
- Empty target defaults to `www.bing.com:443`.
|
||||
@@ -0,0 +1,295 @@
|
||||
# 面板本体一键升级设计
|
||||
|
||||
**日期**: 2026-05-04
|
||||
**状态**: 待审核
|
||||
**作者**: AI Assistant
|
||||
|
||||
## 概述
|
||||
|
||||
在 FLVX 管理面板中增加“面板升级”能力,使管理员可以在网页上检查 GitHub Release 并触发面板本体升级。目标是升级整套面板,而不是只升级转发节点或只替换后端二进制。
|
||||
|
||||
本设计采用 Docker Compose 整套升级方案:后端容器通过受限的 Docker socket 能力更新宿主机部署目录中的 `docker-compose.yml` 和 `.env`,然后启动独立的升级 helper 容器,由 helper 拉取新版 backend/frontend 镜像并重新启动 `backend` 与 `frontend` 服务。
|
||||
|
||||
## sub2api 参考结论
|
||||
|
||||
sub2api 的一键升级不是在宿主机执行 `docker compose pull/up`。它的运行形态是单体 Go 服务:前端构建产物 embed 到后端二进制,容器内只运行 `/app/sub2api`。升级接口下载 GitHub Release 中匹配当前系统和架构的 `sub2api_<version>_<os>_<arch>.tar.gz` 以及 `checksums.txt`,校验后把当前 `os.Executable()` 指向的 `/app/sub2api` 改名为 `/app/sub2api.backup`,再把新二进制原子替换到原路径。重启接口延迟调用 `os.Exit(0)`,依赖 Docker Compose 的 `restart: unless-stopped` 拉起同一个容器。
|
||||
|
||||
这种方式在 sub2api 的 Docker 部署中可行,是因为它的前后端在同一个二进制里。FLVX 当前是 `flux-panel-backend` 与 `vite-frontend` 两个容器,替换 `/app/paneld` 只能升级后端,不能升级前端页面。因此 FLVX 的“面板本体升级”需要更新 Compose 版本和两个镜像,而不是照搬二进制替换。
|
||||
|
||||
## 目标
|
||||
|
||||
1. 管理员可在面板上查看当前版本、最新版本、升级通道和升级能力状态。
|
||||
2. 管理员可一键升级整套面板 backend/frontend。
|
||||
3. 升级复用现有 GitHub Release 和 `FLUX_VERSION` 版本机制。
|
||||
4. 升级复用现有 GitHub 加速配置 `github_proxy_enabled` / `github_proxy_url`。
|
||||
5. 升级过程不接受任意命令、任意 URL 或任意 compose 路径。
|
||||
6. 环境不满足时清晰提示不可用原因,不静默失败。
|
||||
|
||||
## 非目标
|
||||
|
||||
1. 不实现 sub2api 式后端二进制替换作为本次主路径。
|
||||
2. 不支持从非本仓库 Release 下载升级资产。
|
||||
3. 不支持普通用户触发升级。
|
||||
4. 不支持在前端执行 shell 命令。
|
||||
5. 不修改 `install.sh` 或 `panel_install.sh` 的本地安装菜单逻辑;发布流程仍可能覆盖这些脚本。
|
||||
6. 不引入前端测试框架。
|
||||
|
||||
## 影响范围
|
||||
|
||||
### 后端
|
||||
|
||||
- `go-backend/internal/http/handler/handler.go`
|
||||
- `go-backend/internal/http/handler/upgrade.go`
|
||||
- 新增 `go-backend/internal/http/handler/system_upgrade.go`
|
||||
- 新增 `go-backend/internal/http/handler/system_upgrade_test.go`
|
||||
- `go-backend/Dockerfile`
|
||||
|
||||
### 部署模板
|
||||
|
||||
- `docker-compose-v4.yml`
|
||||
- `docker-compose-v6.yml`
|
||||
|
||||
### 前端
|
||||
|
||||
- `vite-frontend/src/api/index.ts`
|
||||
- `vite-frontend/src/api/types.ts`
|
||||
- `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
## 运行前提
|
||||
|
||||
升级能力仅在 Docker Compose 部署中可用,并要求后端容器具备以下条件:
|
||||
|
||||
1. 容器内存在 Docker CLI,且支持 `docker compose version`。
|
||||
2. `/var/run/docker.sock` 挂载到后端容器。
|
||||
3. 宿主部署目录挂载到容器内固定路径,例如 `/opt/flvx-panel`。
|
||||
4. 环境变量 `PANEL_DEPLOY_DIR=/opt/flvx-panel`。
|
||||
5. 环境变量 `PANEL_BACKEND_CONTAINER=flux-panel-backend`,为空时默认使用 `flux-panel-backend`;值必须匹配容器名安全字符集 `[A-Za-z0-9_.-]+`。
|
||||
6. 部署目录内存在 `.env` 和 `docker-compose.yml`。
|
||||
|
||||
如果任一条件不满足,检查接口返回 `capable=false` 和明确的 `reason`,升级按钮禁用。
|
||||
|
||||
## 后端设计
|
||||
|
||||
### API
|
||||
|
||||
新增系统升级接口,路径使用 `/api/v1/system/*`,继续受现有 middleware 管控,仅管理员可访问。
|
||||
|
||||
| 方法 | 路径 | 用途 |
|
||||
|------|------|------|
|
||||
| `POST` | `/api/v1/system/version` | 返回当前版本、升级通道、能力状态和可选最新版本 |
|
||||
| `POST` | `/api/v1/system/check-updates` | 强制查询 GitHub Release,返回最新版本和候选列表 |
|
||||
| `POST` | `/api/v1/system/upgrade` | 执行升级 |
|
||||
|
||||
请求体:
|
||||
|
||||
```json
|
||||
{
|
||||
"channel": "stable",
|
||||
"version": ""
|
||||
}
|
||||
```
|
||||
|
||||
`channel` 使用现有节点升级的通道语义:`stable` 匹配纯数字版本,`dev` 匹配 `alpha` / `beta` / `rc`。`version` 为空时自动选择该通道最新 Release。
|
||||
|
||||
`/api/v1/system/version` 返回:
|
||||
|
||||
```json
|
||||
{
|
||||
"currentVersion": "2.1.9-beta14",
|
||||
"channel": "stable",
|
||||
"latestVersion": "2.1.9",
|
||||
"hasUpdate": true,
|
||||
"capable": true,
|
||||
"reason": "",
|
||||
"deployDir": "/opt/flvx-panel",
|
||||
"composeFile": "/opt/flvx-panel/docker-compose.yml",
|
||||
"backendContainer": "flux-panel-backend"
|
||||
}
|
||||
```
|
||||
|
||||
`/api/v1/system/upgrade` 成功返回:
|
||||
|
||||
```json
|
||||
{
|
||||
"version": "2.1.9",
|
||||
"message": "升级 helper 已启动,面板服务将短暂重启",
|
||||
"commands": [
|
||||
"docker run -d --rm --volumes-from flux-panel-backend ...",
|
||||
"docker compose pull backend frontend",
|
||||
"docker compose up -d backend frontend"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
返回的 `commands` 只用于 UI 展示固定步骤,不包含用户输入或 shell 拼接结果。
|
||||
|
||||
### 版本来源
|
||||
|
||||
当前版本优先从容器环境变量读取:
|
||||
|
||||
1. `FLUX_VERSION`
|
||||
2. `VITE_APP_VERSION` 不在后端容器中可靠存在,不作为后端版本来源。
|
||||
3. 为空时返回 `dev`。
|
||||
|
||||
发布流程已经在 `panel_install.sh` 写入 `.env` 的 `FLUX_VERSION`,Compose 模板需要把该变量传给 backend 容器,保证后端可感知当前版本。
|
||||
|
||||
### Release 查询
|
||||
|
||||
复用现有 `fetchGitHubReleases`、`resolveLatestReleaseByChannel`、`normalizeReleaseChannel`、`releaseChannelFromTag`、`releaseChannelLabel` 和 GitHub 加速配置能力。新增函数只负责筛选系统升级所需资产:
|
||||
|
||||
- `docker-compose-v4.yml`
|
||||
- `docker-compose-v6.yml`
|
||||
|
||||
是否下载 v4/v6 compose 文件通过当前部署目录中的 `docker-compose.yml` 判断:如果网络定义包含 `enable_ipv6: true`,选择 `docker-compose-v6.yml`;否则选择 `docker-compose-v4.yml`。
|
||||
|
||||
### 升级执行器
|
||||
|
||||
新增 `systemUpgradeExecutor`,职责明确分为可测试的小函数:
|
||||
|
||||
1. `checkSystemUpgradeCapability()` 检查 Docker CLI、Docker socket、部署目录、`.env`、`docker-compose.yml`。
|
||||
2. `selectComposeAsset(currentCompose []byte) string` 选择 v4/v6 compose 资产。
|
||||
3. `updateEnvVersion(path, version string) error` 原子更新 `.env` 中的 `FLUX_VERSION`。
|
||||
4. `downloadCompose(version, assetName, dest string) error` 下载新版 compose 模板到临时文件。
|
||||
5. `currentBackendImage(containerName string) (string, error)` 获取当前 backend 容器镜像 ID。
|
||||
6. `startSystemUpgradeHelper(version string) error` 启动独立 helper 容器执行固定升级流程。
|
||||
|
||||
升级流程:
|
||||
|
||||
1. 获取全局升级锁,拒绝并发升级。
|
||||
2. 校验目标版本存在且不是 draft。
|
||||
3. 检查升级能力。
|
||||
4. 备份 `.env` 为 `.env.upgrade.bak`,备份 `docker-compose.yml` 为 `docker-compose.yml.upgrade.bak`。
|
||||
5. 下载目标版本的 compose 文件到部署目录临时文件。
|
||||
6. 原子替换 `docker-compose.yml`。
|
||||
7. 原子更新 `.env` 的 `FLUX_VERSION`。
|
||||
8. 通过 Docker socket 查询当前 backend 容器的镜像 ID。
|
||||
9. 使用当前 backend 镜像启动一个不属于 Compose 项目的临时 helper 容器。
|
||||
10. helper 通过 `--volumes-from flux-panel-backend` 继承部署目录挂载,并显式挂载 `/var/run/docker.sock`。
|
||||
11. helper 在 `PANEL_DEPLOY_DIR` 下执行 `docker compose pull backend frontend`。
|
||||
12. helper 等待 5 秒,让 SQLite WAL 等文件刷盘。
|
||||
13. helper 执行 `docker compose up -d backend frontend`,由 Compose 重建前端和后端。
|
||||
14. 后端接口在 helper 成功启动后立即返回;浏览器随后会经历短暂断线。
|
||||
|
||||
PostgreSQL 模式不主动 pull 或重建 `postgres` 服务,避免无关数据库变动。新版 compose 文件仍保留 postgres 配置供后续手动迁移或重建使用。
|
||||
|
||||
### 命令安全
|
||||
|
||||
后端不暴露通用命令执行能力。后端只直接执行 Docker CLI 的固定参数,用于获取当前镜像和启动 helper:
|
||||
|
||||
```go
|
||||
exec.CommandContext(ctx, "docker", "inspect", "-f", "{{.Image}}", backendContainer)
|
||||
exec.CommandContext(ctx, "docker", "run", "-d", "--rm", "--name", helperName,
|
||||
"--volumes-from", backendContainer,
|
||||
"-v", "/var/run/docker.sock:/var/run/docker.sock",
|
||||
"-e", "PANEL_DEPLOY_DIR=/opt/flvx-panel",
|
||||
"--entrypoint", "/bin/sh", imageID,
|
||||
"-c", helperScript)
|
||||
```
|
||||
|
||||
`helperScript` 由后端固定生成,不拼接用户输入:
|
||||
|
||||
```sh
|
||||
cd "$PANEL_DEPLOY_DIR" && docker compose pull backend frontend && sleep 5 && docker compose up -d backend frontend
|
||||
```
|
||||
|
||||
工作目录固定为 `PANEL_DEPLOY_DIR`。`PANEL_DEPLOY_DIR` 必须是绝对路径,且必须包含 `.env` 和 `docker-compose.yml`。接口输入只允许影响 `channel` 和已验证的 Release `version`。
|
||||
|
||||
### 超时和错误处理
|
||||
|
||||
1. Release 查询超时沿用现有 GitHub API 客户端超时。
|
||||
2. 下载 compose 文件使用 60 秒超时。
|
||||
3. 启动 helper 使用 30 秒超时,helper 内部命令不受原 HTTP 请求生命周期影响。
|
||||
4. 任一步失败时返回错误信息,并尽量保留 `.upgrade.bak` 供人工恢复。
|
||||
5. 如果 `.env` 更新后后续步骤失败,不自动回滚镜像或容器,避免误判导致更大破坏;错误信息提示备份文件位置。
|
||||
|
||||
## 部署模板设计
|
||||
|
||||
`docker-compose-v4.yml` 和 `docker-compose-v6.yml` 的 backend 服务增加:
|
||||
|
||||
```yaml
|
||||
environment:
|
||||
FLUX_VERSION: ${FLUX_VERSION:-dev}
|
||||
PANEL_DEPLOY_DIR: /opt/flvx-panel
|
||||
PANEL_BACKEND_CONTAINER: flux-panel-backend
|
||||
volumes:
|
||||
- sqlite_data:/app/data
|
||||
- /var/run/docker.sock:/var/run/docker.sock
|
||||
- ./:/opt/flvx-panel
|
||||
```
|
||||
|
||||
`go-backend/Dockerfile` 的 runtime 镜像通过多阶段构建从官方 `docker:27-cli` 镜像复制 Docker CLI 和 compose 插件到 Debian runtime 镜像,避免依赖 Debian apt 源中的 Docker 包可用性:
|
||||
|
||||
```dockerfile
|
||||
FROM docker:27-cli AS dockercli
|
||||
FROM debian:bookworm-slim
|
||||
COPY --from=dockercli /usr/local/bin/docker /usr/local/bin/docker
|
||||
COPY --from=dockercli /usr/local/libexec/docker/cli-plugins/docker-compose /usr/local/libexec/docker/cli-plugins/docker-compose
|
||||
```
|
||||
|
||||
实现时保留现有 Go builder 和 `/app/paneld` 入口,仅增加 Docker CLI stage 和复制步骤。
|
||||
|
||||
helper 容器使用当前 backend 容器的镜像 ID 启动,而不是额外依赖 `docker:cli` 镜像。这样不引入新的镜像仓库依赖,并保证 helper 内可用的 Docker CLI 与当前后端一致。
|
||||
|
||||
## 前端设计
|
||||
|
||||
在 `vite-frontend/src/pages/config.tsx` 的基本设置或数据库占用附近增加“面板升级”卡片,避免隐藏在节点页导致误解为“节点升级”。
|
||||
|
||||
展示内容:
|
||||
|
||||
1. 当前版本。
|
||||
2. 最新版本。
|
||||
3. 更新通道选择,复用现有 `stable` / `dev` 语义和 `UpdateReleaseChannel` 本地存储。
|
||||
4. 升级能力状态:可用、不可用原因、Docker socket 高权限提示。
|
||||
5. 操作按钮:检查更新、立即升级。
|
||||
|
||||
交互:
|
||||
|
||||
1. 页面加载时调用 `/system/version`。
|
||||
2. 点击“检查更新”调用 `/system/check-updates`。
|
||||
3. 点击“立即升级”前弹出确认框,明确提示服务会短暂中断,并提示 Docker socket 具备宿主高权限。
|
||||
4. 升级请求只等待 helper 启动,超时设置为 60 秒。
|
||||
5. 成功后 toast 显示“升级已触发,面板将在数十秒内重启”,并可提示用户稍后刷新。
|
||||
|
||||
## 安全边界
|
||||
|
||||
Docker socket 挂载等同于给后端容器宿主机级别控制能力。这是本设计的主要风险。缓解措施:
|
||||
|
||||
1. 仅 `/api/v1/system/*` 管理员接口可触发。
|
||||
2. 不提供任意命令执行接口。
|
||||
3. 不允许用户传入下载 URL。
|
||||
4. 不允许用户传入 compose 路径。
|
||||
5. 只升级本仓库 GitHub Release,且跳过 draft。
|
||||
6. 前端明确展示 Docker socket 权限提示。
|
||||
|
||||
## 测试策略
|
||||
|
||||
### Go 单测
|
||||
|
||||
新增 `system_upgrade_test.go` 覆盖:
|
||||
|
||||
1. `selectComposeAsset` 对 v4/v6 compose 内容的判断。
|
||||
2. `.env` 中已有 `FLUX_VERSION` 时更新值。
|
||||
3. `.env` 中缺少 `FLUX_VERSION` 时追加值。
|
||||
4. 缺少部署目录、`.env`、`docker-compose.yml`、Docker socket 时返回不可用原因。
|
||||
5. helper 命令构造固定命令序列,不拼接用户输入。
|
||||
6. 并发升级锁会拒绝第二个升级请求。
|
||||
|
||||
### 手动/集成验证
|
||||
|
||||
1. `go-backend`: `go test ./...`
|
||||
2. `vite-frontend`: `pnpm run build`
|
||||
3. 本地容器验证:启动 Compose 后检查设置页升级卡片可显示能力状态。
|
||||
4. 在无 Docker socket 的开发环境验证按钮禁用并显示原因。
|
||||
|
||||
## 回滚与恢复
|
||||
|
||||
自动升级失败时不做自动容器回滚。后端会保留:
|
||||
|
||||
1. `.env.upgrade.bak`
|
||||
2. `docker-compose.yml.upgrade.bak`
|
||||
|
||||
人工恢复步骤由错误信息提示:进入部署目录,按需恢复备份文件,再执行 `docker compose up -d backend frontend`。
|
||||
|
||||
## 决策记录
|
||||
|
||||
本设计已确定采用 Docker socket 整套升级方案,不再保留二进制替换作为本次实现路径。Docker socket 的权限风险通过管理员限制、命令白名单和前端提示控制。
|
||||
@@ -0,0 +1,377 @@
|
||||
# FLVX 安全问题修复设计
|
||||
|
||||
**日期**: 2026-05-13
|
||||
**状态**: 待审核
|
||||
**作者**: AI Assistant
|
||||
|
||||
## 概述
|
||||
|
||||
针对 PR #502 提到的安全问题,对 FLVX 后端认证、配置访问控制、配置写入保护、备份导出和 JWT 失效模型做一次集中修复。目标是优先消除高风险漏洞,同时保留当前必须兼容的登录页品牌配置读取和验证码兼容行为。
|
||||
|
||||
本设计采用“高危项一次收口,结构性问题只分析不重构”的策略:本轮修复 MD5 密码存储、未受控配置读取、敏感配置写入、备份配置泄露和 JWT 长期有效且改密后不失效的问题;不修改“无 Cloudflare secret 时允许当前 captcha 兼容行为”,也不重构 `autoMigrateAll()` 与 `migrateSchema()` 的双迁移入口。
|
||||
|
||||
## 背景
|
||||
|
||||
当前主线存在以下已确认问题:
|
||||
|
||||
1. `login`、`open_api/sub_store`、用户改密、管理员创建用户和管理员修改用户密码仍然使用 `security.MD5(...)`。
|
||||
2. `/api/v1/config/get` 在 middleware 的 `shouldSkip()` 中被匿名放行,导致任意调用方可以读取绝大多数配置。
|
||||
3. `updateConfigs()` 与 `updateSingleConfig()` 使用了两套不同的限制逻辑,敏感配置键在单项写接口中未被保护。
|
||||
4. `ExportAll()` 和 `ExportPartial(types=["configs"])` 会直接导出所有配置,包含 `jwt_secret`、`license_key`、`cloudflare_secret_key`。
|
||||
5. JWT 当前有效期为 90 天,且 token 在用户改密、禁用、角色变化后仍可继续使用到过期。
|
||||
|
||||
同时存在两个重要约束:
|
||||
|
||||
1. 登录页和未登录态品牌展示依赖匿名读取 `app_name`、`app_logo`、`app_favicon`、`app_bg_image` 和 `cloudflare_site_key`。
|
||||
2. `tests/contract/migration_contract_test.go` 已把“无 Cloudflare secret 时允许当前 captcha 兼容行为”定义为既有契约,本轮不改变。
|
||||
|
||||
## 目标
|
||||
|
||||
1. 新增和更新后的用户密码不再以 MD5 存储。
|
||||
2. 历史 MD5 用户可在首次成功认证时自动迁移到强哈希。
|
||||
3. 匿名请求不能再读取任意配置,只能读取明确的公开配置白名单。
|
||||
4. 通用配置写接口不能覆盖敏感配置键。
|
||||
5. 备份导出默认不泄露敏感配置明文。
|
||||
6. 用户改密、禁用或角色变化后,旧 JWT 应立即失效。
|
||||
7. 不破坏现有登录页品牌展示和 captcha 兼容行为。
|
||||
|
||||
## 非目标
|
||||
|
||||
1. 不重构 `open_api/sub_store` 的整体认证模型;该接口仍使用现有用户名和密码查询参数语义。
|
||||
2. 不实现完整 refresh token、session 管理后台或 token 黑名单体系。
|
||||
3. 不改变“无 Cloudflare secret 时允许当前 captcha 兼容行为”。
|
||||
4. 不在本轮重构 `autoMigrateAll()` 与 `migrateSchema()` 的启动流程。
|
||||
5. 不引入前端测试框架。
|
||||
|
||||
## 影响范围
|
||||
|
||||
### 后端
|
||||
|
||||
- `go-backend/internal/security/`
|
||||
- `go-backend/internal/auth/jwt.go`
|
||||
- `go-backend/internal/http/middleware/auth.go`
|
||||
- `go-backend/internal/http/handler/handler.go`
|
||||
- `go-backend/internal/http/handler/mutations.go`
|
||||
- `go-backend/internal/store/model/model.go`
|
||||
- `go-backend/internal/store/repo/repository.go`
|
||||
- `go-backend/internal/store/repo/repository_mutations.go`
|
||||
- `go-backend/tests/contract/`
|
||||
- `go-backend/internal/store/repo/*_test.go`
|
||||
|
||||
### 前端
|
||||
|
||||
- `vite-frontend/src/api/index.ts`
|
||||
- `vite-frontend/src/config/site.ts`
|
||||
- `vite-frontend/src/pages/index.tsx`
|
||||
- 任何在未登录态读取品牌配置的组件
|
||||
|
||||
## 设计决策
|
||||
|
||||
### 已确认决策
|
||||
|
||||
1. 本轮采用安全优先策略,允许收紧危险默认行为。
|
||||
2. MD5 密码采用“登录成功时自动迁移”的兼容方案。
|
||||
3. captcha 在未配置 Cloudflare secret 时的兼容行为保持不变。
|
||||
4. JWT 采用“最小可撤销”方案,而不是完整 session 体系。
|
||||
5. 双重迁移系统只分析,不在本轮中修改。
|
||||
|
||||
### 迁移系统分析结论
|
||||
|
||||
`autoMigrateAll()` 与 `migrateSchema()` 当前职责并不相同:
|
||||
|
||||
1. `autoMigrateAll()` 负责表和列结构补齐。
|
||||
2. `migrateSchema()` 负责基于 `schema_version` 的数据修正,以及 PostgreSQL ID 默认值修复等兼容迁移。
|
||||
3. 现有 `repository_migrate_test.go` 已明确覆盖这两部分逻辑,说明它们在现有代码库中是被依赖的互补结构,而不是已确认的重复安全漏洞。
|
||||
|
||||
因此本轮仅记录该分析结论,不把双迁移入口纳入改动范围,避免把安全修复扩展为启动流程重构。
|
||||
|
||||
## 详细设计
|
||||
|
||||
### 1. 密码存储与认证迁移
|
||||
|
||||
在 `internal/security/` 中新增统一密码能力,替代各处直接使用 `security.MD5(...)` 的做法。
|
||||
|
||||
建议新增以下接口:
|
||||
|
||||
```go
|
||||
func HashPassword(plain string) (string, error)
|
||||
func VerifyPassword(storedHash, plain string) (ok bool, legacy bool)
|
||||
func IsLegacyPasswordHash(storedHash string) bool
|
||||
```
|
||||
|
||||
哈希算法使用 `bcrypt`:
|
||||
|
||||
1. `user.pwd` 当前为 `varchar(100)`,足以容纳 bcrypt 哈希。
|
||||
2. 不需要修改密码列长度,改动最小。
|
||||
3. 对当前 Go 后端来说,bcrypt 是最稳妥的强哈希升级路径。
|
||||
|
||||
所有密码入口统一改为走这套能力:
|
||||
|
||||
1. `login`
|
||||
2. `openAPISubStore`
|
||||
3. `updatePassword`
|
||||
4. `userCreate`
|
||||
5. `userUpdate` 中的管理员改密路径
|
||||
|
||||
认证迁移规则:
|
||||
|
||||
1. 如果数据库中存的是 bcrypt,则按 bcrypt 校验。
|
||||
2. 如果数据库中存的是历史 MD5,则先按旧逻辑校验。
|
||||
3. 历史 MD5 校验成功后,立即把 `pwd` 改写为 bcrypt。
|
||||
4. 自动迁移不仅在网页登录时执行,也在 `open_api/sub_store` 成功鉴权时执行,避免只使用订阅接口的老用户永远停留在 MD5。
|
||||
|
||||
默认管理员种子账号仍保留当前默认密码语义和 `requirePasswordChange` 行为,但种子哈希改为 bcrypt,不再在新建数据库中写入 MD5 值。
|
||||
|
||||
### 2. JWT 最小可撤销方案
|
||||
|
||||
本轮不引入 refresh token 和黑名单表,而是做一个可以立即生效的最小撤销闭环。
|
||||
|
||||
#### 数据模型
|
||||
|
||||
在 `user` 表新增字段:
|
||||
|
||||
- `password_changed_at BIGINT NOT NULL DEFAULT 0`
|
||||
|
||||
该字段专门表示密码最后一次变更时间,不能复用现有 `updated_time`,原因是 `updated_time` 还会被流量、状态或其他用户资料更新触发,复用后会让非密码更新错误地使 token 失效。
|
||||
|
||||
#### token 签发与校验
|
||||
|
||||
继续使用现有 `iat` 声明,但把有效期从 90 天收紧到 7 天。
|
||||
|
||||
token 校验分两步:
|
||||
|
||||
1. 先做现有签名和过期时间校验。
|
||||
2. 再读取用户最小认证状态,确认:
|
||||
- 用户仍存在
|
||||
- 用户状态未被禁用
|
||||
- 当前 `role_id` 与 token 中一致
|
||||
- `claims.iat` 不早于 `password_changed_at`
|
||||
|
||||
为避免 middleware 每次都查询完整用户对象,Repository 新增专用读取方法,只返回 token 校验需要的最小字段,例如:
|
||||
|
||||
```go
|
||||
type UserAuthState struct {
|
||||
ID int64
|
||||
RoleID int
|
||||
Status int
|
||||
PasswordChangedAt int64
|
||||
}
|
||||
|
||||
func (r *Repository) GetUserAuthState(userID int64) (*UserAuthState, error)
|
||||
```
|
||||
|
||||
#### 失效语义
|
||||
|
||||
以下场景下,旧 token 应立即失效:
|
||||
|
||||
1. 用户修改密码
|
||||
2. 管理员修改用户密码
|
||||
3. 用户被禁用
|
||||
4. 用户角色发生变化
|
||||
|
||||
这会带来一次明确的兼容收紧:升级完成后,部分历史 token 可能因为寿命策略或认证状态变化而失效,这是安全优先下的可接受行为。
|
||||
|
||||
### 3. 配置读取访问控制
|
||||
|
||||
为了避免继续让 `/api/v1/config/get` 承担“有时匿名、有时鉴权”的混合语义,本设计将公开配置读取拆成单独的 public 端点。
|
||||
|
||||
#### 端点设计
|
||||
|
||||
保留现有受保护端点:
|
||||
|
||||
- `POST /api/v1/config/get`
|
||||
|
||||
新增公开端点:
|
||||
|
||||
- `POST /api/v1/public/config/get`
|
||||
|
||||
middleware 仅对白名单 public 端点放行,不再放行 `/api/v1/config/get`。
|
||||
|
||||
#### 公开白名单
|
||||
|
||||
匿名仅允许读取以下配置:
|
||||
|
||||
1. `app_name`
|
||||
2. `app_logo`
|
||||
3. `app_favicon`
|
||||
4. `app_bg_image`
|
||||
5. `cloudflare_site_key`
|
||||
|
||||
理由:
|
||||
|
||||
1. 登录页与未登录态品牌渲染依赖前四项。
|
||||
2. 登录页在 captcha 开启时需要读取 `cloudflare_site_key`。
|
||||
3. 其他配置不应暴露给匿名方。
|
||||
|
||||
前端调整规则:
|
||||
|
||||
1. 登录页和 `site.ts` 中的未登录态品牌配置读取改走 `/public/config/get`。
|
||||
2. 登录后页面仍使用现有 `/config/get` 或 `/config/list`。
|
||||
3. 已登录页面中的配置读取逻辑不变,只是恢复为真正受 JWT 保护。
|
||||
|
||||
### 4. 配置写保护统一
|
||||
|
||||
当前 `updateConfigs()` 与 `updateSingleConfig()` 各自维护不同限制逻辑,是本次越权写入漏洞的根源。本轮把配置访问规则统一收口为一套辅助函数。
|
||||
|
||||
建议新增配置策略定义:
|
||||
|
||||
```go
|
||||
type ConfigAccessPolicy struct {
|
||||
PublicReadable bool
|
||||
Sensitive bool
|
||||
CommercialOnly bool
|
||||
}
|
||||
```
|
||||
|
||||
由统一函数返回某个 key 的策略,再由:
|
||||
|
||||
1. `public config get`
|
||||
2. `config get`
|
||||
3. `config list`
|
||||
4. `updateConfigs()`
|
||||
5. `updateSingleConfig()`
|
||||
|
||||
共同复用。
|
||||
|
||||
敏感配置键至少包含:
|
||||
|
||||
1. `jwt_secret`
|
||||
2. `license_key`
|
||||
3. `cloudflare_secret_key`
|
||||
|
||||
这些键的写入规则:
|
||||
|
||||
1. 不允许通过通用配置写接口改写。
|
||||
2. 不允许通过公开读取接口读取。
|
||||
3. 非管理员在配置列表接口中也不能获得。
|
||||
|
||||
商业版白名单键继续沿用现有语义,例如:
|
||||
|
||||
1. `app_name`
|
||||
2. `app_logo`
|
||||
3. `app_favicon`
|
||||
4. `hide_footer_brand`
|
||||
|
||||
但其判断逻辑同样统一走同一套策略函数,避免再次出现单接口漏判。
|
||||
|
||||
### 5. 备份导出与导入脱敏
|
||||
|
||||
备份系统改为“默认安全导出”,而不是“完整明文镜像”。
|
||||
|
||||
#### 导出
|
||||
|
||||
`ExportAll()` 和 `ExportPartial(types=["configs"])` 在写入 `backup.Configs` 前都先经过统一过滤函数,移除敏感配置键。
|
||||
|
||||
敏感配置键与配置写保护列表保持一致:
|
||||
|
||||
1. `jwt_secret`
|
||||
2. `license_key`
|
||||
3. `cloudflare_secret_key`
|
||||
|
||||
#### 导入
|
||||
|
||||
导入配置时,即使旧备份中带有上述敏感键,也会在导入前被丢弃,不允许通过备份恢复路径覆盖在线安全配置。
|
||||
|
||||
该设计的取舍如下:
|
||||
|
||||
1. 保留大部分业务配置、节点、转发、用户数据的恢复能力。
|
||||
2. 不再把备份文件当作核心密钥分发载体。
|
||||
3. `UserBackup.Pwd` 仍然保留,以维持用户恢复语义;在本轮密码升级后,这些值将是 bcrypt 哈希,而不是 MD5。
|
||||
|
||||
### 6. captcha 兼容行为
|
||||
|
||||
`captcha_enabled`、`cloudflare_site_key`、`cloudflare_secret_key` 的现有兼容行为保持不变。
|
||||
|
||||
明确保持以下现状:
|
||||
|
||||
1. 当未完整配置 Cloudflare key 时,当前 contract test 约定的兼容路径继续存在。
|
||||
2. 本轮不把 captcha 兼容逻辑从“兼容旧行为”切换为“严格校验”。
|
||||
|
||||
这样可以避免把一轮安全修复扩展成登录流程行为变更,同时与用户已确认的范围保持一致。
|
||||
|
||||
## 错误处理与兼容行为
|
||||
|
||||
### 错误处理
|
||||
|
||||
保持现有 API envelope:`{code, msg, data, ts}`。
|
||||
|
||||
建议的接口行为:
|
||||
|
||||
1. `POST /api/v1/public/config/get` 请求非公开 key 时返回 `403`。
|
||||
2. 受保护配置端点未登录时返回 `401`。
|
||||
3. 登录、订阅接口、改密接口继续返回通用认证失败,不暴露“用户名存在但密码错误”等细节。
|
||||
4. token 因签名错误、过期、改密、禁用或角色变化失效时,统一返回现有 `401` 语义。
|
||||
5. 备份导入中出现敏感配置键时,接口整体仍允许成功导入其他数据,敏感键静默忽略。
|
||||
|
||||
### 保留兼容
|
||||
|
||||
1. 登录页和未登录态品牌展示继续可用。
|
||||
2. 未配置 Cloudflare secret 时的 captcha 兼容逻辑继续保留。
|
||||
3. 历史 MD5 用户仍可继续认证,并在成功后自动迁移。
|
||||
|
||||
### 刻意收紧
|
||||
|
||||
1. 匿名方不再可读取任意配置。
|
||||
2. 通用配置写接口不再能写入敏感键。
|
||||
3. 备份不再导出敏感配置明文。
|
||||
4. 改密、禁用和角色变化会立即使旧 token 失效。
|
||||
|
||||
## 测试设计
|
||||
|
||||
本轮以 Go 单测和 contract test 为主,覆盖以下场景。
|
||||
|
||||
### 密码迁移
|
||||
|
||||
1. 历史 MD5 用户在网页登录成功后,数据库中的 `pwd` 被升级为 bcrypt。
|
||||
2. 历史 MD5 用户在 `open_api/sub_store` 成功鉴权后,同样触发迁移。
|
||||
3. 新建用户后落库的是 bcrypt,而不是 MD5。
|
||||
4. 管理员修改用户密码和用户自助改密后,落库的是 bcrypt。
|
||||
|
||||
### JWT 最小可撤销
|
||||
|
||||
1. 正常 token 仍可访问受保护接口。
|
||||
2. 改密后旧 token 失效。
|
||||
3. 用户被禁用后旧 token 失效。
|
||||
4. 用户角色变化后旧 token 失效。
|
||||
5. 过期 token 失效。
|
||||
|
||||
### 配置访问控制
|
||||
|
||||
1. 匿名访问公开配置成功。
|
||||
2. 匿名访问非公开配置失败。
|
||||
3. 已登录页面需要的普通配置读取仍然可用。
|
||||
4. `updateSingleConfig()` 无法修改敏感键。
|
||||
5. `updateConfigs()` 同样无法修改敏感键。
|
||||
|
||||
### 备份脱敏
|
||||
|
||||
1. `ExportAll()` 不包含敏感配置。
|
||||
2. `ExportPartial(types=["configs"])` 不包含敏感配置。
|
||||
3. 导入带敏感键的备份时,这些键不会被写回数据库。
|
||||
4. 非敏感配置和其他业务数据仍可正常导入导出。
|
||||
|
||||
### 迁移系统回归
|
||||
|
||||
1. 现有 `repository_migrate_test.go` 保持通过。
|
||||
2. 本轮不对 `autoMigrateAll()` 与 `migrateSchema()` 的职责边界做行为性改动。
|
||||
|
||||
## 验收标准
|
||||
|
||||
1. 数据库中不再新增 MD5 密码。
|
||||
2. 历史 MD5 用户可在首次成功认证后自动升级到 bcrypt。
|
||||
3. 匿名调用方不能再读取非公开配置。
|
||||
4. `updateSingleConfig()` 和 `updateConfigs()` 都无法改写敏感配置键。
|
||||
5. 备份导出默认不包含 `jwt_secret`、`license_key`、`cloudflare_secret_key`。
|
||||
6. 改密、禁用和角色变化后,旧 JWT 立即失效。
|
||||
7. 登录页品牌展示和 captcha 兼容行为不被破坏。
|
||||
8. 现有迁移测试和本轮新增安全测试全部通过。
|
||||
|
||||
## PR #502 处置
|
||||
|
||||
PR #502 的价值在于指出了真实问题,但其实现方式只是把讽刺性注释写进生产代码,并未修复漏洞。因此该 PR 不应合并。
|
||||
|
||||
执行阶段的处置方式:
|
||||
|
||||
1. 关闭 PR #502。
|
||||
2. 在关闭说明中指出:问题成立,但修复将通过正式代码与测试提交完成,而不是通过向源文件加入讽刺性注释。
|
||||
3. 后续在新提交中按本设计逐项修复。
|
||||
@@ -0,0 +1,231 @@
|
||||
# FLVX Dependabot 告警修复设计
|
||||
|
||||
**日期**: 2026-05-14
|
||||
**状态**: 待审核
|
||||
**范围**: 仅处理 GitHub Dependabot 依赖告警
|
||||
|
||||
## 概述
|
||||
|
||||
本设计针对 `Sagit-chu/flvx` 当前 open Dependabot alerts 制定依赖修复方案。目标是在不混入业务安全逻辑改造的前提下,消除或最大限度降低依赖漏洞告警,并通过各模块现有构建和测试命令验证兼容性。
|
||||
|
||||
当前告警共 21 条:
|
||||
|
||||
- Critical: 1
|
||||
- High: 6
|
||||
- Medium: 13
|
||||
- Low: 1
|
||||
|
||||
按生态划分:
|
||||
|
||||
- Go: 14
|
||||
- npm: 7
|
||||
|
||||
Go 告警中有一部分因为 `go-gost/go.mod` 和 `go-gost/x/go.mod` 同时被扫描而重复出现;实际修复应按依赖和模块关系聚合处理,而不是按 alert 数逐条机械修改。
|
||||
|
||||
## 目标
|
||||
|
||||
1. 修复 `go-backend` 中 `github.com/jackc/pgx/v5` 的 critical 和 low 告警。
|
||||
2. 修复 `vite-frontend` 中 npm 直接依赖、开发依赖和 lockfile 传递依赖告警。
|
||||
3. 修复 `go-gost` 与 `go-gost/x` 中可升级到 patched version 的 Go 依赖告警。
|
||||
4. 对 Dependabot 未给出 patched version 的依赖进行单独确认,避免盲目大版本升级。
|
||||
5. 保持改动最小化,便于回滚和定位 CI 失败。
|
||||
|
||||
## 非目标
|
||||
|
||||
1. 不处理既有认证、配置读取、备份导出、JWT 失效等业务安全逻辑问题。
|
||||
2. 不合并或修改 `2026-05-13-security-remediation` 相关设计和计划。
|
||||
3. 不进行 `go get -u ./...` 或 `pnpm update` 级别的大范围依赖升级。
|
||||
4. 不引入前端测试框架。
|
||||
5. 不编辑 `install.sh`、`panel_install.sh` 或 generated `.pb.go` 文件。
|
||||
|
||||
## 影响范围
|
||||
|
||||
### Backend
|
||||
|
||||
- `go-backend/go.mod`
|
||||
- `go-backend/go.sum`
|
||||
|
||||
### Frontend
|
||||
|
||||
- `vite-frontend/package.json`
|
||||
- `vite-frontend/pnpm-lock.yaml`
|
||||
|
||||
### Agent
|
||||
|
||||
- `go-gost/go.mod`
|
||||
- `go-gost/go.sum`
|
||||
- `go-gost/x/go.mod`
|
||||
- `go-gost/x/go.sum`
|
||||
|
||||
`go-gost/go.mod` 使用:
|
||||
|
||||
```go
|
||||
replace github.com/go-gost/x => ./x
|
||||
```
|
||||
|
||||
因此 `go-gost/x` 的依赖修复应先完成,再验证 `go-gost` 主模块。
|
||||
|
||||
## 修复策略
|
||||
|
||||
采用“分模块、最小安全升级”策略。
|
||||
|
||||
### 1. go-backend
|
||||
|
||||
Dependabot alerts:
|
||||
|
||||
- `github.com/jackc/pgx/v5 < 5.9.0`
|
||||
- severity: critical
|
||||
- summary: memory-safety vulnerability
|
||||
- `github.com/jackc/pgx/v5 < 5.9.2`
|
||||
- severity: low
|
||||
- summary: SQL injection via placeholder confusion with dollar quoted string literals
|
||||
|
||||
当前版本:
|
||||
|
||||
- `github.com/jackc/pgx/v5 v5.7.3`
|
||||
|
||||
目标版本:
|
||||
|
||||
- `github.com/jackc/pgx/v5 v5.9.2`
|
||||
|
||||
设计说明:
|
||||
|
||||
- 直接升到 `v5.9.2`,同时覆盖 `v5.9.0` 和 `v5.9.2` 的修复要求。
|
||||
- 不调整 GORM PostgreSQL driver,除非 `go mod tidy` 或测试显示必须联动升级。
|
||||
- 验证以 backend 全量测试为准。
|
||||
|
||||
验证命令:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
```
|
||||
|
||||
### 2. vite-frontend
|
||||
|
||||
Dependabot alerts:
|
||||
|
||||
- `postcss < 8.5.10`
|
||||
- appears in `vite-frontend/package.json`
|
||||
- appears in `vite-frontend/pnpm-lock.yaml`
|
||||
- `serialize-javascript <= 7.0.2` and `< 7.0.5`
|
||||
- lockfile includes `serialize-javascript@6.0.2`
|
||||
- package override currently pins `serialize-javascript` to `7.0.3`
|
||||
- `fast-uri <= 3.1.1`
|
||||
- lockfile currently includes `fast-uri@3.1.0`
|
||||
- `@babel/plugin-transform-modules-systemjs <= 7.29.3`
|
||||
- lockfile currently includes `7.29.0`
|
||||
|
||||
目标版本:
|
||||
|
||||
- `postcss >= 8.5.10`
|
||||
- `serialize-javascript >= 7.0.5`
|
||||
- `fast-uri >= 3.1.2`
|
||||
- `@babel/plugin-transform-modules-systemjs >= 7.29.4`
|
||||
|
||||
设计说明:
|
||||
|
||||
- 对直接声明的 `postcss` 更新 `package.json`。
|
||||
- 将 `overrides.serialize-javascript` 从 `7.0.3` 更新到 `7.0.5`。
|
||||
- 对只出现在 lockfile 的传递依赖,优先通过 `pnpm install` 重新解析 lockfile,让上游范围自然选择 patched version。
|
||||
- 如果 lockfile 仍保留 vulnerable 版本,再添加精确 `pnpm.overrides` 或现有 `overrides` 条目,避免无关依赖大升级。
|
||||
- 不引入前端测试框架,验证使用现有 build。
|
||||
|
||||
验证命令:
|
||||
|
||||
```bash
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
可选补充检查:
|
||||
|
||||
```bash
|
||||
(cd vite-frontend && pnpm why postcss serialize-javascript fast-uri @babel/plugin-transform-modules-systemjs)
|
||||
```
|
||||
|
||||
### 3. go-gost/x
|
||||
|
||||
Dependabot alerts:
|
||||
|
||||
- `github.com/sirupsen/logrus < 1.8.3`
|
||||
- severity: high
|
||||
- current: `v1.8.1`
|
||||
- target: at least `v1.8.3`
|
||||
- `github.com/quic-go/quic-go < 0.57.0`
|
||||
- severity: medium
|
||||
- current: `v0.49.1`
|
||||
- target: `v0.57.0`
|
||||
- `github.com/quic-go/webtransport-go <= 0.9.0`
|
||||
- severity: medium
|
||||
- current: `v0.8.1-0.20241018022711-4ac2c9250e66`
|
||||
- target: `v0.10.0`
|
||||
- `github.com/pion/dtls/v2 <= 2.2.12`
|
||||
- severity: medium
|
||||
- current: `v2.2.6`
|
||||
- target: no patched version provided by Dependabot
|
||||
|
||||
设计说明:
|
||||
|
||||
- 先处理 `go-gost/x`,因为它是 `go-gost` 通过 `replace` 使用的本地模块。
|
||||
- 将 `logrus`、`quic-go` 和 `webtransport-go` 升级到 Dependabot 标出的 patched version。
|
||||
- 单独处理 `pion/dtls/v2`,因为 Dependabot 没有提供 `first_patched_version`。
|
||||
- 对 `pion/dtls/v2`,先查询可用 module versions 和 advisory 详情。如果存在 patched `v2` release,使用最小安全修复版本;如果不存在 patched version,则记录残留告警,避免在没有兼容性评估的情况下强行做高风险大版本迁移。
|
||||
- 升级完成后,在 `go-gost/x` 中运行 `go mod tidy` 并编译/测试该模块。
|
||||
|
||||
验证命令:
|
||||
|
||||
```bash
|
||||
(cd go-gost/x && go test ./...)
|
||||
```
|
||||
|
||||
### 4. go-gost
|
||||
|
||||
Dependabot reports the same vulnerable Go dependencies in `go-gost/go.mod`.
|
||||
|
||||
设计说明:
|
||||
|
||||
- `go-gost/x` 修复后,再更新 `go-gost` 的 module requirements,使主模块也解析到 patched versions。
|
||||
- 保留 `replace github.com/go-gost/x => ./x`。
|
||||
- 使用针对具体漏洞依赖的 `go get` 命令,不使用宽泛的 `go get -u`。
|
||||
- 定向升级后运行 `go mod tidy`。
|
||||
|
||||
验证命令:
|
||||
|
||||
```bash
|
||||
(cd go-gost && go test ./...)
|
||||
(cd go-gost && go build .)
|
||||
```
|
||||
|
||||
## 处理顺序
|
||||
|
||||
1. 修复 `go-backend` 的 `pgx/v5`。
|
||||
2. 修复 `vite-frontend` 的 npm dependencies 和 lockfile。
|
||||
3. 修复 `go-gost/x` 的 Go dependencies。
|
||||
4. 同步并验证 `go-gost`。
|
||||
5. 再次查询 Dependabot alerts,确认 alert 数量下降,或记录有意保留的未解决告警。
|
||||
|
||||
这个顺序可以降低耦合:backend 和 frontend 能独立验证,而 `go-gost/x` 因本地 module replacement 必须先于 `go-gost` 处理。
|
||||
|
||||
## 错误处理与回退
|
||||
|
||||
如果定向依赖升级无法解析:
|
||||
|
||||
1. 使用 `go mod why`、`go mod graph` 或 `pnpm why` 检查依赖链。
|
||||
2. 优先添加最小显式 requirement 或 override,以强制解析到 patched version。
|
||||
3. 除非定向解析不可行,否则避免宽泛升级。
|
||||
4. 如果 patched version 不可用,记录准确 advisory、受影响依赖、当前暴露面,以及保留 open 状态的原因。
|
||||
|
||||
如果验证失败:
|
||||
|
||||
1. 将失败范围限制在当前升级的模块内。
|
||||
2. 先分析编译错误或测试失败,再决定是否调整版本。
|
||||
3. 优先选择能通过测试和构建的最低 patched version。
|
||||
4. 不通过删除测试或修改无关应用代码来掩盖失败。
|
||||
|
||||
## 成功标准
|
||||
|
||||
1. `pgx/v5` 升级后,`go-backend` 测试通过。
|
||||
2. npm dependency 和 lockfile 更新后,frontend production build 通过。
|
||||
3. 定向升级后,`go-gost/x` 测试通过。
|
||||
4. 同步 module requirements 后,`go-gost` 测试和构建通过。
|
||||
5. 最终 Dependabot API 查询显示所有可修复告警已关闭或数量明确下降。
|
||||
6. 任何剩余告警都有明确记录;尤其是 `github.com/pion/dtls/v2` 如果不存在 patched version,需要记录原因和下一步动作。
|
||||
@@ -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 后续补齐。
|
||||
@@ -0,0 +1,197 @@
|
||||
# 规则流量清零设计
|
||||
|
||||
## 背景
|
||||
|
||||
Issue #523 希望“规则”页面中每条隧道规则显示的流量使用量支持手动清零。
|
||||
|
||||
当前规则流量保存在 `forward.in_flow` 和 `forward.out_flow`。流量上报时,同一份增量还会累计到用户总流量、用户隧道流量和相关配额统计中。因此,本功能必须将“规则展示计数器清零”与“用户或隧道配额重置”严格区分。
|
||||
|
||||
## 目标
|
||||
|
||||
为单条规则提供手动流量清零能力:
|
||||
|
||||
- 将所选规则的上传流量和下载流量清零。
|
||||
- 管理员可以清零任意规则。
|
||||
- 普通用户只能清零自己的规则。
|
||||
- 清零后,新产生的流量继续从零正常累计。
|
||||
|
||||
## 非目标
|
||||
|
||||
本功能不会:
|
||||
|
||||
- 修改用户总流量 `user.in_flow` 或 `user.out_flow`。
|
||||
- 修改用户隧道流量 `user_tunnel.in_flow` 或 `user_tunnel.out_flow`。
|
||||
- 修改每日或每月配额用量。
|
||||
- 修改历史流量统计。
|
||||
- 重置 nftables 节点计数器或其增量计算基线。
|
||||
- 重启、暂停、恢复或重新部署规则服务。
|
||||
- 增加批量流量清零功能。
|
||||
|
||||
## 后端设计
|
||||
|
||||
### API
|
||||
|
||||
新增接口:
|
||||
|
||||
```text
|
||||
POST /api/v1/forward/reset-flow
|
||||
```
|
||||
|
||||
请求体:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": 123
|
||||
}
|
||||
```
|
||||
|
||||
成功响应沿用统一 envelope:
|
||||
|
||||
```json
|
||||
{
|
||||
"code": 0,
|
||||
"msg": "success",
|
||||
"data": null,
|
||||
"ts": 0
|
||||
}
|
||||
```
|
||||
|
||||
具体 `msg`、`data` 和 `ts` 值继续由现有 response helper 生成。
|
||||
|
||||
### 参数与权限校验
|
||||
|
||||
Handler 执行以下步骤:
|
||||
|
||||
1. 只接受 `POST` 请求。
|
||||
2. 从 JSON 请求体读取正整数规则 ID。
|
||||
3. 调用现有 `resolveForwardAccess`:
|
||||
- 管理员角色可以访问任意存在的规则。
|
||||
- 普通用户仅能访问 `forward.user_id` 等于当前用户 ID 的规则。
|
||||
- 对普通用户访问他人规则的情况,沿用现有逻辑返回“转发不存在”,避免暴露规则存在性。
|
||||
4. 调用 Repository 完成清零。
|
||||
5. 返回统一成功响应。
|
||||
|
||||
### Repository
|
||||
|
||||
新增方法:
|
||||
|
||||
```go
|
||||
func (r *Repository) ResetForwardFlow(forwardID int64, now int64) error
|
||||
```
|
||||
|
||||
该方法只更新指定 `forward` 记录:
|
||||
|
||||
```text
|
||||
in_flow = 0
|
||||
out_flow = 0
|
||||
updated_time = now
|
||||
```
|
||||
|
||||
Repository 不直接操作 Handler 的身份信息,也不更新任何其他表。
|
||||
|
||||
### 并发与后续流量
|
||||
|
||||
清零使用单条 SQL `UPDATE`。agent 流量上报和 nftables 流量采集仍使用原有增量累加逻辑。清零不会重置采集基线,因此下一次采集只会把清零之后新计算出的增量加回规则计数,不会把清零前的累计值整体恢复。
|
||||
|
||||
若清零 SQL 与流量增量 SQL 同时执行,数据库按实际语句执行顺序决定最终值;每条更新本身保持原子性。本功能不引入暂停采集或跨节点同步流程。
|
||||
|
||||
## 前端设计
|
||||
|
||||
### API 封装
|
||||
|
||||
在 `vite-frontend/src/api/index.ts` 新增:
|
||||
|
||||
```ts
|
||||
export const resetForwardFlow = (id: number) =>
|
||||
Network.post("/forward/reset-flow", { id });
|
||||
```
|
||||
|
||||
### 入口
|
||||
|
||||
在规则页面所有单条规则操作入口中增加“流量清零”操作:
|
||||
|
||||
- 分组表格视图。
|
||||
- 精简表格视图。
|
||||
- 卡片视图。
|
||||
|
||||
按钮使用独立的清零/刷新语义图标和提示文本,不复用删除按钮样式。
|
||||
|
||||
当规则的 `inFlow + outFlow` 等于零时,按钮禁用,避免重复请求。
|
||||
|
||||
### 确认交互
|
||||
|
||||
点击按钮后打开确认弹窗,显示规则名称,并明确说明:
|
||||
|
||||
- 仅清零当前规则显示的上传和下载流量。
|
||||
- 不影响用户总流量、用户隧道配额和历史统计。
|
||||
- 操作不可撤销。
|
||||
|
||||
确认期间显示 loading 状态并阻止重复提交。
|
||||
|
||||
### 成功与失败
|
||||
|
||||
- 成功:关闭弹窗,显示成功 toast,并刷新规则列表。
|
||||
- 失败:保留弹窗,显示后端错误信息或通用失败 toast。
|
||||
- 刷新后,该规则上传和下载均显示为零;后续流量继续正常累计。
|
||||
|
||||
## 错误处理
|
||||
|
||||
- 非 POST 请求:返回现有通用请求失败响应。
|
||||
- 请求体无法解析、ID 缺失或 ID 非正数:返回“请求参数错误”。
|
||||
- 规则不存在或普通用户访问他人规则:返回“转发不存在”。
|
||||
- Repository 更新失败:返回包含 Repository 错误信息的统一错误响应。
|
||||
- 前端网络错误:显示“流量清零失败”。
|
||||
|
||||
## 测试策略
|
||||
|
||||
### Repository 测试
|
||||
|
||||
验证:
|
||||
|
||||
- 指定规则的 `in_flow`、`out_flow` 被清零。
|
||||
- 指定规则的 `updated_time` 被更新。
|
||||
- 其他规则的流量不变。
|
||||
- 用户总流量不变。
|
||||
- 用户隧道流量不变。
|
||||
- Repository 未初始化时返回错误。
|
||||
|
||||
### Handler 测试
|
||||
|
||||
验证:
|
||||
|
||||
- 管理员能够清零任意存在的规则。
|
||||
- 普通用户能够清零自己的规则。
|
||||
- 普通用户不能清零他人的规则。
|
||||
- 不存在的规则返回错误。
|
||||
- 无效 ID 返回参数错误。
|
||||
- 非 POST 请求返回请求失败。
|
||||
- 成功请求不修改用户和用户隧道流量。
|
||||
|
||||
### 前端验证
|
||||
|
||||
项目没有配置前端测试框架,因此不新增前端单元测试。使用以下命令验证:
|
||||
|
||||
```bash
|
||||
(cd vite-frontend && pnpm run build)
|
||||
(cd vite-frontend && pnpm run lint)
|
||||
```
|
||||
|
||||
后端使用:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
```
|
||||
|
||||
## 文件范围
|
||||
|
||||
预计修改:
|
||||
|
||||
- `go-backend/internal/http/handler/handler.go`
|
||||
- `go-backend/internal/http/handler/mutations.go`
|
||||
- `go-backend/internal/http/handler/*_test.go`
|
||||
- `go-backend/internal/store/repo/repository_mutations.go`
|
||||
- `go-backend/internal/store/repo/*_test.go`
|
||||
- `vite-frontend/src/api/index.ts`
|
||||
- `vite-frontend/src/pages/forward.tsx`
|
||||
|
||||
不需要数据库迁移或新增依赖。
|
||||
@@ -7,12 +7,17 @@ RUN go mod download
|
||||
COPY . .
|
||||
ARG TARGETOS
|
||||
ARG TARGETARCH
|
||||
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -o /out/paneld ./cmd/paneld
|
||||
ARG KEYGEN_ACCOUNT_ID
|
||||
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -ldflags="-X 'go-backend/internal/license.AccountID=${KEYGEN_ACCOUNT_ID}'" -o /out/paneld ./cmd/paneld
|
||||
|
||||
FROM docker:27-cli AS dockercli
|
||||
|
||||
FROM debian:bookworm-slim
|
||||
WORKDIR /app
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates wget && rm -rf /var/lib/apt/lists/*
|
||||
COPY --from=builder /out/paneld /app/paneld
|
||||
COPY --from=dockercli /usr/local/bin/docker /usr/local/bin/docker
|
||||
COPY --from=dockercli /usr/local/libexec/docker/cli-plugins/docker-compose /usr/local/libexec/docker/cli-plugins/docker-compose
|
||||
|
||||
ENV SERVER_ADDR=:6365
|
||||
EXPOSE 6365
|
||||
|
||||
@@ -90,7 +90,7 @@
|
||||
| 表名 | Model | 特殊处理 |
|
||||
|------|-------|----------|
|
||||
| `user` | `User` | `TableName()` 返回 `"user"` (PG 保留字) |
|
||||
| `forward` | `Forward` | |
|
||||
| `forward` | `Forward` | 增加 `proxy_protocol` 字段 |
|
||||
| `forward_port` | `ForwardPort` | |
|
||||
| `node` | `Node` | |
|
||||
| `speed_limit` | `SpeedLimit` | |
|
||||
|
||||
+2
-2
@@ -6,7 +6,8 @@ require (
|
||||
github.com/glebarez/sqlite v1.11.0
|
||||
github.com/google/uuid v1.6.0
|
||||
github.com/gorilla/websocket v1.5.3
|
||||
github.com/jackc/pgx/v5 v5.7.3
|
||||
github.com/jackc/pgx/v5 v5.9.2
|
||||
golang.org/x/crypto v0.31.0
|
||||
gorm.io/driver/postgres v1.6.0
|
||||
gorm.io/gorm v1.31.1
|
||||
)
|
||||
@@ -22,7 +23,6 @@ require (
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/ncruces/go-strftime v0.1.9 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
golang.org/x/crypto v0.50.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect
|
||||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.43.0 // indirect
|
||||
|
||||
+6
-6
@@ -17,8 +17,8 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
|
||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||
github.com/jackc/pgx/v5 v5.7.3 h1:PO1wNKj/bTAwxSJnO1Z4Ai8j4magtqg2SLNjEDzcXQo=
|
||||
github.com/jackc/pgx/v5 v5.7.3/go.mod h1:ncY89UGWxg82EykZUwSpUKEfccBGGYq1xjrOpsbsfGQ=
|
||||
github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
|
||||
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||
github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E=
|
||||
@@ -36,10 +36,10 @@ github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qq
|
||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||
github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk=
|
||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
||||
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
||||
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
golang.org/x/crypto v0.31.0 h1:ihbySMvVjLAeSH1IbfcRTkD/iNscyz8rGzjF/E5hV6U=
|
||||
golang.org/x/crypto v0.31.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM=
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8=
|
||||
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
|
||||
|
||||
@@ -12,12 +12,13 @@ import (
|
||||
|
||||
const (
|
||||
algorithm = "HmacSHA256"
|
||||
expireTime = 90 * 24 * time.Hour
|
||||
expireTime = 7 * 24 * time.Hour
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
Sub string `json:"sub"`
|
||||
Iat int64 `json:"iat"`
|
||||
IatMs int64 `json:"iat_ms"`
|
||||
Exp int64 `json:"exp"`
|
||||
User string `json:"user"`
|
||||
Name string `json:"name"`
|
||||
@@ -30,11 +31,15 @@ type tokenHeader struct {
|
||||
}
|
||||
|
||||
func GenerateToken(userID int64, username string, roleID int, secret string) (string, error) {
|
||||
now := time.Now()
|
||||
return GenerateTokenAt(userID, username, roleID, secret, time.Now())
|
||||
}
|
||||
|
||||
func GenerateTokenAt(userID int64, username string, roleID int, secret string, now time.Time) (string, error) {
|
||||
header := tokenHeader{Alg: algorithm, Typ: "JWT"}
|
||||
claims := Claims{
|
||||
Sub: strconv.FormatInt(userID, 10),
|
||||
Iat: now.Unix(),
|
||||
IatMs: now.UnixMilli(),
|
||||
Exp: now.Add(expireTime).Unix(),
|
||||
User: username,
|
||||
Name: username,
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
package auth
|
||||
|
||||
type UserAuthState struct {
|
||||
ID int64
|
||||
RoleID int
|
||||
Status int
|
||||
PasswordChangedAt int64
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req nameRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("配置名称不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
configName := strings.ToLower(strings.TrimSpace(req.Name))
|
||||
if configName == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("配置名称不能为空"))
|
||||
return
|
||||
}
|
||||
if !repo.IsPublicConfigKey(configName) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
if value, gated := h.unlicensedPublicBrandValue(configName); gated {
|
||||
response.WriteJSON(w, response.OK(map[string]string{"name": configName, "value": value}))
|
||||
return
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(configName)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if cfg == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("配置不存在"))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(cfg))
|
||||
}
|
||||
|
||||
func (h *Handler) unlicensedPublicBrandValue(configName string) (string, bool) {
|
||||
switch configName {
|
||||
case "app_name", "app_logo", "app_favicon", "hide_footer_brand":
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
|
||||
if isCommercial == "true" {
|
||||
return "", false
|
||||
}
|
||||
if configName == "app_name" {
|
||||
return "FLVX", true
|
||||
}
|
||||
if configName == "hide_footer_brand" {
|
||||
return "false", true
|
||||
}
|
||||
return "", true
|
||||
}
|
||||
@@ -0,0 +1,363 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
seedConfigValue(t, r, "app_name", "FLVX Brand")
|
||||
seedConfigValue(t, r, "app_logo", "logo-data")
|
||||
seedConfigValue(t, r, "app_favicon", "favicon-data")
|
||||
seedConfigValue(t, r, "app_bg_image", "bg-data")
|
||||
seedConfigValue(t, r, "app_bg_image_light", "light-bg-data")
|
||||
seedConfigValue(t, r, "app_bg_image_dark", "dark-bg-data")
|
||||
seedConfigValue(t, r, "cloudflare_site_key", "site-key")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
for name, want := range map[string]string{
|
||||
"app_bg_image_light": "light-bg-data",
|
||||
"app_bg_image_dark": "dark-bg-data",
|
||||
} {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerConfigValue(t, resp, name, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicBrandConfigFallsBackWithoutCommercialLicense(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
seedConfigValue(t, r, "app_name", "Paid Brand")
|
||||
seedConfigValue(t, r, "app_logo", "logo-data")
|
||||
seedConfigValue(t, r, "app_favicon", "favicon-data")
|
||||
seedConfigValue(t, r, "hide_footer_brand", "true")
|
||||
seedConfigValue(t, r, "is_commercial", "false")
|
||||
|
||||
for name, want := range map[string]string{
|
||||
"app_name": "FLVX",
|
||||
"app_logo": "",
|
||||
"app_favicon": "",
|
||||
"hide_footer_brand": "false",
|
||||
} {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerConfigValue(t, resp, name, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicBrandConfigUsesSavedValuesWithCommercialLicense(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
seedConfigValue(t, r, "app_name", "Paid Brand")
|
||||
seedConfigValue(t, r, "is_commercial", "true")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerConfigValue(t, resp, "app_name", "Paid Brand")
|
||||
}
|
||||
|
||||
func TestPublicConfigGetRejectsSensitiveKeys(t *testing.T) {
|
||||
router, _ := setupConfigAccessTestRouter(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"jwt_secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCodeMsg(t, resp, 403, "禁止访问敏感配置")
|
||||
}
|
||||
|
||||
func TestConfigGetAllowsPublicCloudflareSiteKeyWithoutAuthForCachedLoginPage(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
seedConfigValue(t, r, "cloudflare_site_key", "site-key")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"cloudflare_site_key"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerConfigValue(t, resp, "cloudflare_site_key", "site-key")
|
||||
}
|
||||
|
||||
func TestConfigGetRejectsSensitiveKeysWithoutAuth(t *testing.T) {
|
||||
router, _ := setupConfigAccessTestRouter(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"jwt_secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCodeMsg(t, resp, 403, "禁止访问敏感配置")
|
||||
}
|
||||
|
||||
func TestConfigGetAllowsSensitiveKeysForAdmin(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
seedConfigValue(t, r, "jwt_secret", "jwt-secret")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"jwt_secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerConfigValue(t, resp, "jwt_secret", "jwt-secret")
|
||||
}
|
||||
|
||||
func TestConfigGetNeverReturnsLicenseCredentials(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
seedConfigValue(t, r, "license_key", "license-secret")
|
||||
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
|
||||
|
||||
for _, name := range []string{"license_key", "license_machine_id", "machine_fingerprint"} {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerCodeMsg(t, resp, 403, "禁止访问系统授权凭据")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListNeverReturnsLicenseCredentials(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
seedConfigValue(t, r, "license_key", "license-secret")
|
||||
seedConfigValue(t, r, "license_machine_id", "machine-id")
|
||||
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
|
||||
seedConfigValue(t, r, "is_commercial", "true")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/list", nil)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
var out struct {
|
||||
Code int `json:"code"`
|
||||
Data map[string]string `json:"data"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 || out.Data["is_commercial"] != "true" {
|
||||
t.Fatalf("unexpected config response: %+v", out)
|
||||
}
|
||||
if _, ok := out.Data["license_key"]; ok {
|
||||
t.Fatal("license_key must not be returned")
|
||||
}
|
||||
if _, ok := out.Data["license_machine_id"]; ok {
|
||||
t.Fatal("license_machine_id must not be returned")
|
||||
}
|
||||
if _, ok := out.Data["machine_fingerprint"]; ok {
|
||||
t.Fatal("machine_fingerprint must not be returned")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateAllowsSensitiveKeysForAdmin(t *testing.T) {
|
||||
router, _ := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", bytes.NewBufferString(`{"jwt_secret":"rotated-secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
}
|
||||
|
||||
func TestConfigUpdateSingleAllowsSensitiveKeysForAdmin(t *testing.T) {
|
||||
router, _ := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update-single", bytes.NewBufferString(`{"name":"jwt_secret","value":"rotated-secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
}
|
||||
|
||||
func TestConfigUpdateAllowsCloudflareSecretKeyWrite(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", bytes.NewBufferString(`{"cloudflare_secret_key":"turnstile-secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
|
||||
cfg, err := r.GetConfigByName("cloudflare_secret_key")
|
||||
if err != nil {
|
||||
t.Fatalf("get config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "turnstile-secret" {
|
||||
t.Fatalf("expected cloudflare_secret_key to be updated, got %#v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateSingleAllowsCloudflareSecretKeyWrite(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update-single", bytes.NewBufferString(`{"name":"cloudflare_secret_key","value":"turnstile-secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
|
||||
cfg, err := r.GetConfigByName("cloudflare_secret_key")
|
||||
if err != nil {
|
||||
t.Fatalf("get config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "turnstile-secret" {
|
||||
t.Fatalf("expected cloudflare_secret_key to be updated, got %#v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateRejectsLicenseKeyWrite(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", bytes.NewBufferString(`{"license_key":"license-secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
|
||||
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
|
||||
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateSingleRejectsLicenseKeyWrite(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update-single", bytes.NewBufferString(`{"name":"license_key","value":"license-secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
|
||||
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
|
||||
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
|
||||
}
|
||||
}
|
||||
|
||||
func setupConfigAccessTestRouter(t *testing.T) (http.Handler, *repo.Repository) {
|
||||
t.Helper()
|
||||
r, err := repo.Open(t.TempDir() + "/config-access.db")
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = r.Close()
|
||||
})
|
||||
|
||||
h := New(r, "unit-test-secret")
|
||||
mux := http.NewServeMux()
|
||||
h.Register(mux)
|
||||
wrapped := middleware.Recover(mux)
|
||||
wrapped = middleware.JWT(middleware.AuthOptions{JWTSecret: "unit-test-secret", GetUserAuthState: h.GetUserAuthState})(wrapped)
|
||||
wrapped = middleware.RequestLog(wrapped)
|
||||
wrapped = middleware.CORS(wrapped)
|
||||
return wrapped, r
|
||||
}
|
||||
|
||||
func seedConfigValue(t *testing.T, r *repo.Repository, name, value string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`INSERT INTO vite_config(name, value, time) VALUES(?, ?, 0) ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time`, name, value).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustGenerateConfigAccessToken(t *testing.T, userID int64, username string, roleID int) string {
|
||||
t.Helper()
|
||||
token, err := auth.GenerateToken(userID, username, roleID, "unit-test-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
return token
|
||||
}
|
||||
|
||||
func assertHandlerCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expected {
|
||||
t.Fatalf("expected code %d, got %d", expected, out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func assertHandlerCodeMsg(t *testing.T, rec *httptest.ResponseRecorder, expectedCode int, expectedMsg string) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expectedCode || out.Msg != expectedMsg {
|
||||
t.Fatalf("expected (%d,%q), got (%d,%q)", expectedCode, expectedMsg, out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func assertHandlerConfigValue(t *testing.T, rec *httptest.ResponseRecorder, expectedName, expectedValue string) {
|
||||
t.Helper()
|
||||
var out struct {
|
||||
Code int `json:"code"`
|
||||
Data struct {
|
||||
Name string `json:"name"`
|
||||
Value string `json:"value"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d", out.Code)
|
||||
}
|
||||
if out.Data.Name != expectedName || out.Data.Value != expectedValue {
|
||||
t.Fatalf("expected config (%q,%q), got (%q,%q)", expectedName, expectedValue, out.Data.Name, out.Data.Value)
|
||||
}
|
||||
}
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/client"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/ws"
|
||||
)
|
||||
@@ -27,6 +28,16 @@ type nodeRecord = model.NodeRecord
|
||||
|
||||
type chainNodeRecord = model.ChainNodeRecord
|
||||
|
||||
type forwardRuntimeLimiters struct {
|
||||
TrafficLimiter string
|
||||
ConnLimiter string
|
||||
}
|
||||
|
||||
type forwardLimiterConfig struct {
|
||||
Name string
|
||||
Limits []string
|
||||
}
|
||||
|
||||
type diagnosisTarget struct {
|
||||
Address string
|
||||
IP string
|
||||
@@ -48,6 +59,7 @@ type diagnosisWorkItem struct {
|
||||
type diagnosisExecOptions struct {
|
||||
commandTimeout time.Duration
|
||||
pingTimeoutMS int
|
||||
pingCount int
|
||||
timeoutMessage string
|
||||
}
|
||||
|
||||
@@ -230,6 +242,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
|
||||
@@ -264,6 +289,13 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
speed = utSpeed
|
||||
}
|
||||
|
||||
var ipSpeed *int
|
||||
if forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0 {
|
||||
if speedVal, err := h.repo.GetSpeedLimitSpeed(forward.IPSpeedID.Int64); err == nil && speedVal > 0 {
|
||||
ipSpeed = &speedVal
|
||||
}
|
||||
}
|
||||
|
||||
serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID)
|
||||
|
||||
user, err := h.repo.GetUserByID(forward.UserID)
|
||||
@@ -271,19 +303,17 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var cLimiterName string
|
||||
var maxConnToSet int
|
||||
|
||||
if forward.MaxConn > 0 {
|
||||
maxConnToSet = forward.MaxConn
|
||||
cLimiterName = fmt.Sprintf("rule_conn_limit_%d", forward.ID)
|
||||
} else if user != nil && user.MaxConn > 0 {
|
||||
maxConnToSet = user.MaxConn
|
||||
cLimiterName = fmt.Sprintf("user_conn_limit_%d", user.ID)
|
||||
userMaxConn := 0
|
||||
if user != nil && user.MaxConn > 0 {
|
||||
userMaxConn = user.MaxConn
|
||||
}
|
||||
connLimiterConfigs := buildConnLimiterConfigs(forward, userMaxConn)
|
||||
|
||||
for _, fp := range ports {
|
||||
runtimeLimiters := forwardRuntimeLimiters{ConnLimiter: joinLimiterNames(connLimiterConfigs)}
|
||||
trafficLimiterNames := make([]string, 0, 2)
|
||||
if limiterID != nil && speed != nil {
|
||||
totalLimiterName := strconv.FormatInt(*limiterID, 10)
|
||||
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
|
||||
// If the limiter push fails because the node is offline, skip it with a warning
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
@@ -297,10 +327,29 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
trafficLimiterNames = append(trafficLimiterNames, totalLimiterName)
|
||||
}
|
||||
if ipSpeed != nil {
|
||||
ruleLimiterName := fmt.Sprintf("rule_traffic_limit_%d", forward.ID)
|
||||
if err := h.ensureTrafficLimiterOnNode(fp.NodeID, ruleLimiterName, nil, ipSpeed); err != nil {
|
||||
// If the limiter push fails because the node is offline, skip it with a warning
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
node, _ := h.getNodeRecord(fp.NodeID)
|
||||
nodeName := fmt.Sprintf("%d", fp.NodeID)
|
||||
if node != nil && strings.TrimSpace(node.Name) != "" {
|
||||
nodeName = strings.TrimSpace(node.Name)
|
||||
}
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName))
|
||||
continue
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
trafficLimiterNames = append(trafficLimiterNames, ruleLimiterName)
|
||||
}
|
||||
runtimeLimiters.TrafficLimiter = strings.Join(trafficLimiterNames, ",")
|
||||
|
||||
if cLimiterName != "" {
|
||||
if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil {
|
||||
for _, connLimiterConfig := range connLimiterConfigs {
|
||||
if err := h.ensureConnLimiterOnNode(fp.NodeID, connLimiterConfig); err != nil {
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
|
||||
}
|
||||
}
|
||||
@@ -309,7 +358,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, cLimiterName)
|
||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), runtimeLimiters)
|
||||
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
|
||||
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
||||
if isNotFoundError(err) {
|
||||
@@ -324,7 +373,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
}
|
||||
if err != nil && strings.EqualFold(strings.TrimSpace(method), "UpdateService") && isCannotAssignRequestedAddressError(err) {
|
||||
var warning string
|
||||
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID, cLimiterName)
|
||||
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, runtimeLimiters)
|
||||
if err == nil && warning != "" {
|
||||
warnings = append(warnings, warning)
|
||||
}
|
||||
@@ -350,7 +399,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
return warnings, nil
|
||||
}
|
||||
|
||||
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64, cLimiterName string) (string, error) {
|
||||
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, runtimeLimiters forwardRuntimeLimiters) (string, error) {
|
||||
if h == nil || forward == nil || tunnel == nil || node == nil {
|
||||
return "", errors.New("invalid bind fallback context")
|
||||
}
|
||||
@@ -367,7 +416,7 @@ func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunne
|
||||
}
|
||||
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, cLimiterName)
|
||||
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", runtimeLimiters)
|
||||
if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -450,13 +499,39 @@ func (h *Handler) forwardServiceBaseCandidates(forward *forwardRecord) ([]string
|
||||
}
|
||||
|
||||
func (h *Handler) deleteForwardServiceBasesOnNode(nodeID int64, bases []string) error {
|
||||
return deleteForwardServiceCandidates(bases, func(name string) error {
|
||||
payload := map[string]interface{}{
|
||||
"services": []string{name},
|
||||
names := buildForwardServiceDeleteNames(bases)
|
||||
if len(names) == 0 {
|
||||
return nil
|
||||
}
|
||||
payload := map[string]interface{}{"services": names}
|
||||
_, err := h.sendNodeCommand(nodeID, "DeleteService", payload, false, true)
|
||||
return err
|
||||
}
|
||||
|
||||
func buildForwardServiceDeleteNames(bases []string) []string {
|
||||
names := make([]string, 0, len(bases)*3)
|
||||
seen := make(map[string]struct{}, len(bases)*3)
|
||||
appendName := func(name string) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return
|
||||
}
|
||||
_, err := h.sendNodeCommand(nodeID, "DeleteService", payload, false, false)
|
||||
return err
|
||||
})
|
||||
if _, ok := seen[name]; ok {
|
||||
return
|
||||
}
|
||||
seen[name] = struct{}{}
|
||||
names = append(names, name)
|
||||
}
|
||||
for _, base := range bases {
|
||||
base = strings.TrimSpace(base)
|
||||
if base == "" {
|
||||
continue
|
||||
}
|
||||
appendName(base + "_tcp")
|
||||
appendName(base + "_udp")
|
||||
appendName(base)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
func (h *Handler) controlForwardServices(forward *forwardRecord, commandType string, tolerateNotFound bool) error {
|
||||
@@ -486,6 +561,25 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, userTunnelIDs...)
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, allUserTunnelIDs...)
|
||||
bases := buildForwardServiceBaseCandidates(forward.ID, forward.UserID, userTunnelID, candidateTunnelIDs)
|
||||
if strings.EqualFold(strings.TrimSpace(commandType), "DeleteService") {
|
||||
seen := map[int64]struct{}{}
|
||||
for _, fp := range ports {
|
||||
if _, ok := seen[fp.NodeID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[fp.NodeID] = struct{}{}
|
||||
if err := h.deleteForwardServiceBasesOnNode(fp.NodeID, bases); err != nil {
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
continue
|
||||
}
|
||||
if tolerateNotFound && isNotFoundError(err) {
|
||||
continue
|
||||
}
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
seen := map[int64]struct{}{}
|
||||
healed := false
|
||||
for _, fp := range ports {
|
||||
@@ -680,6 +774,9 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
if payload, handled, err := h.diagnoseNftablesForwardRuntime(forward); handled || err != nil {
|
||||
return payload, err
|
||||
}
|
||||
forwardName, workItems, err := h.prepareForwardDiagnosis(forward)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -695,6 +792,111 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func (h *Handler) diagnoseNftablesForwardRuntime(forward *forwardRecord) (map[string]interface{}, bool, error) {
|
||||
if forward == nil {
|
||||
return nil, false, errForwardNotFound
|
||||
}
|
||||
nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
if !nftMode {
|
||||
return nil, false, nil
|
||||
}
|
||||
if len(entryNodeIDs) == 0 {
|
||||
return nil, true, errors.New("nftables 转发缺少入口节点")
|
||||
}
|
||||
targets, err := resolveDiagnosisTargets(forward.RemoteAddr)
|
||||
if err != nil {
|
||||
return nil, true, err
|
||||
}
|
||||
results, err := h.buildNftablesForwardDiagnosisResults(forward, entryNodeIDs[0], targets)
|
||||
if err != nil {
|
||||
return nil, true, err
|
||||
}
|
||||
payload := map[string]interface{}{
|
||||
"forwardName": forward.Name,
|
||||
"timestamp": time.Now().UnixMilli(),
|
||||
"results": results,
|
||||
}
|
||||
return payload, true, nil
|
||||
}
|
||||
|
||||
func (h *Handler) buildNftablesForwardDiagnosisResults(forward *forwardRecord, nodeID int64, targets []diagnosisTarget) ([]map[string]interface{}, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return nil, errors.New("handler not initialized")
|
||||
}
|
||||
node, err := h.getNodeRecord(nodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var binding *model.NftRuleBinding
|
||||
for i := range bindings {
|
||||
if bindings[i].ForwardID == forward.ID {
|
||||
binding = &bindings[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
target := diagnosisTarget{}
|
||||
if len(targets) > 0 {
|
||||
target = targets[0]
|
||||
}
|
||||
|
||||
status := "missing"
|
||||
message := "nftables 规则未下发"
|
||||
success := false
|
||||
inPort := 0
|
||||
protocols := ""
|
||||
targetAddr := strings.TrimSpace(forward.RemoteAddr)
|
||||
ruleHash := ""
|
||||
if binding != nil {
|
||||
status = strings.ToLower(strings.TrimSpace(binding.Status))
|
||||
inPort = binding.InPort
|
||||
protocols = strings.TrimSpace(binding.Protocols)
|
||||
targetAddr = strings.TrimSpace(binding.TargetAddr)
|
||||
ruleHash = strings.TrimSpace(binding.RuleHash)
|
||||
if status == "" {
|
||||
status = "pending"
|
||||
}
|
||||
if status == runtimenft.StatusApplied {
|
||||
success = true
|
||||
message = "nftables 规则已下发"
|
||||
} else if strings.TrimSpace(binding.LastError) != "" {
|
||||
message = binding.LastError
|
||||
} else {
|
||||
message = "nftables 规则未完成下发"
|
||||
}
|
||||
}
|
||||
|
||||
packetLoss := 100
|
||||
if success {
|
||||
packetLoss = 0
|
||||
}
|
||||
result := map[string]interface{}{
|
||||
"success": success,
|
||||
"nodeName": node.Name,
|
||||
"nodeId": strconv.FormatInt(nodeID, 10),
|
||||
"targetIp": target.IP,
|
||||
"targetPort": target.Port,
|
||||
"description": fmt.Sprintf("nftables规则(%s)->目标(%s)", node.Name, defaultString(target.Address, targetAddr)),
|
||||
"averageTime": 0,
|
||||
"packetLoss": packetLoss,
|
||||
"message": message,
|
||||
"fromChainType": 1,
|
||||
"forwardMode": "nftables",
|
||||
"nftRuleStatus": status,
|
||||
"nftRuleHash": ruleHash,
|
||||
"inPort": inPort,
|
||||
"protocols": protocols,
|
||||
"targetAddr": targetAddr,
|
||||
}
|
||||
return []map[string]interface{}{result}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []diagnosisWorkItem, error) {
|
||||
if forward == nil {
|
||||
return "", nil, errForwardNotFound
|
||||
@@ -899,6 +1101,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
|
||||
ipPreference := h.repo.GetTunnelIPPreference(tunnelID)
|
||||
protocol := strings.ToLower(strings.TrimSpace(tunnel.Protocol))
|
||||
probeTarget := effectiveTunnelProbeTargetValues(tunnel.ProbeTargetHost, tunnel.ProbeTargetPort)
|
||||
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
|
||||
workItems := make([]diagnosisWorkItem, 0, len(chainRows)*2)
|
||||
|
||||
@@ -908,8 +1111,8 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
|
||||
workItems = append(workItems, diagnosisWorkItem{
|
||||
fromNodeID: inNode.NodeID,
|
||||
targetIP: "www.bing.com",
|
||||
targetPort: 443,
|
||||
targetIP: probeTarget.Host,
|
||||
targetPort: probeTarget.Port,
|
||||
description: description,
|
||||
protocol: "tcp",
|
||||
metadata: map[string]interface{}{
|
||||
@@ -1000,8 +1203,8 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
description := fmt.Sprintf("出口(%s)->外网", outNode.NodeName)
|
||||
workItems = append(workItems, diagnosisWorkItem{
|
||||
fromNodeID: outNode.NodeID,
|
||||
targetIP: "www.bing.com",
|
||||
targetPort: 443,
|
||||
targetIP: probeTarget.Host,
|
||||
targetPort: probeTarget.Port,
|
||||
description: description,
|
||||
protocol: "tcp",
|
||||
metadata: map[string]interface{}{
|
||||
@@ -1014,8 +1217,8 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
|
||||
workItems = append(workItems, diagnosisWorkItem{
|
||||
fromNodeID: inNode.NodeID,
|
||||
targetIP: "www.bing.com",
|
||||
targetPort: 443,
|
||||
targetIP: probeTarget.Host,
|
||||
targetPort: probeTarget.Port,
|
||||
description: description,
|
||||
protocol: "tcp",
|
||||
metadata: map[string]interface{}{
|
||||
@@ -1394,10 +1597,14 @@ func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int, options diag
|
||||
if options.pingTimeoutMS <= 0 {
|
||||
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
|
||||
}
|
||||
pingCount := options.pingCount
|
||||
if pingCount <= 0 {
|
||||
pingCount = 4
|
||||
}
|
||||
res, err := h.sendNodeCommandWithTimeout(nodeID, "TcpPing", map[string]interface{}{
|
||||
"ip": ip,
|
||||
"port": port,
|
||||
"count": 4,
|
||||
"count": pingCount,
|
||||
"timeout": options.pingTimeoutMS,
|
||||
}, options.commandTimeout, false, false)
|
||||
if err != nil {
|
||||
@@ -1424,12 +1631,16 @@ func (h *Handler) tcpPingViaRemoteNode(node *nodeRecord, ip string, port int, op
|
||||
if options.pingTimeoutMS <= 0 {
|
||||
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
|
||||
}
|
||||
pingCount := options.pingCount
|
||||
if pingCount <= 0 {
|
||||
pingCount = 4
|
||||
}
|
||||
|
||||
fc := client.NewFederationClientWithTimeout(options.commandTimeout)
|
||||
return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{
|
||||
IP: strings.TrimSpace(ip),
|
||||
Port: port,
|
||||
Count: 4,
|
||||
Count: pingCount,
|
||||
Timeout: options.pingTimeoutMS,
|
||||
Protocol: "tcp",
|
||||
})
|
||||
@@ -1659,11 +1870,12 @@ func compactErrorMessage(msg string) string {
|
||||
return strings.Join(strings.Fields(strings.ToLower(msg)), "")
|
||||
}
|
||||
|
||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, cLimiterName string) []map[string]interface{} {
|
||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, runtimeLimiters forwardRuntimeLimiters) []map[string]interface{} {
|
||||
protocols := []string{"tcp", "udp"}
|
||||
services := make([]map[string]interface{}, 0, 2)
|
||||
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||
strategy := strings.TrimSpace(forward.Strategy)
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(forward.ProxyProtocol, forward.ProxyProtocolReceive, forward.ProxyProtocolSend)
|
||||
if strategy == "" {
|
||||
strategy = "fifo"
|
||||
}
|
||||
@@ -1702,8 +1914,22 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
},
|
||||
},
|
||||
}
|
||||
if cLimiterName != "" {
|
||||
service["climiter"] = cLimiterName
|
||||
if runtimeLimiters.ConnLimiter != "" {
|
||||
service["climiter"] = runtimeLimiters.ConnLimiter
|
||||
}
|
||||
if runtimeLimiters.TrafficLimiter != "" {
|
||||
service["limiter"] = runtimeLimiters.TrafficLimiter
|
||||
}
|
||||
if proxyProtocolReceive > 0 {
|
||||
serviceMetadata := ensureServiceMetadata(service)
|
||||
serviceMetadata["proxyProtocol"] = proxyProtocolReceive
|
||||
}
|
||||
if proxyProtocolSend > 0 {
|
||||
handlerConfig := service["handler"].(map[string]interface{})
|
||||
if handlerConfig["metadata"] == nil {
|
||||
handlerConfig["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = proxyProtocolSend
|
||||
}
|
||||
if protocol == "udp" {
|
||||
listenerMetadata := map[string]interface{}{
|
||||
@@ -1716,10 +1942,8 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
|
||||
}
|
||||
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
|
||||
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
|
||||
}
|
||||
if limiterID != nil && *limiterID > 0 {
|
||||
service["limiter"] = strconv.FormatInt(*limiterID, 10)
|
||||
serviceMetadata := ensureServiceMetadata(service)
|
||||
serviceMetadata["interface"] = node.InterfaceName
|
||||
}
|
||||
services = append(services, service)
|
||||
}
|
||||
@@ -1738,6 +1962,25 @@ func buildForwarderNodes(targets []string) []map[string]interface{} {
|
||||
return nodes
|
||||
}
|
||||
|
||||
func ensureServiceMetadata(service map[string]interface{}) map[string]interface{} {
|
||||
if service["metadata"] == nil {
|
||||
service["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
metadata, ok := service["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
metadata = map[string]interface{}{}
|
||||
service["metadata"] = metadata
|
||||
}
|
||||
return metadata
|
||||
}
|
||||
|
||||
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
|
||||
if send == 0 && legacy > 0 {
|
||||
send = legacy
|
||||
}
|
||||
return receive, send
|
||||
}
|
||||
|
||||
func processServerAddress(serverAddr string) string {
|
||||
serverAddr = normalizeServerAddressInput(serverAddr)
|
||||
if serverAddr == "" {
|
||||
@@ -1820,22 +2063,16 @@ func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error {
|
||||
limitStr := fmt.Sprintf("$ %d", maxConn)
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"name": limiterName,
|
||||
"limits": []string{limitStr},
|
||||
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, cfg forwardLimiterConfig) error {
|
||||
if cfg.Name == "" || len(cfg.Limits) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{"name": cfg.Name, "limits": cfg.Limits}
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
|
||||
if !isAlreadyExistsMessage(err.Error()) {
|
||||
return fmt.Errorf("连接限制器下发失败: %w", err)
|
||||
}
|
||||
updatePayload := map[string]interface{}{
|
||||
"limiter": limiterName,
|
||||
"data": payload,
|
||||
}
|
||||
updatePayload := map[string]interface{}{"limiter": cfg.Name, "data": payload}
|
||||
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
|
||||
return fmt.Errorf("连接限制器更新失败: %w", updateErr)
|
||||
}
|
||||
@@ -1843,15 +2080,59 @@ func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxC
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildConnLimiterConfigs(forward *forwardRecord, userMaxConn int) []forwardLimiterConfig {
|
||||
if forward == nil {
|
||||
return nil
|
||||
}
|
||||
if forward.MaxConn > 0 {
|
||||
limits := []string{fmt.Sprintf("$ %d", forward.MaxConn)}
|
||||
if forward.IPMaxConn > 0 {
|
||||
limits = append(limits, fmt.Sprintf("$$ %d", forward.IPMaxConn))
|
||||
}
|
||||
return []forwardLimiterConfig{{Name: fmt.Sprintf("rule_conn_limit_%d", forward.ID), Limits: limits}}
|
||||
}
|
||||
configs := make([]forwardLimiterConfig, 0, 2)
|
||||
if userMaxConn > 0 {
|
||||
configs = append(configs, forwardLimiterConfig{Name: fmt.Sprintf("user_conn_limit_%d", forward.UserID), Limits: []string{fmt.Sprintf("$ %d", userMaxConn)}})
|
||||
}
|
||||
if forward.IPMaxConn > 0 {
|
||||
configs = append(configs, forwardLimiterConfig{Name: fmt.Sprintf("rule_conn_limit_%d", forward.ID), Limits: []string{fmt.Sprintf("$$ %d", forward.IPMaxConn)}})
|
||||
}
|
||||
return configs
|
||||
}
|
||||
|
||||
func joinLimiterNames(configs []forwardLimiterConfig) string {
|
||||
names := make([]string, 0, len(configs))
|
||||
for _, cfg := range configs {
|
||||
if cfg.Name != "" {
|
||||
names = append(names, cfg.Name)
|
||||
}
|
||||
}
|
||||
return strings.Join(names, ",")
|
||||
}
|
||||
|
||||
func speedToLimitLine(key string, speed int) string {
|
||||
rate := float64(speed) / 8.0
|
||||
return fmt.Sprintf("%s %.1fMB %.1fMB", key, rate, rate)
|
||||
}
|
||||
|
||||
func buildTrafficLimiterPayload(name string, totalSpeed *int, ipSpeed *int) map[string]interface{} {
|
||||
limits := make([]string, 0, 3)
|
||||
if totalSpeed != nil && *totalSpeed > 0 {
|
||||
limits = append(limits, speedToLimitLine("$", *totalSpeed))
|
||||
}
|
||||
if ipSpeed != nil && *ipSpeed > 0 {
|
||||
limits = append(limits, speedToLimitLine("0.0.0.0/0", *ipSpeed), speedToLimitLine("::/0", *ipSpeed))
|
||||
}
|
||||
return map[string]interface{}{"name": name, "limits": limits}
|
||||
}
|
||||
|
||||
func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) {
|
||||
rate := float64(speed) / 8.0
|
||||
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
|
||||
name := strconv.FormatInt(limiterID, 10)
|
||||
|
||||
return name, map[string]interface{}{
|
||||
"name": name,
|
||||
"limits": []string{limitStr},
|
||||
"limits": []string{speedToLimitLine("$", speed)},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1879,3 +2160,20 @@ func (h *Handler) upsertLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureTrafficLimiterOnNode(nodeID int64, name string, totalSpeed *int, ipSpeed *int) error {
|
||||
payload := buildTrafficLimiterPayload(name, totalSpeed, ipSpeed)
|
||||
limits, _ := payload["limits"].([]string)
|
||||
if name == "" || len(limits) == 0 {
|
||||
return nil
|
||||
}
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
|
||||
if !isAlreadyExistsMessage(err.Error()) {
|
||||
return fmt.Errorf("限速规则下发失败: %w", err)
|
||||
}
|
||||
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateLimiters", buildLimiterUpdatePayload(name, payload), false, false); updateErr != nil {
|
||||
return fmt.Errorf("限速规则更新失败: %w", updateErr)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -201,6 +201,52 @@ func TestDeleteForwardServiceCandidatesDeletesAllMatchingVariants(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceDeleteNamesBatchesAndDeduplicatesVariants(t *testing.T) {
|
||||
bases := []string{"57_7_7", "57_7_0", "57_7_7"}
|
||||
got := buildForwardServiceDeleteNames(bases)
|
||||
want := []string{"57_7_7_tcp", "57_7_7_udp", "57_7_7", "57_7_0_tcp", "57_7_0_udp", "57_7_0"}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("expected %v, got %v", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemovedTunnelRuntimeNodeIDsSeparatesChainAndServiceRoles(t *testing.T) {
|
||||
oldRows := []chainNodeRecord{
|
||||
{NodeID: 1, ChainType: 1},
|
||||
{NodeID: 2, ChainType: 2},
|
||||
{NodeID: 3, ChainType: 3},
|
||||
{NodeID: 5, ChainType: 2},
|
||||
{NodeID: 6, ChainType: 3},
|
||||
}
|
||||
newRows := []chainNodeRecord{
|
||||
{NodeID: 2, ChainType: 3},
|
||||
{NodeID: 3, ChainType: 3},
|
||||
{NodeID: 5, ChainType: 1},
|
||||
}
|
||||
|
||||
removedChains := removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsChain)
|
||||
if want := []int64{1, 2}; !reflect.DeepEqual(removedChains, want) {
|
||||
t.Fatalf("expected removed chains %v, got %v", want, removedChains)
|
||||
}
|
||||
|
||||
removedServices := removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsService)
|
||||
if want := []int64{5, 6}; !reflect.DeepEqual(removedServices, want) {
|
||||
t.Fatalf("expected removed services %v, got %v", want, removedServices)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelForwardRuntimeNeedsSyncOnlyWhenTypeOrEntriesChange(t *testing.T) {
|
||||
if tunnelForwardRuntimeNeedsSync(2, 2, []int64{1, 2}, []int64{2, 1}) {
|
||||
t.Fatalf("same tunnel type and same entry set should not resync forwards")
|
||||
}
|
||||
if !tunnelForwardRuntimeNeedsSync(1, 2, []int64{1}, []int64{1}) {
|
||||
t.Fatalf("type change should resync forwards")
|
||||
}
|
||||
if !tunnelForwardRuntimeNeedsSync(2, 2, []int64{1}, []int64{1, 2}) {
|
||||
t.Fatalf("entry set change should resync forwards")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateForwardPortAvailabilityRejectsOtherForwardOccupancy(t *testing.T) {
|
||||
h := &Handler{repo: nil}
|
||||
node := &nodeRecord{ID: 9, Name: "test-node"}
|
||||
@@ -378,7 +424,7 @@ func TestRetryTunnelServiceAddWithCleanupReturnsCleanupError(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -393,7 +439,7 @@ func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -409,7 +455,7 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.
|
||||
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -464,7 +510,7 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -478,6 +524,58 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildConnLimiterConfigCombinesTotalAndPerIP(t *testing.T) {
|
||||
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, MaxConn: 100, IPMaxConn: 5}, 37)
|
||||
want := []forwardLimiterConfig{{Name: "rule_conn_limit_42", Limits: []string{"$ 100", "$$ 5"}}}
|
||||
if !reflect.DeepEqual(cfgs, want) {
|
||||
t.Fatalf("expected %+v, got %+v", want, cfgs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildConnLimiterConfigUsesUserTotalWithRulePerIP(t *testing.T) {
|
||||
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, IPMaxConn: 5}, 37)
|
||||
want := []forwardLimiterConfig{
|
||||
{Name: "user_conn_limit_9", Limits: []string{"$ 37"}},
|
||||
{Name: "rule_conn_limit_42", Limits: []string{"$$ 5"}},
|
||||
}
|
||||
if !reflect.DeepEqual(cfgs, want) {
|
||||
t.Fatalf("expected %+v, got %+v", want, cfgs)
|
||||
}
|
||||
if got := joinLimiterNames(cfgs); got != "user_conn_limit_9,rule_conn_limit_42" {
|
||||
t.Fatalf("expected composite limiter names, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTrafficLimiterPayloadUsesOnlyPerIPRulesWhenTotalIsSeparate(t *testing.T) {
|
||||
payload := buildTrafficLimiterPayload("rule_traffic_limit_42", nil, intPtr(40))
|
||||
wantLimits := []string{"0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"}
|
||||
if payload["name"] != "rule_traffic_limit_42" {
|
||||
t.Fatalf("expected name rule_traffic_limit_42, got %v", payload["name"])
|
||||
}
|
||||
if !reflect.DeepEqual(payload["limits"], wantLimits) {
|
||||
t.Fatalf("expected limits %v, got %v", wantLimits, payload["limits"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigsUsesRuntimeLimiterNames(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{TrafficLimiter: "rule_traffic_limit_42", ConnLimiter: "rule_conn_limit_42"})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
for _, service := range services {
|
||||
if service["limiter"] != "rule_traffic_limit_42" {
|
||||
t.Fatalf("expected traffic limiter rule_traffic_limit_42, got %v", service["limiter"])
|
||||
}
|
||||
if service["climiter"] != "rule_conn_limit_42" {
|
||||
t.Fatalf("expected conn limiter rule_conn_limit_42, got %v", service["climiter"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func intPtr(v int) *int { return &v }
|
||||
|
||||
func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -85,6 +85,81 @@ type federationRuntimeReleaseRoleRequest struct {
|
||||
ResourceKey string `json:"resourceKey"`
|
||||
}
|
||||
|
||||
func federationRuntimeChainName(bindingID string) string {
|
||||
bindingID = strings.TrimSpace(bindingID)
|
||||
if bindingID == "" {
|
||||
return ""
|
||||
}
|
||||
return "fed_chain_" + bindingID
|
||||
}
|
||||
|
||||
func buildFederationMiddleChainConfig(chainName string, runtimeID int64, protocol, strategy string, targets []federationRuntimeTarget, interfaceName string) (map[string]interface{}, error) {
|
||||
chainName = strings.TrimSpace(chainName)
|
||||
if chainName == "" {
|
||||
return nil, fmt.Errorf("chain name is required")
|
||||
}
|
||||
if len(targets) == 0 {
|
||||
return nil, fmt.Errorf("targets are required for middle role")
|
||||
}
|
||||
protocol = defaultString(protocol, "tls")
|
||||
nodeItems := make([]map[string]interface{}, 0, len(targets))
|
||||
for i, target := range targets {
|
||||
host := strings.TrimSpace(target.Host)
|
||||
if host == "" || target.Port <= 0 {
|
||||
return nil, fmt.Errorf("Invalid target")
|
||||
}
|
||||
targetProtocol := defaultString(target.Protocol, protocol)
|
||||
connector := map[string]interface{}{
|
||||
"type": "relay",
|
||||
}
|
||||
if isTCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"nodelay": true,
|
||||
"mux.keepaliveInterval": "15s",
|
||||
"mux.keepaliveTimeout": "45s",
|
||||
}
|
||||
}
|
||||
if isKCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"connectTimeout": "30s",
|
||||
}
|
||||
}
|
||||
nodeItems = append(nodeItems, map[string]interface{}{
|
||||
"name": fmt.Sprintf("node_%d", i+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, target.Port)),
|
||||
"connector": connector,
|
||||
"dialer": buildTunnelDialerConfig(targetProtocol),
|
||||
})
|
||||
}
|
||||
|
||||
chainData := map[string]interface{}{
|
||||
"name": chainName,
|
||||
"hops": []map[string]interface{}{
|
||||
{
|
||||
"name": fmt.Sprintf("hop_%d", runtimeID),
|
||||
"selector": map[string]interface{}{
|
||||
"strategy": runtimeTunnelStrategy(strategy),
|
||||
"maxFails": 1,
|
||||
"failTimeout": int64(600000000000),
|
||||
},
|
||||
"nodes": nodeItems,
|
||||
},
|
||||
},
|
||||
}
|
||||
if strings.TrimSpace(interfaceName) != "" {
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
hops[0]["interface"] = interfaceName
|
||||
}
|
||||
return chainData, nil
|
||||
}
|
||||
|
||||
func updateChainPayload(chainName string, chainData map[string]interface{}) map[string]interface{} {
|
||||
return map[string]interface{}{
|
||||
"chain": chainName,
|
||||
"data": chainData,
|
||||
}
|
||||
}
|
||||
|
||||
type federationRuntimeDiagnoseRequest struct {
|
||||
IP string `json:"ip"`
|
||||
Port int `json:"port"`
|
||||
@@ -145,8 +220,8 @@ type remoteUsageNodeItem struct {
|
||||
|
||||
func buildFederationServiceConfig(serviceName, addr, protocol, role, chainName string, targetCount int, interfaceName string) map[string]interface{} {
|
||||
service := map[string]interface{}{
|
||||
"name": serviceName,
|
||||
"addr": addr,
|
||||
"name": serviceName,
|
||||
"addr": addr,
|
||||
"handler": map[string]interface{}{
|
||||
"type": "relay",
|
||||
},
|
||||
@@ -1056,7 +1131,43 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
return
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(share.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
protocol := defaultString(req.Protocol, runtime.Protocol)
|
||||
strategy := defaultString(req.Strategy, "round")
|
||||
chainName := defaultString(runtime.ChainName, federationRuntimeChainName(runtime.BindingID))
|
||||
if chainName == "" {
|
||||
chainName = federationRuntimeChainName(fmt.Sprintf("%d", runtime.ID))
|
||||
}
|
||||
serviceName := fmt.Sprintf("fed_svc_%d", runtime.ID)
|
||||
if runtime.Applied == 1 && strings.TrimSpace(runtime.BindingID) != "" {
|
||||
if req.Role == "middle" && len(req.Targets) > 0 {
|
||||
chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName)
|
||||
if buildErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(buildErr.Error()))
|
||||
return
|
||||
}
|
||||
if _, err := h.sendNodeCommand(share.NodeID, "UpdateChains", updateChainPayload(chainName, chainData), false, false); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
targetBytes, _ := json.Marshal(req.Targets)
|
||||
runtime.Role = req.Role
|
||||
runtime.ChainName = chainName
|
||||
runtime.Protocol = protocol
|
||||
runtime.Strategy = strategy
|
||||
runtime.Target = string(targetBytes)
|
||||
runtime.Status = 1
|
||||
runtime.UpdatedTime = time.Now().UnixMilli()
|
||||
if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"bindingId": runtime.BindingID,
|
||||
"allocatedPort": runtime.Port,
|
||||
@@ -1076,71 +1187,12 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
}
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(share.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
protocol := defaultString(req.Protocol, runtime.Protocol)
|
||||
strategy := defaultString(req.Strategy, "round")
|
||||
chainName := fmt.Sprintf("fed_chain_%d", runtime.ID)
|
||||
serviceName := fmt.Sprintf("fed_svc_%d", runtime.ID)
|
||||
|
||||
if req.Role == "middle" {
|
||||
if len(req.Targets) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("targets are required for middle role"))
|
||||
chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName)
|
||||
if buildErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(buildErr.Error()))
|
||||
return
|
||||
}
|
||||
nodeItems := make([]map[string]interface{}, 0, len(req.Targets))
|
||||
for i, target := range req.Targets {
|
||||
host := strings.TrimSpace(target.Host)
|
||||
if host == "" || target.Port <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid target"))
|
||||
return
|
||||
}
|
||||
targetProtocol := defaultString(target.Protocol, protocol)
|
||||
connector := map[string]interface{}{
|
||||
"type": "relay",
|
||||
}
|
||||
if isTCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"nodelay": true,
|
||||
"mux.keepaliveInterval": "15s",
|
||||
"mux.keepaliveTimeout": "45s",
|
||||
}
|
||||
}
|
||||
if isKCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"connectTimeout": "30s",
|
||||
}
|
||||
}
|
||||
nodeItems = append(nodeItems, map[string]interface{}{
|
||||
"name": fmt.Sprintf("node_%d", i+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, target.Port)),
|
||||
"connector": connector,
|
||||
"dialer": buildTunnelDialerConfig(targetProtocol),
|
||||
})
|
||||
}
|
||||
|
||||
chainData := map[string]interface{}{
|
||||
"name": chainName,
|
||||
"hops": []map[string]interface{}{
|
||||
{
|
||||
"name": fmt.Sprintf("hop_%d", runtime.ID),
|
||||
"selector": map[string]interface{}{
|
||||
"strategy": strategy,
|
||||
"maxFails": 1,
|
||||
"failTimeout": int64(600000000000),
|
||||
},
|
||||
"nodes": nodeItems,
|
||||
},
|
||||
},
|
||||
}
|
||||
if strings.TrimSpace(node.InterfaceName) != "" {
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
hops[0]["interface"] = node.InterfaceName
|
||||
}
|
||||
if _, err := h.sendNodeCommand(share.NodeID, "AddChains", chainData, true, false); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
|
||||
@@ -281,6 +281,63 @@ func TestBuildFederationServiceConfig_NonTLSProtocol_NoNodelay(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRuntimeChainNameDerivesFromBindingID(t *testing.T) {
|
||||
if got := federationRuntimeChainName("12"); got != "fed_chain_12" {
|
||||
t.Fatalf("expected fed_chain_12, got %q", got)
|
||||
}
|
||||
if got := federationRuntimeChainName(" 12 "); got != "fed_chain_12" {
|
||||
t.Fatalf("expected trimmed fed_chain_12, got %q", got)
|
||||
}
|
||||
if got := federationRuntimeChainName(""); got != "" {
|
||||
t.Fatalf("expected blank binding ID to stay blank, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFederationMiddleChainConfigUsesExistingChainNameAndBestStrategy(t *testing.T) {
|
||||
chainData, err := buildFederationMiddleChainConfig("fed_chain_12", 12, "tls", tunnelStrategyBest, []federationRuntimeTarget{
|
||||
{Host: "10.0.0.31", Port: 30031, Protocol: "tls"},
|
||||
{Host: "10.0.0.30", Port: 30030, Protocol: "tls"},
|
||||
}, "")
|
||||
if err != nil {
|
||||
t.Fatalf("build chain: %v", err)
|
||||
}
|
||||
if chainData["name"] != "fed_chain_12" {
|
||||
t.Fatalf("expected existing chain name, got %v", chainData["name"])
|
||||
}
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
selector := hops[0]["selector"].(map[string]interface{})
|
||||
if selector["strategy"] != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best strategy to map to fifo, got %v", selector["strategy"])
|
||||
}
|
||||
nodes := hops[0]["nodes"].([]map[string]interface{})
|
||||
if nodes[0]["addr"] != "10.0.0.31:30031" || nodes[1]["addr"] != "10.0.0.30:30030" {
|
||||
t.Fatalf("expected target order to be preserved, got %+v", nodes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateChainPayloadWrapsChainDataForAgentUpdate(t *testing.T) {
|
||||
chainData := map[string]interface{}{
|
||||
"name": "fed_chain_12",
|
||||
"hops": []map[string]interface{}{},
|
||||
}
|
||||
|
||||
payload := updateChainPayload("fed_chain_12", chainData)
|
||||
if len(payload) != 2 {
|
||||
t.Fatalf("expected exact wrapper with 2 keys, got %+v", payload)
|
||||
}
|
||||
if payload["chain"] != "fed_chain_12" {
|
||||
t.Fatalf("expected chain name in wrapper, got %v", payload["chain"])
|
||||
}
|
||||
wrappedData, ok := payload["data"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected wrapped chain data map, got %T", payload["data"])
|
||||
}
|
||||
chainData["name"] = "fed_chain_12_updated"
|
||||
if wrappedData["name"] != "fed_chain_12_updated" {
|
||||
t.Fatalf("expected wrapper to preserve chainData identity, got %+v", wrappedData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRuntimeReservePortRejectsWhenShareFlowExceeded(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log"
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -12,12 +13,27 @@ import (
|
||||
)
|
||||
|
||||
const bytesPerGB int64 = 1024 * 1024 * 1024
|
||||
const bytesPerMiB int64 = 1024 * 1024
|
||||
|
||||
func flowLimitBytes(flowGB, flowMiB int64) int64 {
|
||||
if flowMiB > 0 {
|
||||
if flowMiB > math.MaxInt64/bytesPerMiB {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return flowMiB * bytesPerMiB
|
||||
}
|
||||
if flowGB > math.MaxInt64/bytesPerGB {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return flowGB * bytesPerGB
|
||||
}
|
||||
|
||||
type userTunnelPolicy struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
TunnelID int64
|
||||
Flow int64
|
||||
FlowMiB int64
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
ExpTime int64
|
||||
@@ -358,7 +374,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
return errors.New("账号已过期")
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return errors.New("流量已超额,禁止开启转发")
|
||||
@@ -400,7 +416,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
return errors.New("该隧道已过期")
|
||||
}
|
||||
|
||||
utFlowLimit := policy.Flow * bytesPerGB
|
||||
utFlowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
|
||||
utCurrent := policy.InFlow + policy.OutFlow
|
||||
if utCurrent >= utFlowLimit {
|
||||
return errors.New("该隧道流量已超额,禁止开启转发")
|
||||
@@ -425,7 +441,7 @@ func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return true
|
||||
@@ -441,7 +457,7 @@ func shouldPauseUserTunnel(policy *userTunnelPolicy, now int64) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := policy.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
|
||||
current := policy.InFlow + policy.OutFlow
|
||||
if current >= flowLimit {
|
||||
return true
|
||||
@@ -465,7 +481,7 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er
|
||||
}
|
||||
return &userTunnelPolicy{
|
||||
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
|
||||
Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
ExpTime: ut.ExpTime, Status: ut.Status, Num: ut.Num,
|
||||
}, nil
|
||||
}
|
||||
@@ -659,9 +675,21 @@ func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName str
|
||||
}
|
||||
|
||||
func (h *Handler) speedLimiterExists(name string) bool {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
const forwardRulePrefix = "rule_traffic_limit_"
|
||||
if strings.HasPrefix(name, forwardRulePrefix) {
|
||||
forwardID, err := strconv.ParseInt(strings.TrimPrefix(name, forwardRulePrefix), 10, 64)
|
||||
if err != nil || forwardID <= 0 {
|
||||
return false
|
||||
}
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
return err == nil && forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0
|
||||
}
|
||||
|
||||
id, err := strconv.ParseInt(name, 10, 64)
|
||||
if err != nil || id <= 0 {
|
||||
return false
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestSpeedLimiterExistsPreservesForwardRuleLimiter(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
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, ip_speed_id)
|
||||
VALUES(8, 1, 'user', 'forward', 1, '127.0.0.1:80', 'fifo', 0, 0, 1, 1, 1, 0, 3),
|
||||
(9, 1, 'user', 'forward-without-ip-limit', 1, '127.0.0.1:81', 'fifo', 0, 0, 1, 1, 1, 0, NULL)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
if !h.speedLimiterExists("rule_traffic_limit_8") {
|
||||
t.Fatal("expected runtime limiter for existing forward to be preserved")
|
||||
}
|
||||
if h.speedLimiterExists("rule_traffic_limit_9") {
|
||||
t.Fatal("expected runtime limiter for forward without per-IP speed limit to be treated as orphaned")
|
||||
}
|
||||
if h.speedLimiterExists("rule_traffic_limit_10") {
|
||||
t.Fatal("expected runtime limiter for missing forward to be treated as orphaned")
|
||||
}
|
||||
if h.speedLimiterExists("rule_traffic_limit_invalid") {
|
||||
t.Fatal("expected malformed runtime limiter name to be treated as orphaned")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type flowPolicyTarget struct {
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
}
|
||||
|
||||
type flowUploadBatch struct {
|
||||
flowDeltas []repo.FlowUploadCounterDelta
|
||||
quotaUsage map[int64]int64
|
||||
policyTargets []flowPolicyTarget
|
||||
forwardTraffic map[int64]tunnelTrafficDelta
|
||||
orphanServices map[string]struct{}
|
||||
peerShareForwardItems map[string]flowItem
|
||||
peerShareRuntimeItems map[int64]flowItem
|
||||
}
|
||||
|
||||
func (h *Handler) buildFlowUploadBatch(items []flowItem, 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 _, item := range items {
|
||||
serviceName := strings.TrimSpace(item.N)
|
||||
if serviceName == "" || serviceName == "web_api" {
|
||||
continue
|
||||
}
|
||||
if runtimeID, ok := parsePeerShareRuntimeServiceID(serviceName); ok {
|
||||
merged := batch.peerShareRuntimeItems[runtimeID]
|
||||
merged.N = serviceName
|
||||
merged.U += item.U
|
||||
merged.D += item.D
|
||||
batch.peerShareRuntimeItems[runtimeID] = merged
|
||||
continue
|
||||
}
|
||||
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
normalized := normalizeForwardRuntimeServiceName(serviceName)
|
||||
merged := batch.peerShareForwardItems[normalized]
|
||||
merged.N = normalized
|
||||
merged.U += item.U
|
||||
merged.D += item.D
|
||||
batch.peerShareForwardItems[normalized] = merged
|
||||
|
||||
meta, exists := metas[forwardID]
|
||||
if !exists {
|
||||
batch.orphanServices[serviceName] = struct{}{}
|
||||
continue
|
||||
}
|
||||
|
||||
raw := batch.forwardTraffic[forwardID]
|
||||
raw.bytesIn += item.D
|
||||
raw.bytesOut += item.U
|
||||
batch.forwardTraffic[forwardID] = raw
|
||||
|
||||
scaledIn := int64(float64(item.D)*meta.TrafficRatio) * meta.TunnelFlow
|
||||
scaledOut := int64(float64(item.U)*meta.TrafficRatio) * meta.TunnelFlow
|
||||
if idx, ok := flowSeen[forwardID]; ok {
|
||||
batch.flowDeltas[idx].InFlow += scaledIn
|
||||
batch.flowDeltas[idx].OutFlow += scaledOut
|
||||
} else {
|
||||
flowSeen[forwardID] = len(batch.flowDeltas)
|
||||
batch.flowDeltas = append(batch.flowDeltas, repo.FlowUploadCounterDelta{
|
||||
ForwardID: forwardID,
|
||||
UserID: userID,
|
||||
UserTunnelID: userTunnelID,
|
||||
InFlow: scaledIn,
|
||||
OutFlow: scaledOut,
|
||||
})
|
||||
}
|
||||
batch.quotaUsage[userID] += scaledIn + scaledOut
|
||||
|
||||
target := flowPolicyTarget{UserID: userID, UserTunnelID: 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 (h *Handler) applyFlowUploadBatch(nodeID int64, batch flowUploadBatch, now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
h.applyFlowDeltasWithFallback(nodeID, batch.flowDeltas)
|
||||
for userID, quota := range h.applyQuotaUsageWithFallback(nodeID, batch.quotaUsage, now) {
|
||||
h.enforceUserQuotaIfNeeded(userID, quota)
|
||||
}
|
||||
for _, target := range batch.policyTargets {
|
||||
if target.UserID <= 0 || target.UserTunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
h.enforceFlowPolicies(target.UserID, target.UserTunnelID)
|
||||
}
|
||||
for serviceName := range batch.orphanServices {
|
||||
h.sendDeleteOrphanedForwardService(nodeID, serviceName)
|
||||
}
|
||||
for serviceName, item := range batch.peerShareForwardItems {
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(serviceName)
|
||||
if ok {
|
||||
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
|
||||
}
|
||||
}
|
||||
for runtimeID, item := range batch.peerShareRuntimeItems {
|
||||
h.processPeerShareFlow(runtimeID, item)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) applyFlowDeltasWithFallback(nodeID int64, deltas []repo.FlowUploadCounterDelta) {
|
||||
if h == nil || h.repo == nil || len(deltas) == 0 {
|
||||
return
|
||||
}
|
||||
if err := h.repo.ApplyFlowUploadDeltasBatch(deltas); err == nil {
|
||||
return
|
||||
} else {
|
||||
log.Printf("flow upload write failed op=flow.batch_apply node_id=%d err=%v", nodeID, err)
|
||||
}
|
||||
for _, delta := range deltas {
|
||||
if err := h.repo.AddFlow(delta.ForwardID, delta.UserID, delta.UserTunnelID, delta.InFlow, delta.OutFlow); err != nil {
|
||||
log.Printf("flow upload write failed op=flow.single_apply node_id=%d forward_id=%d user_id=%d user_tunnel_id=%d err=%v", nodeID, delta.ForwardID, delta.UserID, delta.UserTunnelID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) applyQuotaUsageWithFallback(nodeID int64, usages map[int64]int64, now time.Time) map[int64]*model.UserQuotaView {
|
||||
if h == nil || h.repo == nil || len(usages) == 0 {
|
||||
return map[int64]*model.UserQuotaView{}
|
||||
}
|
||||
quotaViews, err := h.repo.AddUserQuotaUsageBatch(usages, now)
|
||||
if err == nil {
|
||||
return quotaViews
|
||||
}
|
||||
log.Printf("flow upload write failed op=quota.batch_apply node_id=%d err=%v", nodeID, err)
|
||||
|
||||
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] })
|
||||
|
||||
quotaViews = make(map[int64]*model.UserQuotaView, len(userIDs))
|
||||
for _, userID := range userIDs {
|
||||
quota, singleErr := h.repo.AddUserQuotaUsage(userID, usages[userID], now)
|
||||
if singleErr != nil {
|
||||
log.Printf("flow upload write failed op=quota.single_apply node_id=%d user_id=%d err=%v", nodeID, userID, singleErr)
|
||||
continue
|
||||
}
|
||||
if quota != nil {
|
||||
quotaViews[userID] = quota
|
||||
}
|
||||
}
|
||||
return quotaViews
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets(t *testing.T) {
|
||||
h := &Handler{}
|
||||
metas := map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {
|
||||
ForwardID: 20,
|
||||
TunnelID: 1,
|
||||
TrafficRatio: 2,
|
||||
TunnelFlow: 3,
|
||||
},
|
||||
}
|
||||
|
||||
batch := h.buildFlowUploadBatch([]flowItem{
|
||||
{N: "20_2_10", U: 70, D: 50},
|
||||
{N: "20_2_10_tcp", U: 40, D: 30},
|
||||
{N: "99_2_10", U: 12, D: 8},
|
||||
{N: "fed_svc_17", U: 9, D: 1},
|
||||
}, metas)
|
||||
|
||||
if len(batch.flowDeltas) != 1 {
|
||||
t.Fatalf("expected 1 flow delta, got %d", len(batch.flowDeltas))
|
||||
}
|
||||
delta := batch.flowDeltas[0]
|
||||
if delta.ForwardID != 20 || delta.UserID != 2 || delta.UserTunnelID != 10 {
|
||||
t.Fatalf("unexpected flow delta identity: %#v", delta)
|
||||
}
|
||||
if delta.InFlow != 480 || delta.OutFlow != 660 {
|
||||
t.Fatalf("expected scaled flow in=480 out=660, got in=%d out=%d", delta.InFlow, delta.OutFlow)
|
||||
}
|
||||
if batch.quotaUsage[2] != 1140 {
|
||||
t.Fatalf("expected quota usage 1140, got %d", batch.quotaUsage[2])
|
||||
}
|
||||
if len(batch.policyTargets) != 1 {
|
||||
t.Fatalf("expected 1 policy target, got %d", len(batch.policyTargets))
|
||||
}
|
||||
if batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("unexpected policy target: %#v", batch.policyTargets[0])
|
||||
}
|
||||
traffic := batch.forwardTraffic[20]
|
||||
if traffic.bytesIn != 80 || traffic.bytesOut != 110 {
|
||||
t.Fatalf("expected raw traffic in=80 out=110, got in=%d out=%d", traffic.bytesIn, traffic.bytesOut)
|
||||
}
|
||||
if _, ok := batch.orphanServices["99_2_10"]; !ok {
|
||||
t.Fatalf("expected orphan service cleanup target for 99_2_10")
|
||||
}
|
||||
if item, ok := batch.peerShareForwardItems["99_2_10"]; !ok || item.U != 12 || item.D != 8 {
|
||||
t.Fatalf("expected orphan forward to remain eligible for peer-share accounting, got %#v ok=%v", item, ok)
|
||||
}
|
||||
if item, ok := batch.peerShareForwardItems["20_2_10"]; !ok || item.U != 110 || item.D != 80 {
|
||||
t.Fatalf("expected merged peer-share forward item, got %#v ok=%v", item, ok)
|
||||
}
|
||||
if item, ok := batch.peerShareRuntimeItems[17]; !ok || item.U != 9 || item.D != 1 {
|
||||
t.Fatalf("expected merged peer-share runtime item, got %#v ok=%v", item, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFlowUploadBatchContinuesPolicyAndPeerShareSideEffectsWhenQuotaBatchFails(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "flow-upload-batch-quota-fail.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
if err := r.DB().Create(&model.User{ID: 2, User: "flow-user", Pwd: "pwd", RoleID: 1, ExpTime: 2727251700000, Flow: 99999, Num: 99999, CreatedTime: nowMs, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed user: %v", err)
|
||||
}
|
||||
if err := r.DB().Create(&model.Tunnel{ID: 1, Name: "tunnel-1", TrafficRatio: 1, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Create(&model.UserTunnel{ID: 10, UserID: 2, TunnelID: 1, Num: 99999, Flow: 0, ExpTime: 2727251700000, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed user tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Create(&model.Forward{ID: 20, UserID: 2, UserName: "flow-user", Name: "forward-20", TunnelID: 1, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed forward: %v", err)
|
||||
}
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{Name: "share", NodeID: 1, Token: "token", MaxBandwidth: 0, CurrentFlow: 0, PortRangeStart: 31000, PortRangeEnd: 31010, IsActive: 1, CreatedTime: nowMs, UpdatedTime: nowMs}); err != nil {
|
||||
t.Fatalf("create peer share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load peer share: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, 1, "svc-r1", "svc-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 31001, "", 1, 1, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert peer share runtime: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
CREATE TRIGGER fail_user_quota_insert
|
||||
BEFORE INSERT ON user_quota
|
||||
BEGIN
|
||||
SELECT RAISE(FAIL, 'quota insert blocked for test');
|
||||
END;
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create quota failure trigger: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.applyFlowUploadBatch(1, flowUploadBatch{
|
||||
flowDeltas: []repo.FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 80, OutFlow: 120}},
|
||||
quotaUsage: map[int64]int64{2: 200},
|
||||
policyTargets: []flowPolicyTarget{{UserID: 2, UserTunnelID: 10}},
|
||||
peerShareForwardItems: map[string]flowItem{"20_2_10": {N: "20_2_10", U: 120, D: 80}},
|
||||
}, now)
|
||||
|
||||
if got := mustQueryInt(t, r, `SELECT status FROM forward WHERE id = 20`); got != 0 {
|
||||
t.Fatalf("expected flow-policy enforcement to pause forward after quota failure, got status=%d", got)
|
||||
}
|
||||
updatedShare, err := r.GetPeerShare(share.ID)
|
||||
if err != nil || updatedShare == nil {
|
||||
t.Fatalf("reload peer share: %v", err)
|
||||
}
|
||||
if updatedShare.CurrentFlow != 200 {
|
||||
t.Fatalf("expected peer-share flow accounting to continue after quota failure, got %d", updatedShare.CurrentFlow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFlowUploadBatchContinuesPeerShareSideEffectsWhenFlowBatchFails(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "flow-upload-batch-flow-fail.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
if err := r.DB().Create(&model.User{ID: 2, User: "flow-user", Pwd: "pwd", RoleID: 1, ExpTime: 2727251700000, Flow: 99999, Num: 99999, CreatedTime: nowMs, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed user: %v", err)
|
||||
}
|
||||
if err := r.DB().Create(&model.Tunnel{ID: 1, Name: "tunnel-1", TrafficRatio: 1, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Create(&model.UserTunnel{ID: 10, UserID: 2, TunnelID: 1, Num: 99999, Flow: 0, ExpTime: 2727251700000, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed user tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Create(&model.Forward{ID: 20, UserID: 2, UserName: "flow-user", Name: "forward-20", TunnelID: 1, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed forward: %v", err)
|
||||
}
|
||||
if err := r.DB().Create(&model.Forward{ID: 21, UserID: 2, UserName: "flow-user", Name: "forward-21", TunnelID: 1, RemoteAddr: "1.1.1.1:81", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
|
||||
t.Fatalf("seed second forward: %v", err)
|
||||
}
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{Name: "share", NodeID: 1, Token: "token", MaxBandwidth: 0, CurrentFlow: 0, PortRangeStart: 31000, PortRangeEnd: 31010, IsActive: 1, CreatedTime: nowMs, UpdatedTime: nowMs}); err != nil {
|
||||
t.Fatalf("create peer share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load peer share: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, 1, "svc-r1", "svc-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 31001, "", 1, 1, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert peer share runtime: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
CREATE TRIGGER fail_forward_flow_update
|
||||
BEFORE UPDATE ON forward
|
||||
WHEN NEW.id = 21 AND (NEW.in_flow != OLD.in_flow OR NEW.out_flow != OLD.out_flow)
|
||||
BEGIN
|
||||
SELECT RAISE(FAIL, 'forward flow update blocked for test');
|
||||
END;
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create flow failure trigger: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.applyFlowUploadBatch(1, flowUploadBatch{
|
||||
flowDeltas: []repo.FlowUploadCounterDelta{
|
||||
{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 80, OutFlow: 120},
|
||||
{ForwardID: 21, UserID: 2, UserTunnelID: 10, InFlow: 30, OutFlow: 40},
|
||||
},
|
||||
quotaUsage: map[int64]int64{2: 200},
|
||||
policyTargets: []flowPolicyTarget{{UserID: 2, UserTunnelID: 10}},
|
||||
peerShareForwardItems: map[string]flowItem{"20_2_10": {N: "20_2_10", U: 120, D: 80}},
|
||||
}, now)
|
||||
|
||||
if got := mustQueryInt(t, r, `SELECT status FROM forward WHERE id = 20`); got != 0 {
|
||||
t.Fatalf("expected flow-policy enforcement to pause forward after flow batch failure, got status=%d", got)
|
||||
}
|
||||
updatedShare, err := r.GetPeerShare(share.ID)
|
||||
if err != nil || updatedShare == nil {
|
||||
t.Fatalf("reload peer share: %v", err)
|
||||
}
|
||||
if updatedShare.CurrentFlow != 200 {
|
||||
t.Fatalf("expected peer-share flow accounting to continue after flow batch failure, got %d", updatedShare.CurrentFlow)
|
||||
}
|
||||
if got := mustQueryInt(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 80 {
|
||||
t.Fatalf("expected flow fallback to persist forward 20 in_flow=80, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, r, `SELECT in_flow FROM forward WHERE id = 21`); got != 0 {
|
||||
t.Fatalf("expected failed forward 21 delta to remain unapplied, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, r, `SELECT in_flow FROM user WHERE id = 2`); got != 80 {
|
||||
t.Fatalf("expected flow fallback to preserve successful user totals, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 80 {
|
||||
t.Fatalf("expected flow fallback to preserve successful user_tunnel totals, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFlowUploadBatchFallsBackToPerUserQuotaUpdates(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "flow-upload-batch-quota-fallback.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
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)`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user 2: %v", err)
|
||||
}
|
||||
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(3, 'u3', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user 3: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 0, 0, 0, 0, ?, ?, 0, 0, '', ?, ?), (3, 0, 0, 0, 0, ?, ?, 0, 0, '', ?, ?)`, dayKey, monthKey, nowMs, nowMs, dayKey, monthKey, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user quotas: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
CREATE TRIGGER fail_user_3_quota_update
|
||||
BEFORE UPDATE ON user_quota
|
||||
WHEN NEW.user_id = 3 AND (NEW.daily_used_bytes != OLD.daily_used_bytes OR NEW.monthly_used_bytes != OLD.monthly_used_bytes)
|
||||
BEGIN
|
||||
SELECT RAISE(FAIL, 'quota update blocked for user 3');
|
||||
END;
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create quota fallback trigger: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.applyFlowUploadBatch(1, flowUploadBatch{quotaUsage: map[int64]int64{2: 200, 3: 300}}, now)
|
||||
|
||||
if got := mustQueryInt(t, r, `SELECT daily_used_bytes FROM user_quota WHERE user_id = 2`); got != 200 {
|
||||
t.Fatalf("expected quota fallback to persist user 2 usage, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, r, `SELECT daily_used_bytes FROM user_quota WHERE user_id = 3`); got != 0 {
|
||||
t.Fatalf("expected failed user 3 quota delta to remain unapplied, got %d", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildForwardServiceConfigsAppliesProxyProtocolReceiveAndSendIndependently(t *testing.T) {
|
||||
forward := &forwardRecord{
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocolReceive: 1,
|
||||
ProxyProtocolSend: 2,
|
||||
}
|
||||
tunnel := &tunnelRecord{Type: 1}
|
||||
node := &nodeRecord{
|
||||
InterfaceName: "eth0",
|
||||
TCPListenAddr: "0.0.0.0",
|
||||
UDPListenAddr: "0.0.0.0",
|
||||
}
|
||||
|
||||
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
|
||||
for _, service := range services {
|
||||
serviceMetadata, ok := service["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected metadata map, got %T", service["metadata"])
|
||||
}
|
||||
if serviceMetadata["interface"] != "eth0" {
|
||||
t.Fatalf("expected interface metadata eth0, got %v", serviceMetadata["interface"])
|
||||
}
|
||||
if serviceMetadata["proxyProtocol"] != 1 {
|
||||
t.Fatalf("expected service proxyProtocol 1 for receive mode, got %v", serviceMetadata["proxyProtocol"])
|
||||
}
|
||||
|
||||
handlerConfig, ok := service["handler"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected handler config map, got %T", service["handler"])
|
||||
}
|
||||
handlerMetadata, ok := handlerConfig["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected handler metadata map, got %T", handlerConfig["metadata"])
|
||||
}
|
||||
if handlerMetadata["proxyProtocol"] != 2 {
|
||||
t.Fatalf("expected handler proxyProtocol 2, got %v", handlerMetadata["proxyProtocol"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigsKeepsLegacyProxyProtocolAsSend(t *testing.T) {
|
||||
forward := &forwardRecord{
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocol: 2,
|
||||
}
|
||||
tunnel := &tunnelRecord{Type: 1}
|
||||
node := &nodeRecord{
|
||||
TCPListenAddr: "0.0.0.0",
|
||||
UDPListenAddr: "0.0.0.0",
|
||||
}
|
||||
|
||||
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
|
||||
for _, service := range services {
|
||||
serviceMetadata, _ := service["metadata"].(map[string]interface{})
|
||||
if _, ok := serviceMetadata["proxyProtocol"]; ok {
|
||||
t.Fatalf("legacy proxyProtocol should not enable receive mode: %v", serviceMetadata)
|
||||
}
|
||||
|
||||
handlerConfig, ok := service["handler"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected handler config map, got %T", service["handler"])
|
||||
}
|
||||
handlerMetadata, ok := handlerConfig["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected handler metadata map, got %T", handlerConfig["metadata"])
|
||||
}
|
||||
if handlerMetadata["proxyProtocol"] != 2 {
|
||||
t.Fatalf("expected legacy proxyProtocol to send version 2, got %v", handlerMetadata["proxyProtocol"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Create(&model.Forward{
|
||||
UserID: 2,
|
||||
UserName: "rollback-user",
|
||||
Name: "rollback-forward",
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "9.9.9.9:443",
|
||||
Strategy: "fifo",
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
IPMaxConn: 5,
|
||||
IPSpeedID: sql.NullInt64{Int64: 21, Valid: true},
|
||||
ProxyProtocol: 2,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create forward: %v", err)
|
||||
}
|
||||
|
||||
forwardID := mustLastInsertID(t, r, "rollback-forward")
|
||||
if err := r.DB().Model(&model.Forward{}).Where("id = ?", forwardID).Updates(map[string]interface{}{
|
||||
"name": "changed-forward",
|
||||
"ip_max_conn": 0,
|
||||
"ip_speed_id": nil,
|
||||
"proxy_protocol": 0,
|
||||
"updated_time": now + 1,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("mutate forward: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.rollbackForwardMutation(&forwardRecord{
|
||||
ID: forwardID,
|
||||
UserID: 2,
|
||||
UserName: "rollback-user",
|
||||
Name: "rollback-forward",
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "9.9.9.9:443",
|
||||
Strategy: "fifo",
|
||||
Status: 1,
|
||||
IPMaxConn: 5,
|
||||
IPSpeedID: sql.NullInt64{Int64: 21, Valid: true},
|
||||
ProxyProtocol: 2,
|
||||
}, nil)
|
||||
|
||||
var record model.Forward
|
||||
if err := r.DB().Where("id = ?", forwardID).First(&record).Error; err != nil {
|
||||
t.Fatalf("query forward: %v", err)
|
||||
}
|
||||
if record.ProxyProtocol != 2 {
|
||||
t.Fatalf("expected proxyProtocol restored to 2, got %d", record.ProxyProtocol)
|
||||
}
|
||||
if record.IPMaxConn != 5 {
|
||||
t.Fatalf("expected ipMaxConn restored to 5, got %d", record.IPMaxConn)
|
||||
}
|
||||
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 21 {
|
||||
t.Fatalf("expected ipSpeedId restored to 21, got %+v", record.IPSpeedID)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestForwardResetFlowPermissionsAndIsolation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
actorID int64
|
||||
actorRole int
|
||||
forwardID int64
|
||||
wantCode int
|
||||
wantInFlow int64
|
||||
wantOutFlow int64
|
||||
}{
|
||||
{name: "admin resets another user's rule", actorID: 1, actorRole: 0, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "owner resets own rule", actorID: 2, actorRole: 1, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "user cannot reset another user's rule", actorID: 3, actorRole: 1, forwardID: 20, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
{name: "missing rule is rejected", actorID: 1, actorRole: 0, forwardID: 999, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
h, r := setupForwardResetFlowHandler(t)
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, tt.forwardID, tt.actorID, tt.actorRole)
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
h.forwardResetFlow(res, req)
|
||||
|
||||
if got := decodeForwardResetFlowCode(t, res); got != tt.wantCode {
|
||||
t.Fatalf("code = %d, want %d; body=%s", got, tt.wantCode, res.Body.String())
|
||||
}
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM forward WHERE id = 20", tt.wantInFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM forward WHERE id = 20", tt.wantOutFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM user WHERE id = 2", 700)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM user_tunnel WHERE id = 10", 600)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResetFlowRejectsInvalidRequests(t *testing.T) {
|
||||
h, _ := setupForwardResetFlowHandler(t)
|
||||
|
||||
t.Run("non post", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodGet, 20, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid id", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, 0, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func setupForwardResetFlowHandler(t *testing.T) (*Handler, *repo.Repository) {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "forward-reset-handler.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
statements := []string{
|
||||
`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, 'owner', 'pwd', 1, 0, 100, 700, 900, 0, 10, 1000, 1000, 1)`,
|
||||
`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(3, 'other', 'pwd', 1, 0, 100, 0, 0, 0, 10, 1000, 1000, 1)`,
|
||||
`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 'tunnel', 1, 1, 'tls', 1, 1000, 1000, 1, NULL, 0)`,
|
||||
`INSERT INTO user_tunnel(id, user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, 10, 100, 500, 600, 0, 0, 1)`,
|
||||
`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, 'owner', 'target', 1, '127.0.0.1:80', 'fifo', 111, 222, 1000, 1000, 1, 0)`,
|
||||
}
|
||||
for _, statement := range statements {
|
||||
if err := r.DB().Exec(statement).Error; err != nil {
|
||||
t.Fatalf("seed database: %v", err)
|
||||
}
|
||||
}
|
||||
return New(r, "test-secret"), r
|
||||
}
|
||||
|
||||
func newForwardResetFlowRequest(t *testing.T, method string, forwardID, actorID int64, roleID int) *http.Request {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(map[string]int64{"id": forwardID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(method, "/api/v1/forward/reset-flow", bytes.NewReader(body))
|
||||
claims := auth.Claims{Sub: strconv.FormatInt(actorID, 10), RoleID: roleID}
|
||||
return req.WithContext(context.WithValue(req.Context(), middleware.ClaimsContextKey, claims))
|
||||
}
|
||||
|
||||
func decodeForwardResetFlowCode(t *testing.T, res *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
}
|
||||
if err := json.Unmarshal(res.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("decode response: %v; body=%s", err, res.Body.String())
|
||||
}
|
||||
return payload.Code
|
||||
}
|
||||
|
||||
func assertForwardResetFlowDBValue(t *testing.T, r *repo.Repository, query string, want int64) {
|
||||
t.Helper()
|
||||
var got int64
|
||||
if err := r.DB().Raw(query).Scan(&got).Error; err != nil {
|
||||
t.Fatalf("query %q: %v", query, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("query %q returned %d, want %d", query, got, want)
|
||||
}
|
||||
}
|
||||
@@ -5,8 +5,10 @@ import (
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
@@ -19,8 +21,9 @@ import (
|
||||
"go-backend/internal/health"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"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"
|
||||
@@ -29,27 +32,36 @@ 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
|
||||
|
||||
jobsMu sync.Mutex
|
||||
jobsCancel context.CancelFunc
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
jobsMu sync.Mutex
|
||||
jobsCancel context.CancelFunc
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
fingerprintMu sync.Mutex
|
||||
licenseValidationMu sync.Mutex
|
||||
|
||||
upgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
upgradeMu sync.Mutex
|
||||
systemUpgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
nodeOnlineRedeployAt map[int64]time.Time
|
||||
nodeOnlineRedeployQueued map[int64]struct{}
|
||||
nodeOnlineRedeploying map[int64]struct{}
|
||||
|
||||
qualityProber *tunnelQualityProber
|
||||
bestExit *bestExitManager
|
||||
}
|
||||
|
||||
const monitorTunnelQualityEnabledConfigKey = "monitor_tunnel_quality_enabled"
|
||||
const allowLocalRemoteAddrConfigKey = "allow_local_remote_addr"
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username"`
|
||||
@@ -95,13 +107,18 @@ const (
|
||||
|
||||
func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
h := &Handler{
|
||||
repo: repo,
|
||||
jwtSecret: jwtSecret,
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
metrics: metrics.NewIngestionService(repo),
|
||||
healthCheck: nil,
|
||||
captchaTokens: make(map[string]int64),
|
||||
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||
repo: repo,
|
||||
jwtSecret: jwtSecret,
|
||||
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),
|
||||
nodeOnlineRedeployQueued: make(map[int64]struct{}),
|
||||
nodeOnlineRedeploying: make(map[int64]struct{}),
|
||||
bestExit: newBestExitManager(),
|
||||
}
|
||||
h.healthCheck = health.NewChecker(repo, h.wsServer)
|
||||
h.qualityProber = newTunnelQualityProber(h)
|
||||
@@ -124,6 +141,7 @@ func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
}
|
||||
h.metrics.RecordNodeMetric(nodeID, metricInfo)
|
||||
})
|
||||
h.wsServer.SetUserAuthStateLookup(h.GetUserAuthState)
|
||||
return h
|
||||
}
|
||||
|
||||
@@ -131,6 +149,10 @@ func (h *Handler) WebSocketHandler() http.Handler {
|
||||
return h.wsServer
|
||||
}
|
||||
|
||||
func (h *Handler) GetUserAuthState(userID int64) (*auth.UserAuthState, error) {
|
||||
return h.repo.GetUserAuthState(userID)
|
||||
}
|
||||
|
||||
func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/user/login", h.login)
|
||||
mux.HandleFunc("/api/v1/user/list", h.userList)
|
||||
@@ -140,10 +162,15 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/user/reset", h.userResetFlow)
|
||||
mux.HandleFunc("/api/v1/user/quota/reset", h.userQuotaReset)
|
||||
mux.HandleFunc("/api/v1/user/groups", h.userGroups)
|
||||
mux.HandleFunc("/api/v1/public/config/get", h.getPublicConfigByName)
|
||||
mux.HandleFunc("/api/v1/config/get", h.getConfigByName)
|
||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig)
|
||||
mux.HandleFunc("/api/v1/system/storage", h.storageSummary)
|
||||
mux.HandleFunc("/api/v1/system/version", h.systemVersion)
|
||||
mux.HandleFunc("/api/v1/system/check-updates", h.systemCheckUpdates)
|
||||
mux.HandleFunc("/api/v1/system/upgrade", h.systemUpgrade)
|
||||
mux.HandleFunc("/api/v1/license/activate", h.licenseActivate)
|
||||
mux.HandleFunc("/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/backup/import", h.backupImport)
|
||||
@@ -168,6 +195,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)
|
||||
@@ -193,6 +223,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/forward/force-delete", h.forwardForceDelete)
|
||||
mux.HandleFunc("/api/v1/forward/pause", h.forwardPause)
|
||||
mux.HandleFunc("/api/v1/forward/resume", h.forwardResume)
|
||||
mux.HandleFunc("/api/v1/forward/reset-flow", h.forwardResetFlow)
|
||||
mux.HandleFunc("/api/v1/forward/diagnose", h.forwardDiagnose)
|
||||
mux.HandleFunc("/api/v1/forward/diagnose/stream", h.forwardDiagnoseStream)
|
||||
mux.HandleFunc("/api/v1/forward/update-order", h.forwardUpdateOrder)
|
||||
@@ -318,7 +349,8 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("账号或密码错误"))
|
||||
return
|
||||
}
|
||||
if user.Pwd != security.MD5(req.Password) {
|
||||
passwordMatched, passwordWasLegacy := security.VerifyPassword(user.Pwd, req.Password)
|
||||
if !passwordMatched {
|
||||
response.WriteJSON(w, response.ErrDefault("账号或密码错误"))
|
||||
return
|
||||
}
|
||||
@@ -326,8 +358,20 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("账号被停用"))
|
||||
return
|
||||
}
|
||||
issueAt := time.Now()
|
||||
if passwordWasLegacy {
|
||||
updatedAt := time.Now().UnixMilli()
|
||||
hashedPassword, err := security.HashPassword(req.Password)
|
||||
if err != nil {
|
||||
log.Printf("legacy password rehash skipped user_id=%d path=login err=%v", user.ID, err)
|
||||
} else if err := h.repo.UpdateUserPassword(user.ID, hashedPassword, updatedAt); err != nil {
|
||||
log.Printf("legacy password rehash update skipped user_id=%d path=login err=%v", user.ID, err)
|
||||
} else {
|
||||
issueAt = time.UnixMilli(updatedAt + 1)
|
||||
}
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(user.ID, user.User, user.RoleID, h.jwtSecret)
|
||||
token, err := auth.GenerateTokenAt(user.ID, user.User, user.RoleID, h.jwtSecret, issueAt)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -358,13 +402,21 @@ func (h *Handler) getConfigByName(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
configName := strings.ToLower(strings.TrimSpace(req.Name))
|
||||
switch configName {
|
||||
case "license_key", "cloudflare_secret_key", "jwt_secret":
|
||||
if configName == "license_key" || configName == "license_machine_id" || configName == "machine_fingerprint" {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问系统授权凭据"))
|
||||
return
|
||||
}
|
||||
if repo.IsSensitiveConfigKey(configName) && !isAdminRequest(r) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(req.Name)
|
||||
if _, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims); !ok && !repo.IsPublicConfigKey(configName) {
|
||||
response.WriteJSON(w, response.Err(401, "未登录或token已过期"))
|
||||
return
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(configName)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -389,11 +441,13 @@ func (h *Handler) getConfigs(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
ctxClaims := r.Context().Value(middleware.ClaimsContextKey)
|
||||
if claims, ok := ctxClaims.(auth.Claims); !ok || claims.RoleID != 0 {
|
||||
delete(cfgMap, "license_key")
|
||||
delete(cfgMap, "cloudflare_secret_key")
|
||||
delete(cfgMap, "jwt_secret")
|
||||
claims, isAdmin := ctxClaims.(auth.Claims)
|
||||
if !isAdmin || claims.RoleID != 0 {
|
||||
cfgMap = repo.FilterSensitiveConfigs(cfgMap)
|
||||
}
|
||||
delete(cfgMap, "license_key")
|
||||
delete(cfgMap, "license_machine_id")
|
||||
delete(cfgMap, "machine_fingerprint")
|
||||
response.WriteJSON(w, response.OK(cfgMap))
|
||||
}
|
||||
|
||||
@@ -463,6 +517,7 @@ func (h *Handler) tunnelList(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.attachBestExitStates(items)
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
@@ -540,10 +595,27 @@ func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if user == nil || user.Pwd != security.MD5(password) {
|
||||
if user == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("鉴权失败"))
|
||||
return
|
||||
}
|
||||
passwordMatched, passwordWasLegacy := security.VerifyPassword(user.Pwd, password)
|
||||
if !passwordMatched {
|
||||
response.WriteJSON(w, response.ErrDefault("鉴权失败"))
|
||||
return
|
||||
}
|
||||
if user.Status == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("账号被停用"))
|
||||
return
|
||||
}
|
||||
if passwordWasLegacy {
|
||||
hashedPassword, err := security.HashPassword(password)
|
||||
if err != nil {
|
||||
log.Printf("legacy password rehash skipped user_id=%d path=sub_store err=%v", user.ID, err)
|
||||
} else if err := h.repo.UpdateUserPassword(user.ID, hashedPassword, time.Now().UnixMilli()); err != nil {
|
||||
log.Printf("legacy password rehash update skipped user_id=%d path=sub_store err=%v", user.ID, err)
|
||||
}
|
||||
}
|
||||
|
||||
const giga = int64(1024 * 1024 * 1024)
|
||||
headerValue := ""
|
||||
@@ -647,6 +719,7 @@ func (h *Handler) userTunnelList(w http.ResponseWriter, r *http.Request) {
|
||||
"tunnelName": t.TunnelName,
|
||||
"status": t.Status,
|
||||
"flow": t.Flow,
|
||||
"flowMiB": t.FlowMiB,
|
||||
"num": t.Num,
|
||||
"expTime": t.ExpTime,
|
||||
"flowResetTime": t.FlowResetTime,
|
||||
@@ -796,11 +869,16 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
if err == nil && strings.TrimSpace(raw) != "" {
|
||||
var items []flowItem
|
||||
if json.Unmarshal([]byte(raw), &items) == nil {
|
||||
nowMs := time.Now().UnixMilli()
|
||||
h.recordTunnelMetricsFromFlowItems(node.ID, items, nowMs)
|
||||
for _, item := range items {
|
||||
h.processFlowItem(node.ID, item)
|
||||
now := time.Now()
|
||||
forwardIDs := collectFlowUploadForwardIDs(items)
|
||||
metas, metaErr := h.repo.GetFlowUploadForwardMetas(forwardIDs)
|
||||
if metaErr != nil {
|
||||
log.Printf("flow upload metadata lookup failed node_id=%d err=%v", node.ID, metaErr)
|
||||
metas = map[int64]repo.FlowUploadForwardMeta{}
|
||||
}
|
||||
batch := h.buildFlowUploadBatch(items, metas)
|
||||
h.recordTunnelMetricsFromForwardBatch(node.ID, batch.forwardTraffic, metas, now.UnixMilli())
|
||||
h.applyFlowUploadBatch(node.ID, batch, now)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -809,10 +887,16 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
func (h *Handler) getOrCreateMachineFingerprint() (string, error) {
|
||||
fp, _ := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
h.fingerprintMu.Lock()
|
||||
defer h.fingerprintMu.Unlock()
|
||||
|
||||
fp, err := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
if fp != "" {
|
||||
return fp, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return "", err
|
||||
}
|
||||
|
||||
newFp := uuid.New().String()
|
||||
now := time.Now().UnixMilli()
|
||||
@@ -839,56 +923,35 @@ func (h *Handler) licenseActivate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("授权码不能为空"))
|
||||
return
|
||||
}
|
||||
h.licenseValidationMu.Lock()
|
||||
defer h.licenseValidationMu.Unlock()
|
||||
|
||||
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
|
||||
|
||||
fingerprint, err := h.getOrCreateMachineFingerprint()
|
||||
valResp, err := h.validateLicenseForMachine(key)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("生成设备指纹失败"))
|
||||
return
|
||||
}
|
||||
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("连接授权服务器失败: "+err.Error()))
|
||||
log.Printf("license activation failed: %v", err)
|
||||
response.WriteJSON(w, response.ErrDefault(licenseValidationErrorMessage(err)))
|
||||
return
|
||||
}
|
||||
|
||||
if !valResp.Meta.Valid {
|
||||
if valResp.Meta.Code == "NO_MACHINES" || valResp.Meta.Code == "NO_MACHINE" || valResp.Meta.Code == "MACHINE_SCOPE_REQUIRED" || valResp.Meta.Code == "FINGERPRINT_SCOPE_MISMATCH" {
|
||||
// Needs machine activation
|
||||
client.Token = key
|
||||
err = client.ActivateMachine(valResp.Data.ID, fingerprint)
|
||||
if err != nil {
|
||||
// Translate specific error messages or log them
|
||||
response.WriteJSON(w, response.ErrDefault("设备绑定失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
// Validation might still fail with scope if we don't query via machine id, but since activate machine succeeded
|
||||
// we can consider the license valid for our simple usecase
|
||||
} else {
|
||||
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.UpsertConfig("license_key", key, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpsertConfig("is_commercial", "true", now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
expiry := valResp.Data.Attributes.Expiry
|
||||
if expiry == "" {
|
||||
expiry = "never"
|
||||
}
|
||||
if err := h.repo.UpsertConfig("license_expiry", expiry, now); err != nil {
|
||||
licenseState := map[string]string{
|
||||
"license_key": key,
|
||||
"is_commercial": "true",
|
||||
"license_expiry": expiry,
|
||||
}
|
||||
if valResp.MachineID != "" {
|
||||
licenseState["license_machine_id"] = valResp.MachineID
|
||||
}
|
||||
if err := h.repo.UpsertConfigs(licenseState, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -926,6 +989,14 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
if repo.IsSensitiveConfigKey(key) && !isAdminRequest(r) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
if repo.IsSystemManagedConfigKey(key) {
|
||||
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
|
||||
return
|
||||
}
|
||||
|
||||
if protectedKeys[key] && isCommercial != "true" {
|
||||
response.WriteJSON(w, response.ErrDefault("需要商业版授权"))
|
||||
@@ -942,6 +1013,7 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.notifyTunnelQualityConfigChanged(key)
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
@@ -963,6 +1035,14 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("配置名称不能为空"))
|
||||
return
|
||||
}
|
||||
if repo.IsSensitiveConfigKey(name) && !isAdminRequest(r) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
if repo.IsSystemManagedConfigKey(name) {
|
||||
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
|
||||
return
|
||||
}
|
||||
|
||||
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
|
||||
if (name == "app_name" || name == "app_logo" || name == "app_favicon" || name == "hide_footer_brand") && isCommercial != "true" {
|
||||
@@ -985,10 +1065,19 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.notifyTunnelQualityConfigChanged(name)
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func isAdminRequest(r *http.Request) bool {
|
||||
if r == nil {
|
||||
return false
|
||||
}
|
||||
claims, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims)
|
||||
return ok && claims.RoleID == 0
|
||||
}
|
||||
|
||||
func normalizeAndValidateConfigValue(key, value string) (string, error) {
|
||||
switch strings.TrimSpace(key) {
|
||||
case "app_logo", "app_favicon":
|
||||
@@ -1023,11 +1112,25 @@ func normalizeAndValidateConfigValue(key, value string) (string, error) {
|
||||
default:
|
||||
return "", fmt.Errorf("隧道质量检测开关配置值无效")
|
||||
}
|
||||
case monitoring.ConfigMonitorRetentionDays:
|
||||
return monitoring.NormalizeMonitoringRetentionDays(value)
|
||||
case monitoring.ConfigTunnelQualityProbeIntervalSec:
|
||||
return monitoring.NormalizeTunnelQualityProbeIntervalSeconds(value)
|
||||
default:
|
||||
return value, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) notifyTunnelQualityConfigChanged(key string) {
|
||||
if h == nil || h.qualityProber == nil {
|
||||
return
|
||||
}
|
||||
switch strings.TrimSpace(key) {
|
||||
case monitorTunnelQualityEnabledConfigKey, monitoring.ConfigTunnelQualityProbeIntervalSec:
|
||||
h.qualityProber.NotifyConfigChanged()
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) isTunnelQualityMonitoringEnabled() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return true
|
||||
@@ -1041,6 +1144,19 @@ func (h *Handler) isTunnelQualityMonitoringEnabled() bool {
|
||||
return strings.TrimSpace(strings.ToLower(cfg.Value)) != "false"
|
||||
}
|
||||
|
||||
func (h *Handler) allowLocalRemoteAddr() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(allowLocalRemoteAddrConfigKey)
|
||||
if err != nil || cfg == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
return strings.TrimSpace(strings.ToLower(cfg.Value)) == "true"
|
||||
}
|
||||
|
||||
func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -1098,6 +1214,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
"tunnelName": t.TunnelName,
|
||||
"tunnelFlow": t.TunnelFlow,
|
||||
"flow": t.Flow,
|
||||
"flowMiB": t.FlowMiB,
|
||||
"inFlow": t.InFlow,
|
||||
"outFlow": t.OutFlow,
|
||||
"num": t.Num,
|
||||
@@ -1147,6 +1264,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
"user": user.User,
|
||||
"status": user.Status,
|
||||
"flow": user.Flow,
|
||||
"flowMiB": user.FlowMiB,
|
||||
"inFlow": user.InFlow,
|
||||
"outFlow": user.OutFlow,
|
||||
"num": user.Num,
|
||||
@@ -1218,7 +1336,8 @@ func (h *Handler) updatePassword(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if user.Pwd != security.MD5(req.CurrentPassword) {
|
||||
passwordMatched, _ := security.VerifyPassword(user.Pwd, req.CurrentPassword)
|
||||
if !passwordMatched {
|
||||
response.WriteJSON(w, response.ErrDefault("当前密码错误"))
|
||||
return
|
||||
}
|
||||
@@ -1233,7 +1352,12 @@ func (h *Handler) updatePassword(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserNameAndPassword(userID, req.NewUsername, security.MD5(req.NewPassword), time.Now().UnixMilli()); err != nil {
|
||||
hashedPassword, err := security.HashPassword(req.NewPassword)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpdateUserNameAndPassword(userID, req.NewUsername, hashedPassword, time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -2,11 +2,12 @@ package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/license"
|
||||
)
|
||||
|
||||
var nftablesTrafficCollectInterval = 30 * time.Second
|
||||
|
||||
func (h *Handler) StartBackgroundJobs() {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
@@ -20,7 +21,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,10 +31,12 @@ 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) {
|
||||
defer h.jobsWG.Done()
|
||||
h.validateLicenseJob()
|
||||
ticker := time.NewTicker(12 * time.Hour)
|
||||
defer ticker.Stop()
|
||||
|
||||
@@ -51,22 +54,25 @@ func (h *Handler) validateLicenseJob() {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
|
||||
h.licenseValidationMu.Lock()
|
||||
defer h.licenseValidationMu.Unlock()
|
||||
|
||||
key, _ := h.repo.GetViteConfigValue("license_key")
|
||||
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
|
||||
|
||||
if key == "" || isCommercial != "true" {
|
||||
if key == "" {
|
||||
return // Nothing to validate
|
||||
}
|
||||
|
||||
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
|
||||
valResp, err := h.validateLicenseForMachine(key)
|
||||
|
||||
if err != nil {
|
||||
// Network error or timeout. Grace period by not revoking immediately here.
|
||||
// Network and decode failures have no validation response, so retain the
|
||||
// current state as a grace period. A rejected machine binding still has
|
||||
// the original invalid response and must not stay commercially enabled.
|
||||
if licenseValidationErrorIsDefinitive(valResp, err) {
|
||||
now := time.Now().UnixMilli()
|
||||
_ = h.repo.UpsertConfig("is_commercial", "false", now)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -80,7 +86,14 @@ func (h *Handler) validateLicenseJob() {
|
||||
if expiry == "" {
|
||||
expiry = "never"
|
||||
}
|
||||
_ = h.repo.UpsertConfig("license_expiry", expiry, now)
|
||||
licenseState := map[string]string{
|
||||
"is_commercial": "true",
|
||||
"license_expiry": expiry,
|
||||
}
|
||||
if valResp.MachineID != "" {
|
||||
licenseState["license_machine_id"] = valResp.MachineID
|
||||
}
|
||||
_ = h.repo.UpsertConfigs(licenseState, now)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -128,6 +141,54 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) {
|
||||
h.qualityProber.Start(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
h.runNftablesStartupReconcile(ctx)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
h.runNftablesTrafficCollectJob(time.Now())
|
||||
}
|
||||
|
||||
interval := nftablesTrafficCollectInterval
|
||||
if interval <= 0 {
|
||||
interval = 30 * time.Second
|
||||
}
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
h.runNftablesTrafficCollectJob(time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runNftablesStartupReconcile(ctx context.Context) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
nodes, err := h.repo.ListNftablesNodesForCollection()
|
||||
if err != nil {
|
||||
log.Printf("nftables startup reconcile failed op=list_nodes err=%v", err)
|
||||
return
|
||||
}
|
||||
for _, node := range nodes {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
if err := h.syncNftablesNode(node.NodeID); err != nil {
|
||||
log.Printf("nftables startup reconcile failed node_id=%d err=%v", node.NodeID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/license"
|
||||
)
|
||||
|
||||
var newLicenseClient = license.NewKeygenClient
|
||||
|
||||
func keygenAccountID() string {
|
||||
if value := strings.TrimSpace(license.AccountID); value != "" {
|
||||
return value
|
||||
}
|
||||
return strings.TrimSpace(os.Getenv("KEYGEN_ACCOUNT_ID"))
|
||||
}
|
||||
|
||||
func licenseNeedsMachineActivation(code string) bool {
|
||||
switch strings.ToUpper(strings.TrimSpace(code)) {
|
||||
case "NO_MACHINES", "NO_MACHINE", "MACHINE_SCOPE_REQUIRED", "FINGERPRINT_SCOPE_MISMATCH":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func licenseValidationErrorIsDefinitive(validation *license.ValidateResponse, err error) bool {
|
||||
if err == nil {
|
||||
return validation != nil && !validation.Meta.Valid
|
||||
}
|
||||
var apiErr *license.APIError
|
||||
if !errors.As(err, &apiErr) {
|
||||
return false
|
||||
}
|
||||
if apiErr.StatusCode == http.StatusTooManyRequests || apiErr.StatusCode >= http.StatusInternalServerError {
|
||||
return false
|
||||
}
|
||||
return apiErr.Operation == "activate machine" && validation != nil && !validation.Meta.Valid
|
||||
}
|
||||
|
||||
func licenseValidationErrorMessage(err error) string {
|
||||
if strings.Contains(err.Error(), "keygen account id is not configured") {
|
||||
return "授权服务配置错误"
|
||||
}
|
||||
var apiErr *license.APIError
|
||||
if errors.As(err, &apiErr) {
|
||||
if apiErr.HasCode("MACHINE_LIMIT_EXCEEDED") {
|
||||
return "授权设备数量已达上限"
|
||||
}
|
||||
if apiErr.StatusCode == http.StatusUnauthorized || apiErr.StatusCode == http.StatusForbidden {
|
||||
return "授权码无效或无权绑定设备"
|
||||
}
|
||||
}
|
||||
return "连接授权服务器失败,请稍后重试"
|
||||
}
|
||||
|
||||
func (h *Handler) validateLicenseForMachine(key string) (*license.ValidateResponse, error) {
|
||||
fingerprint, err := h.getOrCreateMachineFingerprint()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prepare machine fingerprint: %w", err)
|
||||
}
|
||||
storedMachineID, _ := h.repo.GetViteConfigValue("license_machine_id")
|
||||
|
||||
accountID := keygenAccountID()
|
||||
if accountID == "" {
|
||||
return nil, fmt.Errorf("keygen account id is not configured")
|
||||
}
|
||||
client := newLicenseClient(accountID, "")
|
||||
var validation *license.ValidateResponse
|
||||
if storedMachineID != "" {
|
||||
validation, err = client.ValidateKeyWithMachine(key, fingerprint, storedMachineID)
|
||||
} else {
|
||||
validation, err = client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validation.Meta.Valid || !licenseNeedsMachineActivation(validation.Meta.Code) {
|
||||
if validation.Meta.Valid {
|
||||
validation.MachineID = storedMachineID
|
||||
}
|
||||
return validation, nil
|
||||
}
|
||||
|
||||
client.Token = key
|
||||
machineID, err := client.ActivateMachine(validation.Data.ID, fingerprint)
|
||||
if err != nil {
|
||||
return validation, err
|
||||
}
|
||||
if machineID == "" {
|
||||
machineID, err = client.GetMachineID(fingerprint)
|
||||
if err != nil {
|
||||
return validation, fmt.Errorf("retrieve activated machine: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
validation, err = client.ValidateKeyWithMachine(key, fingerprint, machineID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validation.Meta.Valid {
|
||||
validation.MachineID = machineID
|
||||
}
|
||||
return validation, nil
|
||||
}
|
||||
@@ -0,0 +1,358 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestValidateLicenseJobRepairsMissingMachineBinding(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
|
||||
var validations atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
if validations.Add(1) == 1 {
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
return
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
|
||||
fingerprint, err := r.GetViteConfigValue("machine_fingerprint")
|
||||
if err != nil || strings.TrimSpace(fingerprint) == "" {
|
||||
t.Fatalf("expected persisted machine fingerprint, got value=%q err=%v", fingerprint, err)
|
||||
}
|
||||
if got := validations.Load(); got != 2 {
|
||||
t.Fatalf("validation calls = %d, want 2", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobAcceptsExistingMachineActivation(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
|
||||
var validations atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
if validations.Add(1) == 1 {
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
return
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "never")
|
||||
if got := validations.Load(); got != 2 {
|
||||
t.Fatalf("validation calls = %d, want 2", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseActivateRequiresSuccessfulPostActivationValidation(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
|
||||
var validations atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
code := "NO_MACHINE"
|
||||
if validations.Add(1) > 1 {
|
||||
code = "FINGERPRINT_SCOPE_MISMATCH"
|
||||
}
|
||||
_, _ = fmt.Fprintf(w, `{"meta":{"valid":false,"code":%q},"data":{"id":"license-id","attributes":{}}}`, code)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", bytes.NewBufferString(`{"license_key":"license-secret"}`))
|
||||
res := httptest.NewRecorder()
|
||||
h.licenseActivate(res, req)
|
||||
|
||||
if !strings.Contains(res.Body.String(), "FINGERPRINT_SCOPE_MISMATCH") {
|
||||
t.Fatalf("expected post-activation validation failure, got %s", res.Body.String())
|
||||
}
|
||||
assertLicenseConfig(t, r, "is_commercial", "false")
|
||||
for _, name := range []string{"license_key", "license_expiry"} {
|
||||
if value, err := r.GetViteConfigValue(name); err == nil || value != "" {
|
||||
t.Fatalf("%s should not be persisted, got value=%q err=%v", name, value, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseActivatePersistsValidatedState(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", bytes.NewBufferString(`{"license_key":"license-secret"}`))
|
||||
res := httptest.NewRecorder()
|
||||
h.licenseActivate(res, req)
|
||||
|
||||
if !strings.Contains(res.Body.String(), `"code":0`) {
|
||||
t.Fatalf("expected activation success, got %s", res.Body.String())
|
||||
}
|
||||
assertLicenseConfig(t, r, "license_key", "license-secret")
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobDowngradesWhenMachineBindingIsRejected(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "false")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobRestoresCommercialStateWhenLicenseRecovers(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "false", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
if strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, req)
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "never")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobUsesStoredMachineScope(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
seedLicenseConfig(t, r, "license_machine_id", "machine-id", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
if !strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
|
||||
t.Fatalf("unexpected request %s %s", req.Method, req.URL.Path)
|
||||
}
|
||||
var body struct {
|
||||
Meta struct {
|
||||
Scope map[string]string `json:"scope"`
|
||||
} `json:"meta"`
|
||||
}
|
||||
if err := json.NewDecoder(req.Body).Decode(&body); err != nil {
|
||||
t.Fatalf("decode request: %v", err)
|
||||
}
|
||||
if body.Meta.Scope["machine"] != "machine-id" || body.Meta.Scope["fingerprint"] != "fingerprint" {
|
||||
t.Fatalf("unexpected validation scope: %+v", body.Meta.Scope)
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_machine_id", "machine-id")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobKeepsStateOnMachineLookupServerFailure(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"SERVICE_UNAVAILABLE"}]}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobKeepsStateOnMachineLookupNotFound(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
http.NotFound(w, req)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
}
|
||||
|
||||
func TestLicenseValidationErrorMessageDoesNotExposeKeygenResponse(t *testing.T) {
|
||||
err := &license.APIError{
|
||||
Operation: "activate machine",
|
||||
StatusCode: http.StatusUnprocessableEntity,
|
||||
Body: `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED","detail":"private detail"}]}`,
|
||||
}
|
||||
message := licenseValidationErrorMessage(err)
|
||||
if message != "授权设备数量已达上限" || strings.Contains(message, "private detail") {
|
||||
t.Fatalf("unexpected public error message %q", message)
|
||||
}
|
||||
}
|
||||
|
||||
func openLicenseTestRepository(t *testing.T) *repo.Repository {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "license.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("repo.Open() error = %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
return r
|
||||
}
|
||||
|
||||
func seedLicenseConfig(t *testing.T, r *repo.Repository, name, value string, now int64) {
|
||||
t.Helper()
|
||||
if err := r.UpsertConfig(name, value, now); err != nil {
|
||||
t.Fatalf("UpsertConfig(%q) error = %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertLicenseConfig(t *testing.T, r *repo.Repository, name, want string) {
|
||||
t.Helper()
|
||||
got, err := r.GetViteConfigValue(name)
|
||||
if err != nil {
|
||||
t.Fatalf("GetViteConfigValue(%q) error = %v", name, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("config %q = %q, want %q", name, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func restoreLicenseClientFactory(t *testing.T, baseURL string) {
|
||||
t.Helper()
|
||||
t.Setenv("KEYGEN_ACCOUNT_ID", "account-id")
|
||||
previous := newLicenseClient
|
||||
newLicenseClient = func(accountID, token string) *license.KeygenClient {
|
||||
client := license.NewKeygenClient(accountID, token)
|
||||
client.BaseURL = baseURL
|
||||
return client
|
||||
}
|
||||
t.Cleanup(func() { newLicenseClient = previous })
|
||||
}
|
||||
@@ -234,8 +234,22 @@ func (h *Handler) monitorTunnelQualityHandler(w http.ResponseWriter, r *http.Req
|
||||
return
|
||||
}
|
||||
|
||||
targetsByTunnelID := map[int64]tunnelProbeTarget{}
|
||||
if tunnels, listErr := h.repo.ListTunnels(); listErr == nil {
|
||||
for _, item := range tunnels {
|
||||
id := asInt64(item["id"], 0)
|
||||
if id > 0 {
|
||||
targetsByTunnelID[id] = effectiveTunnelProbeTargetValues(asString(item["probeTargetHost"]), asInt(item["probeTargetPort"], 0))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
snapshots := make([]tunnelQualitySnapshot, 0, len(qualities))
|
||||
for _, q := range qualities {
|
||||
target := targetsByTunnelID[q.TunnelID]
|
||||
if target.Host == "" {
|
||||
target = defaultTunnelProbeTarget()
|
||||
}
|
||||
snapshots = append(snapshots, tunnelQualitySnapshot{
|
||||
TunnelID: q.TunnelID,
|
||||
EntryToExitLatency: q.EntryToExitLatency,
|
||||
@@ -246,6 +260,8 @@ func (h *Handler) monitorTunnelQualityHandler(w http.ResponseWriter, r *http.Req
|
||||
ErrorMessage: q.ErrorMessage,
|
||||
Timestamp: q.Timestamp,
|
||||
ChainDetails: q.ChainDetails,
|
||||
ProbeTargetHost: target.Host,
|
||||
ProbeTargetPort: target.Port,
|
||||
})
|
||||
}
|
||||
response.WriteJSON(w, response.OK(snapshots))
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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,711 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"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 {
|
||||
mu sync.Mutex
|
||||
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.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.lastConfig = cfg
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
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: runtimenft.PlanHashes(plan),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.clearHit++
|
||||
return f.clearErr
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.collectHit++
|
||||
f.lastConfig = cfg
|
||||
if f.collectErr != nil {
|
||||
return nil, f.collectErr
|
||||
}
|
||||
return f.counterSamples, nil
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) reconcileCount() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return f.reconcileHit
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) collectCount() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return f.collectHit
|
||||
}
|
||||
|
||||
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 TestStartBackgroundJobsReconcilesNftablesRulesAtStartup(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-startup-tunnel", fixture.nodeID)
|
||||
seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
h.StartBackgroundJobs()
|
||||
t.Cleanup(h.StopBackgroundJobs)
|
||||
|
||||
waitForCondition(t, time.Second, func() bool {
|
||||
return manager.reconcileCount() > 0
|
||||
}, "nftables startup reconcile")
|
||||
}
|
||||
|
||||
func TestStartBackgroundJobsCollectsNftablesTrafficImmediatelyAndUsesFastInterval(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},
|
||||
}}
|
||||
h.nftablesManager = manager
|
||||
|
||||
oldInterval := nftablesTrafficCollectInterval
|
||||
nftablesTrafficCollectInterval = 20 * time.Millisecond
|
||||
t.Cleanup(func() { nftablesTrafficCollectInterval = oldInterval })
|
||||
|
||||
h.StartBackgroundJobs()
|
||||
t.Cleanup(h.StopBackgroundJobs)
|
||||
|
||||
waitForCondition(t, time.Second, func() bool {
|
||||
return manager.collectCount() >= 2
|
||||
}, "immediate and repeated nftables traffic collection")
|
||||
}
|
||||
|
||||
func TestNftablesTrafficCollectIntervalDefaultsToThirtySeconds(t *testing.T) {
|
||||
if nftablesTrafficCollectInterval != 30*time.Second {
|
||||
t.Fatalf("expected default nftables traffic collection interval 30s, got %s", nftablesTrafficCollectInterval)
|
||||
}
|
||||
}
|
||||
|
||||
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 TestNodeUpdateSkipsAgentProtocolCommandForNftablesNode(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",
|
||||
"http": 1,
|
||||
"tls": 1,
|
||||
"socks": 1,
|
||||
"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.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 TestDiagnoseForwardRuntimeReturnsNftablesRuleStatus(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-diagnose-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
now := time.Now().UnixMilli()
|
||||
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)
|
||||
}
|
||||
|
||||
payload, err := h.diagnoseForwardRuntime(context.Background(), &forwardRecord{
|
||||
ID: forward.ID,
|
||||
Name: forward.Name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: "203.0.113.9:8080",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("diagnose forward: %v", err)
|
||||
}
|
||||
results, ok := payload["results"].([]map[string]interface{})
|
||||
if !ok || len(results) != 1 {
|
||||
t.Fatalf("expected one nftables diagnosis result, got %#v", payload["results"])
|
||||
}
|
||||
result := results[0]
|
||||
if result["forwardMode"] != "nftables" || result["nftRuleStatus"] != runtimenft.StatusApplied {
|
||||
t.Fatalf("expected nftables applied result, got %#v", result)
|
||||
}
|
||||
if result["success"] != true {
|
||||
t.Fatalf("expected nftables diagnosis success, got %#v", result)
|
||||
}
|
||||
if !strings.Contains(asString(result["message"]), "已下发") {
|
||||
t.Fatalf("expected applied message, got %#v", result["message"])
|
||||
}
|
||||
}
|
||||
|
||||
func waitForCondition(t *testing.T, timeout time.Duration, condition func() bool, description string) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
if condition() {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("timed out waiting for %s", description)
|
||||
}
|
||||
|
||||
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, 0, 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
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
@@ -11,11 +12,37 @@ var DisableSafeRemoteAddrCheckForTesting = false
|
||||
|
||||
// IsSafeRemoteAddr checks if a given address is safe to connect to (prevents SSRF/Open Proxy).
|
||||
// It resolves domains to IPs to prevent DNS rebinding attacks pointing to internal networks.
|
||||
// Supports multiple addresses separated by commas or newlines (one per line).
|
||||
func IsSafeRemoteAddr(addr string) error {
|
||||
if DisableSafeRemoteAddrCheckForTesting {
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, part := range splitRemoteParts(addr) {
|
||||
if err := checkSingleRemoteAddr(part); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// splitRemoteParts splits a multi-address string by commas and newlines.
|
||||
func splitRemoteParts(addr string) []string {
|
||||
addr = strings.ReplaceAll(addr, "\n", ",")
|
||||
addr = strings.ReplaceAll(addr, "\r", ",")
|
||||
parts := strings.Split(addr, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
if part != "" {
|
||||
out = append(out, part)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// checkSingleRemoteAddr validates a single address.
|
||||
func checkSingleRemoteAddr(addr string) error {
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "missing port in address") {
|
||||
@@ -27,12 +54,12 @@ func IsSafeRemoteAddr(addr string) error {
|
||||
|
||||
ips, err := net.LookupIP(host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("could not resolve address: %v", err)
|
||||
return fmt.Errorf("could not resolve address %q: %v", addr, err)
|
||||
}
|
||||
|
||||
for _, ip := range ips {
|
||||
if ip.IsLoopback() || ip.IsPrivate() {
|
||||
return fmt.Errorf("address resolves to internal IP: %s", ip.String())
|
||||
return fmt.Errorf("address %q resolves to internal IP: %s", addr, ip.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -49,6 +76,12 @@ func IsValidNodeAddress(addr string) error {
|
||||
if strings.ContainsAny(addr, "/?") {
|
||||
return fmt.Errorf("address must not contain path or query parameters")
|
||||
}
|
||||
// A bare IPv6 literal contains multiple colons, so net.SplitHostPort treats
|
||||
// it as a malformed host:port pair. Accept IP literals before attempting
|
||||
// host:port parsing; netip also handles scoped IPv6 addresses.
|
||||
if _, err := netip.ParseAddr(addr); err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
_, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
package handler
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestIssue515IsValidNodeAddressAcceptsBareIPv6(t *testing.T) {
|
||||
for _, addr := range []string{
|
||||
"2001:db8::1",
|
||||
"::1",
|
||||
"fe80::1%eth0",
|
||||
} {
|
||||
t.Run(addr, func(t *testing.T) {
|
||||
if err := IsValidNodeAddress(addr); err != nil {
|
||||
t.Fatalf("expected bare IPv6 address %q to be accepted: %v", addr, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsValidNodeAddressKeepsExistingAddressForms(t *testing.T) {
|
||||
for _, addr := range []string{
|
||||
"203.0.113.10",
|
||||
"node.example.com",
|
||||
"node.example.com:6365",
|
||||
"[2001:db8::1]:6365",
|
||||
} {
|
||||
t.Run(addr, func(t *testing.T) {
|
||||
if err := IsValidNodeAddress(addr); err != nil {
|
||||
t.Fatalf("expected node address %q to be accepted: %v", addr, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsValidNodeAddressRejectsURLComponents(t *testing.T) {
|
||||
for _, addr := range []string{
|
||||
"https://node.example.com",
|
||||
"node.example.com/path",
|
||||
"node.example.com?transport=tcp",
|
||||
} {
|
||||
t.Run(addr, func(t *testing.T) {
|
||||
if err := IsValidNodeAddress(addr); err == nil {
|
||||
t.Fatalf("expected node address %q to be rejected", addr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func (h *Handler) storageSummary(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if h == nil || h.repo == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "repository not initialized"))
|
||||
return
|
||||
}
|
||||
|
||||
summary, err := h.repo.DatabaseStorageSummary()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(summary))
|
||||
}
|
||||
@@ -0,0 +1,657 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
const (
|
||||
panelDeployDirEnv = "PANEL_DEPLOY_DIR"
|
||||
panelBackendContainerEnv = "PANEL_BACKEND_CONTAINER"
|
||||
defaultPanelDeployDir = "/opt/flvx-panel"
|
||||
defaultPanelBackendName = "flux-panel-backend"
|
||||
dockerSocketPath = "/var/run/docker.sock"
|
||||
maxSystemUpgradeComposeAssetBytes = 1 << 20
|
||||
systemUpgradeMessage = "升级 helper 已启动,面板服务将短暂重启"
|
||||
systemUpgradeConflictError = "已有面板升级任务执行中"
|
||||
)
|
||||
|
||||
var safeBackendContainerPattern = regexp.MustCompile(`^[A-Za-z0-9_.-]+$`)
|
||||
var enableIPv6ComposePattern = regexp.MustCompile(`(?im)^\s*enable_ipv6\s*:\s*['"]?true['"]?\s*(?:#.*)?$`)
|
||||
var systemUpgradeReleaseBaseURL = githubHTMLBase
|
||||
var systemUpgradeAPIBaseURL = githubAPIBase
|
||||
var systemUpgradeHTTPGet = func(client *http.Client, url string) (*http.Response, error) {
|
||||
return client.Get(url)
|
||||
}
|
||||
|
||||
type systemUpgradeExecutor struct {
|
||||
deployDir string
|
||||
backendContainer string
|
||||
}
|
||||
|
||||
type systemUpgradeCapabilityData struct {
|
||||
Capable bool `json:"capable"`
|
||||
Reasons []string `json:"reasons"`
|
||||
DeployDir string `json:"deployDir"`
|
||||
BackendContainer string `json:"backendContainer"`
|
||||
}
|
||||
|
||||
type systemUpgradeReleaseData struct {
|
||||
Version string `json:"version"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"publishedAt"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
|
||||
type systemUpgradeVersionData struct {
|
||||
CurrentVersion string `json:"currentVersion"`
|
||||
LatestVersion string `json:"latestVersion"`
|
||||
HasUpdate bool `json:"hasUpdate"`
|
||||
Channel string `json:"channel"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
Capability systemUpgradeCapabilityData `json:"capability"`
|
||||
}
|
||||
|
||||
type systemUpgradeCheckData struct {
|
||||
CurrentVersion string `json:"currentVersion"`
|
||||
LatestVersion string `json:"latestVersion"`
|
||||
HasUpdate bool `json:"hasUpdate"`
|
||||
Channel string `json:"channel"`
|
||||
Capability systemUpgradeCapabilityData `json:"capability"`
|
||||
Releases []systemUpgradeReleaseData `json:"releases"`
|
||||
}
|
||||
|
||||
type systemUpgradeRunData struct {
|
||||
Version string `json:"version"`
|
||||
Channel string `json:"channel"`
|
||||
ComposeAsset string `json:"composeAsset"`
|
||||
HelperContainer string `json:"helperContainer"`
|
||||
BackendImageID string `json:"backendImageId"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
type systemUpgradeRequest struct {
|
||||
Version string `json:"version"`
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
|
||||
func newSystemUpgradeExecutor() *systemUpgradeExecutor {
|
||||
deployDir := strings.TrimSpace(os.Getenv(panelDeployDirEnv))
|
||||
if deployDir == "" {
|
||||
deployDir = defaultPanelDeployDir
|
||||
}
|
||||
backendContainer := strings.TrimSpace(os.Getenv(panelBackendContainerEnv))
|
||||
if backendContainer == "" {
|
||||
backendContainer = defaultPanelBackendName
|
||||
}
|
||||
return &systemUpgradeExecutor{deployDir: deployDir, backendContainer: backendContainer}
|
||||
}
|
||||
|
||||
func currentPanelVersion() string {
|
||||
version := strings.TrimSpace(os.Getenv("FLUX_VERSION"))
|
||||
if version == "" {
|
||||
return "dev"
|
||||
}
|
||||
return version
|
||||
}
|
||||
|
||||
func validateBackendContainerName(value string) error {
|
||||
if value == "" {
|
||||
return fmt.Errorf("backend container name is empty")
|
||||
}
|
||||
if !safeBackendContainerPattern.MatchString(value) {
|
||||
return fmt.Errorf("unsafe backend container name: %s", value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUpgradeVersion(value string) error {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fmt.Errorf("upgrade version is empty")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r < 0x20 || r == 0x7f {
|
||||
return fmt.Errorf("unsafe upgrade version: contains control character")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) composePath() string {
|
||||
return filepath.Join(e.deployDir, "docker-compose.yml")
|
||||
}
|
||||
func (e *systemUpgradeExecutor) envPath() string { return filepath.Join(e.deployDir, ".env") }
|
||||
|
||||
func (e *systemUpgradeExecutor) capability(ctx context.Context) systemUpgradeCapabilityData {
|
||||
reasons := make([]string, 0)
|
||||
if !filepath.IsAbs(e.deployDir) {
|
||||
reasons = append(reasons, "部署目录必须是绝对路径")
|
||||
}
|
||||
if err := validateBackendContainerName(e.backendContainer); err != nil {
|
||||
reasons = append(reasons, err.Error())
|
||||
}
|
||||
if out, err := exec.CommandContext(ctx, "docker", "--version").CombinedOutput(); err != nil {
|
||||
reasons = append(reasons, fmt.Sprintf("docker CLI不可用: %v: %s", err, strings.TrimSpace(string(out))))
|
||||
}
|
||||
if info, err := os.Stat(dockerSocketPath); err != nil {
|
||||
reasons = append(reasons, "docker socket不可用: "+err.Error())
|
||||
} else if info.IsDir() {
|
||||
reasons = append(reasons, "docker socket路径不是文件")
|
||||
}
|
||||
if info, err := os.Stat(e.composePath()); err != nil {
|
||||
reasons = append(reasons, "部署docker-compose.yml不可用: "+err.Error())
|
||||
} else if info.IsDir() {
|
||||
reasons = append(reasons, "部署docker-compose.yml不是文件")
|
||||
}
|
||||
if info, err := os.Stat(e.envPath()); err != nil {
|
||||
reasons = append(reasons, "部署.env不可用: "+err.Error())
|
||||
} else if info.IsDir() {
|
||||
reasons = append(reasons, "部署.env不是文件")
|
||||
}
|
||||
if out, err := exec.CommandContext(ctx, "docker", "compose", "version").CombinedOutput(); err != nil {
|
||||
reasons = append(reasons, fmt.Sprintf("docker compose不可用: %v: %s", err, strings.TrimSpace(string(out))))
|
||||
}
|
||||
if _, err := e.currentBackendImage(ctx); err != nil {
|
||||
reasons = append(reasons, err.Error())
|
||||
}
|
||||
|
||||
return systemUpgradeCapabilityData{
|
||||
Capable: len(reasons) == 0,
|
||||
Reasons: reasons,
|
||||
DeployDir: e.deployDir,
|
||||
BackendContainer: e.backendContainer,
|
||||
}
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) selectComposeAsset(current []byte) string {
|
||||
if enableIPv6ComposePattern.Match(current) {
|
||||
return "docker-compose-v6.yml"
|
||||
}
|
||||
return "docker-compose-v4.yml"
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) helperScript() string {
|
||||
return `set -eu
|
||||
LOGFILE="$PANEL_DEPLOY_DIR/upgrade.log"
|
||||
log() { echo "[$(date '+%Y-%m-%d %H:%M:%S')] $*" | tee -a "$LOGFILE"; }
|
||||
|
||||
cd "$PANEL_DEPLOY_DIR"
|
||||
echo "" > "$LOGFILE"
|
||||
log "开始面板升级"
|
||||
log "工作目录: $(pwd)"
|
||||
|
||||
if [ ! -f docker-compose.yml ]; then
|
||||
log "错误: docker-compose.yml 不存在"
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -f .env ]; then
|
||||
log "错误: .env 不存在"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log "拉取新镜像..."
|
||||
if ! docker compose pull backend frontend >> "$LOGFILE" 2>&1; then
|
||||
log "错误: 拉取镜像失败"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log "等待旧容器释放资源..."
|
||||
sleep 3
|
||||
|
||||
log "重启服务(force-recreate)..."
|
||||
if ! docker compose up -d --force-recreate --remove-orphans backend frontend >> "$LOGFILE" 2>&1; then
|
||||
log "错误: 重启服务失败"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log "升级完成"
|
||||
`
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) buildHelperRunArgs(imageID, helperName string) ([]string, error) {
|
||||
if err := validateBackendContainerName(e.backendContainer); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{
|
||||
"run", "-d", "--rm", "--name", helperName,
|
||||
"--volumes-from", e.backendContainer,
|
||||
"-v", dockerSocketPath + ":" + dockerSocketPath,
|
||||
"-e", panelDeployDirEnv + "=" + e.deployDir,
|
||||
"--entrypoint", "/bin/sh", imageID,
|
||||
"-c", e.helperScript(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) updateEnvVersion(envPath, version string) error {
|
||||
if err := validateUpgradeVersion(version); err != nil {
|
||||
return err
|
||||
}
|
||||
mode, err := fileModeOrDefault(envPath, 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
lines := strings.Split(string(data), "\n")
|
||||
replaced := false
|
||||
for i, line := range lines {
|
||||
if strings.HasPrefix(line, "FLUX_VERSION=") {
|
||||
lines[i] = "FLUX_VERSION=" + version
|
||||
replaced = true
|
||||
}
|
||||
}
|
||||
if !replaced {
|
||||
trimmed := strings.TrimRight(strings.Join(lines, "\n"), "\n")
|
||||
if trimmed == "" {
|
||||
trimmed = "FLUX_VERSION=" + version
|
||||
} else {
|
||||
trimmed += "\nFLUX_VERSION=" + version
|
||||
}
|
||||
return writeFileWithMode(envPath, []byte(trimmed+"\n"), mode)
|
||||
}
|
||||
content := strings.TrimRight(strings.Join(lines, "\n"), "\n") + "\n"
|
||||
return writeFileWithMode(envPath, []byte(content), mode)
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) backupFile(path string) (string, error) {
|
||||
mode, err := fileModeOrDefault(path, 0o600)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
backupPath := path + ".upgrade.bak"
|
||||
if err := writeFileWithMode(backupPath, data, mode); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return backupPath, nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) restoreBackup(path string) error {
|
||||
backupPath := path + ".upgrade.bak"
|
||||
mode, err := fileModeOrDefault(backupPath, 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := os.ReadFile(backupPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeFileWithMode(path, data, mode)
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) restoreUpgradeBackups(paths ...string) error {
|
||||
var errs []string
|
||||
for _, path := range paths {
|
||||
if err := e.restoreBackup(path); err != nil {
|
||||
errs = append(errs, fmt.Sprintf("%s: %v", path, err))
|
||||
}
|
||||
}
|
||||
if len(errs) > 0 {
|
||||
return fmt.Errorf("%s", strings.Join(errs, "; "))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) replaceCompose(path string, data []byte) error {
|
||||
if len(bytes.TrimSpace(data)) == 0 {
|
||||
return fmt.Errorf("compose asset is empty")
|
||||
}
|
||||
mode, err := fileModeOrDefault(path, 0o644)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeFileWithMode(path, data, mode)
|
||||
}
|
||||
|
||||
func (h *Handler) buildSystemUpgradeDownloadURL(version, filename string) string {
|
||||
enabled, proxyURL := h.getGithubProxyConfig()
|
||||
base := fmt.Sprintf("%s/%s/releases/download/%s/%s", strings.TrimRight(systemUpgradeReleaseBaseURL, "/"), githubRepo, version, filename)
|
||||
if enabled {
|
||||
return fmt.Sprintf("%s/%s", proxyURL, base)
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
func (h *Handler) fetchSystemUpgradeReleases(perPage int) ([]githubRelease, error) {
|
||||
if perPage <= 0 {
|
||||
perPage = 20
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 15 * time.Second}
|
||||
url := fmt.Sprintf("%s/repos/%s/releases?per_page=%d", strings.TrimRight(systemUpgradeAPIBaseURL, "/"), githubRepo, perPage)
|
||||
if enabled, proxyURL := h.getGithubProxyConfig(); enabled {
|
||||
url = fmt.Sprintf("%s/%s", proxyURL, url)
|
||||
}
|
||||
|
||||
resp, err := systemUpgradeHTTPGet(client, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("请求GitHub API失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
return nil, fmt.Errorf("GitHub API返回 %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var releases []githubRelease
|
||||
if err := json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
return nil, fmt.Errorf("解析GitHub API响应失败: %v", err)
|
||||
}
|
||||
|
||||
return releases, nil
|
||||
}
|
||||
|
||||
func (h *Handler) resolveSystemUpgradeLatestReleaseByChannel(channel string) (string, error) {
|
||||
normalizedChannel := normalizeReleaseChannel(channel)
|
||||
releases, err := h.fetchSystemUpgradeReleases(50)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
for _, r := range releases {
|
||||
if r.Draft {
|
||||
continue
|
||||
}
|
||||
tag := strings.TrimSpace(r.TagName)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
if releaseChannelFromTag(tag) == normalizedChannel {
|
||||
return tag, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("未找到%s版本号", releaseChannelLabel(normalizedChannel))
|
||||
}
|
||||
|
||||
func fileModeOrDefault(path string, fallback os.FileMode) (os.FileMode, error) {
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return fallback, nil
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
return info.Mode().Perm(), nil
|
||||
}
|
||||
|
||||
func writeFileWithMode(path string, data []byte, mode os.FileMode) error {
|
||||
if err := os.WriteFile(path, data, mode); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Chmod(path, mode)
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) currentBackendImage(ctx context.Context) (string, error) {
|
||||
if err := validateBackendContainerName(e.backendContainer); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out, err := exec.CommandContext(ctx, "docker", "inspect", "-f", "{{.Image}}", e.backendContainer).CombinedOutput()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("inspect backend image failed: %v: %s", err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
imageID := strings.TrimSpace(string(out))
|
||||
if imageID == "" {
|
||||
return "", fmt.Errorf("backend image id is empty")
|
||||
}
|
||||
return imageID, nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) startHelper(ctx context.Context, imageID, helperName string) (string, error) {
|
||||
args, err := e.buildHelperRunArgs(imageID, helperName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
out, err := exec.CommandContext(ctx, "docker", args...).CombinedOutput()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("start helper failed: %v: %s", err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
containerID := strings.TrimSpace(string(out))
|
||||
if containerID == "" {
|
||||
containerID = helperName
|
||||
}
|
||||
return containerID, nil
|
||||
}
|
||||
|
||||
func (h *Handler) downloadReleaseAsset(version, filename string) ([]byte, error) {
|
||||
url := h.buildSystemUpgradeDownloadURL(version, filename)
|
||||
client := &http.Client{Timeout: 60 * time.Second}
|
||||
resp, err := systemUpgradeHTTPGet(client, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("下载%s失败: %v", filename, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
||||
return nil, fmt.Errorf("下载%s返回 %d: %s", filename, resp.StatusCode, strings.TrimSpace(string(body)))
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxSystemUpgradeComposeAssetBytes+1))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取%s失败: %v", filename, err)
|
||||
}
|
||||
if len(body) > maxSystemUpgradeComposeAssetBytes {
|
||||
return nil, fmt.Errorf("下载%s过大", filename)
|
||||
}
|
||||
if len(bytes.TrimSpace(body)) == 0 {
|
||||
return nil, fmt.Errorf("下载%s内容为空", filename)
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
func releasesForChannel(releases []githubRelease, channel string) []systemUpgradeReleaseData {
|
||||
channel = normalizeReleaseChannel(channel)
|
||||
items := make([]systemUpgradeReleaseData, 0, len(releases))
|
||||
for _, r := range releases {
|
||||
if r.Draft {
|
||||
continue
|
||||
}
|
||||
tag := strings.TrimSpace(r.TagName)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
itemChannel := releaseChannelFromTag(tag)
|
||||
if itemChannel != channel {
|
||||
continue
|
||||
}
|
||||
items = append(items, systemUpgradeReleaseData{
|
||||
Version: tag,
|
||||
Name: r.Name,
|
||||
PublishedAt: r.PublishedAt,
|
||||
Prerelease: itemChannel == releaseChannelDev,
|
||||
Channel: itemChannel,
|
||||
})
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func decodeSystemUpgradeRequest(r *http.Request, req *systemUpgradeRequest) error {
|
||||
defer r.Body.Close()
|
||||
body, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(bytes.TrimSpace(body)) == 0 {
|
||||
return nil
|
||||
}
|
||||
decoder := json.NewDecoder(bytes.NewReader(body))
|
||||
decoder.DisallowUnknownFields()
|
||||
return decoder.Decode(req)
|
||||
}
|
||||
|
||||
func systemUpgradeVersionResponse(current, channel, latest string, lookupErr error, capability systemUpgradeCapabilityData) systemUpgradeVersionData {
|
||||
data := systemUpgradeVersionData{
|
||||
CurrentVersion: current,
|
||||
LatestVersion: latest,
|
||||
HasUpdate: latest != "" && latest != current,
|
||||
Channel: channel,
|
||||
Capability: capability,
|
||||
}
|
||||
if lookupErr != nil {
|
||||
data.LatestVersion = ""
|
||||
data.HasUpdate = false
|
||||
data.Reason = lookupErr.Error()
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
func (h *Handler) systemVersion(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
channel := releaseChannelStable
|
||||
current := currentPanelVersion()
|
||||
exec := newSystemUpgradeExecutor()
|
||||
capability := exec.capability(r.Context())
|
||||
latest, err := h.resolveSystemUpgradeLatestReleaseByChannel(channel)
|
||||
response.WriteJSON(w, response.OK(systemUpgradeVersionResponse(current, channel, latest, err, capability)))
|
||||
}
|
||||
|
||||
func (h *Handler) systemCheckUpdates(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req systemUpgradeRequest
|
||||
if err := decodeSystemUpgradeRequest(r, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
channel := normalizeReleaseChannel(req.Channel)
|
||||
current := currentPanelVersion()
|
||||
exec := newSystemUpgradeExecutor()
|
||||
capability := exec.capability(r.Context())
|
||||
|
||||
githubReleases, err := h.fetchSystemUpgradeReleases(50)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
releases := releasesForChannel(githubReleases, channel)
|
||||
latest := ""
|
||||
if len(releases) > 0 {
|
||||
latest = releases[0].Version
|
||||
}
|
||||
response.WriteJSON(w, response.OK(systemUpgradeCheckData{
|
||||
CurrentVersion: current,
|
||||
LatestVersion: latest,
|
||||
HasUpdate: latest != "" && latest != current,
|
||||
Channel: channel,
|
||||
Capability: capability,
|
||||
Releases: releases,
|
||||
}))
|
||||
}
|
||||
|
||||
func (h *Handler) systemUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.systemUpgradeMu.TryLock() {
|
||||
response.WriteJSON(w, response.ErrDefault(systemUpgradeConflictError))
|
||||
return
|
||||
}
|
||||
defer h.systemUpgradeMu.Unlock()
|
||||
|
||||
var req systemUpgradeRequest
|
||||
if err := decodeSystemUpgradeRequest(r, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
channel := normalizeReleaseChannel(req.Channel)
|
||||
exec := newSystemUpgradeExecutor()
|
||||
capability := exec.capability(r.Context())
|
||||
if !capability.Capable {
|
||||
response.WriteJSON(w, response.ErrDefault("当前环境不支持面板自升级: "+strings.Join(capability.Reasons, "; ")))
|
||||
return
|
||||
}
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = h.resolveSystemUpgradeLatestReleaseByChannel(channel)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新%s失败: %v", releaseChannelLabel(channel), err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
imageID, err := exec.currentBackendImage(r.Context())
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
composePath := exec.composePath()
|
||||
envPath := exec.envPath()
|
||||
composeData, err := os.ReadFile(composePath)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "读取compose失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
composeAsset := exec.selectComposeAsset(composeData)
|
||||
newCompose, err := h.downloadReleaseAsset(version, composeAsset)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err := exec.backupFile(composePath); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "备份compose失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err := exec.backupFile(envPath); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "备份.env失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
if err := exec.replaceCompose(composePath, newCompose); err != nil {
|
||||
if restoreErr := exec.restoreUpgradeBackups(composePath, envPath); restoreErr != nil {
|
||||
err = fmt.Errorf("%v; 回滚失败: %v", err, restoreErr)
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, "替换compose失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
if err := exec.updateEnvVersion(envPath, version); err != nil {
|
||||
if restoreErr := exec.restoreUpgradeBackups(composePath, envPath); restoreErr != nil {
|
||||
err = fmt.Errorf("%v; 回滚失败: %v", err, restoreErr)
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, "更新版本配置失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
helperName := fmt.Sprintf("flvx-upgrade-helper-%d", time.Now().Unix())
|
||||
helperContainer, err := exec.startHelper(r.Context(), imageID, helperName)
|
||||
if err != nil {
|
||||
if restoreErr := exec.restoreUpgradeBackups(composePath, envPath); restoreErr != nil {
|
||||
err = fmt.Errorf("%v; 回滚失败: %v", err, restoreErr)
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(systemUpgradeRunData{
|
||||
Version: version,
|
||||
Channel: channel,
|
||||
ComposeAsset: composeAsset,
|
||||
HelperContainer: helperContainer,
|
||||
BackendImageID: imageID,
|
||||
Message: systemUpgradeMessage,
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,437 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestSelectComposeAssetUsesIPv6Template(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
compose := []byte("networks:\n gost-network:\n enable_ipv6: true\n")
|
||||
|
||||
if got := exec.selectComposeAsset(compose); got != "docker-compose-v6.yml" {
|
||||
t.Fatalf("selectComposeAsset() = %q, want %q", got, "docker-compose-v6.yml")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadReleaseAssetUsesGithubProxyWhenEnabled(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "test.db")
|
||||
repoStore, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("repo.Open() error = %v", err)
|
||||
}
|
||||
defer repoStore.Close()
|
||||
|
||||
h := &Handler{repo: repoStore}
|
||||
originalBase := systemUpgradeReleaseBaseURL
|
||||
systemUpgradeReleaseBaseURL = "https://example.invalid"
|
||||
t.Cleanup(func() { systemUpgradeReleaseBaseURL = originalBase })
|
||||
|
||||
originalGet := systemUpgradeHTTPGet
|
||||
defer func() { systemUpgradeHTTPGet = originalGet }()
|
||||
|
||||
var gotURL string
|
||||
systemUpgradeHTTPGet = func(client *http.Client, url string) (*http.Response, error) {
|
||||
gotURL = url
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader("services:\n backend:\n image: test\n")),
|
||||
}, nil
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := repoStore.UpsertConfig("github_proxy_enabled", "true", now); err != nil {
|
||||
t.Fatalf("UpsertConfig() github_proxy_enabled error = %v", err)
|
||||
}
|
||||
if err := repoStore.UpsertConfig("github_proxy_url", "https://proxy.example.com", now); err != nil {
|
||||
t.Fatalf("UpsertConfig() github_proxy_url error = %v", err)
|
||||
}
|
||||
|
||||
data, err := h.downloadReleaseAsset("2.1.9", "docker-compose-v4.yml")
|
||||
if err != nil {
|
||||
t.Fatalf("downloadReleaseAsset() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(string(data), "backend") {
|
||||
t.Fatalf("downloadReleaseAsset() data = %q, want compose data", string(data))
|
||||
}
|
||||
|
||||
wantURL := "https://proxy.example.com/https://example.invalid/Sagit-chu/flvx/releases/download/2.1.9/docker-compose-v4.yml"
|
||||
if gotURL != wantURL {
|
||||
t.Fatalf("download URL = %q, want %q", gotURL, wantURL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadReleaseAssetRejectsOversizedBody(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "test.db")
|
||||
repoStore, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("repo.Open() error = %v", err)
|
||||
}
|
||||
defer repoStore.Close()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := repoStore.UpsertConfig("github_proxy_enabled", "false", now); err != nil {
|
||||
t.Fatalf("UpsertConfig() github_proxy_enabled error = %v", err)
|
||||
}
|
||||
|
||||
originalBase := systemUpgradeReleaseBaseURL
|
||||
systemUpgradeReleaseBaseURL = "https://example.invalid"
|
||||
t.Cleanup(func() { systemUpgradeReleaseBaseURL = originalBase })
|
||||
|
||||
originalGet := systemUpgradeHTTPGet
|
||||
defer func() { systemUpgradeHTTPGet = originalGet }()
|
||||
systemUpgradeHTTPGet = func(client *http.Client, url string) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(bytes.NewReader(bytes.Repeat([]byte("a"), maxSystemUpgradeComposeAssetBytes+1))),
|
||||
}, nil
|
||||
}
|
||||
|
||||
h := &Handler{repo: repoStore}
|
||||
_, err = h.downloadReleaseAsset("2.1.9", "docker-compose-v4.yml")
|
||||
if err == nil || !strings.Contains(err.Error(), "过大") {
|
||||
t.Fatalf("downloadReleaseAsset() error = %v, want oversized error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectComposeAssetUsesIPv6TemplateForYAMLVariants(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
for _, compose := range [][]byte{
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6:true\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: True\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: \"true\"\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: 'true'\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: true # comment\n"),
|
||||
} {
|
||||
if got := exec.selectComposeAsset(compose); got != "docker-compose-v6.yml" {
|
||||
t.Fatalf("selectComposeAsset(%q) = %q, want %q", string(compose), got, "docker-compose-v6.yml")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectComposeAssetFallsBackToIPv4Template(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
compose := []byte("services:\n backend:\n image: test\n")
|
||||
|
||||
if got := exec.selectComposeAsset(compose); got != "docker-compose-v4.yml" {
|
||||
t.Fatalf("selectComposeAsset() = %q, want %q", got, "docker-compose-v4.yml")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionReplacesExistingValue(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, "2.1.9"); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
|
||||
want := "FLUX_VERSION=2.1.9\nJWT_SECRET=test\n"
|
||||
if string(data) != want {
|
||||
t.Fatalf("env content = %q, want %q", string(data), want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionAppendsMissingValue(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("JWT_SECRET=test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, "2.1.9"); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
|
||||
want := "JWT_SECRET=test\nFLUX_VERSION=2.1.9\n"
|
||||
if string(data) != want {
|
||||
t.Fatalf("env content = %q, want %q", string(data), want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionRejectsUnsafeValue(t *testing.T) {
|
||||
for _, version := range []string{"", "2.1.9\nJWT_SECRET=bad", "2.1.9\rbad", "2.1.9\x00bad", "2.1.9\x1fbad"} {
|
||||
t.Run(version, func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
original := []byte("JWT_SECRET=test\n")
|
||||
if err := os.WriteFile(envPath, original, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, version); err == nil {
|
||||
t.Fatal("expected unsafe version to fail validation")
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if string(data) != string(original) {
|
||||
t.Fatalf("env content changed to %q, want %q", string(data), string(original))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionAcceptsVersionLabels(t *testing.T) {
|
||||
for _, version := range []string{"2.1.9", "2.1.9-beta14", "v-test"} {
|
||||
t.Run(version, func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("JWT_SECRET=test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, version); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionPreservesFileMode(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, "2.1.9"); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
|
||||
info, err := os.Stat(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat() error = %v", err)
|
||||
}
|
||||
if got := info.Mode().Perm(); got != 0o600 {
|
||||
t.Fatalf("env mode = %o, want 0600", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateBackendContainerNameRejectsUnsafeValue(t *testing.T) {
|
||||
if err := validateBackendContainerName("flux-panel-backend;rm -rf /"); err == nil {
|
||||
t.Fatal("expected unsafe container name to fail validation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildHelperRunArgsUsesDetachedContainer(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
args, err := exec.buildHelperRunArgs("sha256:abc", "flvx-upgrade-helper")
|
||||
if err != nil {
|
||||
t.Fatalf("buildHelperRunArgs() error = %v", err)
|
||||
}
|
||||
want := []string{
|
||||
"run", "-d", "--rm", "--name", "flvx-upgrade-helper",
|
||||
"--volumes-from", "flux-panel-backend",
|
||||
"-v", "/var/run/docker.sock:/var/run/docker.sock",
|
||||
"-e", "PANEL_DEPLOY_DIR=/opt/flvx-panel",
|
||||
"--entrypoint", "/bin/sh", "sha256:abc",
|
||||
"-c", exec.helperScript(),
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(args, want) {
|
||||
t.Fatalf("buildHelperRunArgs() = %#v, want %#v", args, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildHelperRunArgsRejectsUnsafeBackendContainer(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend;rm -rf /"}
|
||||
if _, err := exec.buildHelperRunArgs("sha256:abc", "flvx-upgrade-helper"); err == nil {
|
||||
t.Fatal("expected unsafe backend container name to fail validation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemVersionRejectsWrongMethod(t *testing.T) {
|
||||
h := &Handler{}
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/system/version", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
h.systemVersion(rr, req)
|
||||
|
||||
if !strings.Contains(rr.Body.String(), "请求失败") {
|
||||
t.Fatalf("expected wrong-method response, got %s", rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemUpgradeRejectsConcurrentRequests(t *testing.T) {
|
||||
h := &Handler{}
|
||||
h.systemUpgradeMu.Lock()
|
||||
defer h.systemUpgradeMu.Unlock()
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/upgrade", strings.NewReader(`{"channel":"stable"}`))
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
h.systemUpgrade(rr, req)
|
||||
|
||||
if !strings.Contains(rr.Body.String(), systemUpgradeConflictError) {
|
||||
t.Fatalf("expected conflict message, got %s", rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemUpgradeFailsFastBeforeMutatingFiles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
composePath := filepath.Join(dir, "docker-compose.yml")
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(composePath, []byte("services:\n backend:\n image: test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() compose error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(envPath, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() env error = %v", err)
|
||||
}
|
||||
|
||||
fakeDockerDir := t.TempDir()
|
||||
fakeDockerPath := filepath.Join(fakeDockerDir, "docker")
|
||||
fakeDockerScript := "#!/bin/sh\ncase \"$1\" in\n --version)\n echo 'Docker version 27.0.0'\n exit 0\n ;;&\n compose)\n if [ \"$2\" = version ]; then\n echo 'Docker Compose version v2.33.0'\n exit 0\n fi\n exit 0\n ;;&\n inspect)\n echo 'No such object: flux-panel-backend' >&2\n exit 1\n ;;&\n *)\n exit 0\n ;;&\n esac\n"
|
||||
if err := os.WriteFile(fakeDockerPath, []byte(fakeDockerScript), 0o755); err != nil {
|
||||
t.Fatalf("WriteFile() fake docker error = %v", err)
|
||||
}
|
||||
t.Setenv("PATH", fakeDockerDir+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
t.Setenv(panelDeployDirEnv, dir)
|
||||
t.Setenv(panelBackendContainerEnv, "flux-panel-backend")
|
||||
|
||||
h := &Handler{}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/upgrade", strings.NewReader(`{"channel":"stable"}`))
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
h.systemUpgrade(rr, req)
|
||||
|
||||
if !strings.Contains(rr.Body.String(), "当前环境不支持面板自升级") {
|
||||
t.Fatalf("expected fail-fast capability error, got %s", rr.Body.String())
|
||||
}
|
||||
if _, err := os.Stat(composePath + ".upgrade.bak"); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected no compose backup, got err=%v", err)
|
||||
}
|
||||
if _, err := os.Stat(envPath + ".upgrade.bak"); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected no env backup, got err=%v", err)
|
||||
}
|
||||
composeData, err := os.ReadFile(composePath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() compose error = %v", err)
|
||||
}
|
||||
if string(composeData) != "services:\n backend:\n image: test\n" {
|
||||
t.Fatalf("compose mutated unexpectedly: %q", string(composeData))
|
||||
}
|
||||
envData, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() env error = %v", err)
|
||||
}
|
||||
if string(envData) != "FLUX_VERSION=2.1.8\nJWT_SECRET=test\n" {
|
||||
t.Fatalf("env mutated unexpectedly: %q", string(envData))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpgradeBackupUsesStablePathAndRestoreRestoresOriginal(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "docker-compose.yml")
|
||||
if err := os.WriteFile(path, []byte("original"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
backupPath, err := exec.backupFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("backupFile() error = %v", err)
|
||||
}
|
||||
if backupPath != path+".upgrade.bak" {
|
||||
t.Fatalf("backup path = %q, want %q", backupPath, path+".upgrade.bak")
|
||||
}
|
||||
if err := os.WriteFile(path, []byte("mutated"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
if err := exec.restoreBackup(path); err != nil {
|
||||
t.Fatalf("restoreBackup() error = %v", err)
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if string(data) != "original" {
|
||||
t.Fatalf("restored content = %q, want original", string(data))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRestoreBackupPreservesOriginalFileMode(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(path, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if _, err := exec.backupFile(path); err != nil {
|
||||
t.Fatalf("backupFile() error = %v", err)
|
||||
}
|
||||
if err := os.Remove(path); err != nil {
|
||||
t.Fatalf("Remove() error = %v", err)
|
||||
}
|
||||
if err := exec.restoreBackup(path); err != nil {
|
||||
t.Fatalf("restoreBackup() error = %v", err)
|
||||
}
|
||||
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat() error = %v", err)
|
||||
}
|
||||
if got := info.Mode().Perm(); got != 0o600 {
|
||||
t.Fatalf("restored mode = %o, want 0600", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSystemUpgradeRequestRejectsTruncatedJSON(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/check-updates", strings.NewReader(`{"channel":"stable"`))
|
||||
var payload systemUpgradeRequest
|
||||
|
||||
if err := decodeSystemUpgradeRequest(req, &payload); err == nil {
|
||||
t.Fatal("expected truncated JSON to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSystemUpgradeRequestAllowsEmptyBody(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/check-updates", strings.NewReader(""))
|
||||
var payload systemUpgradeRequest
|
||||
|
||||
if err := decodeSystemUpgradeRequest(req, &payload); err != nil {
|
||||
t.Fatalf("expected empty body to be accepted, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemUpgradeVersionDataSurfacesLookupFailureReason(t *testing.T) {
|
||||
data, err := json.Marshal(systemUpgradeVersionData{Reason: "GitHub unavailable"})
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(string(data), `"reason":"GitHub unavailable"`) {
|
||||
t.Fatalf("expected reason field in JSON, got %s", string(data))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// flowMiB is optional so older clients can keep sending the GB-based flow field.
|
||||
// A positive value takes precedence and preserves sub-GB limits exactly.
|
||||
func parseTrafficLimit(req map[string]interface{}, defaultGB int64) (flowGB, flowMiB int64, err error) {
|
||||
flowGB = asInt64(req["flow"], defaultGB)
|
||||
if flowGB < 0 {
|
||||
return 0, 0, fmt.Errorf("流量限制不能小于0")
|
||||
}
|
||||
raw, present := req["flowMiB"]
|
||||
if !present {
|
||||
return flowGB, 0, nil
|
||||
}
|
||||
flowMiB, err = strconv.ParseInt(asString(raw), 10, 64)
|
||||
if err != nil || flowMiB < 0 || flowMiB > math.MaxInt64/bytesPerMiB {
|
||||
return 0, 0, fmt.Errorf("流量限制超出范围")
|
||||
}
|
||||
if flowMiB > 0 {
|
||||
flowGB = (flowMiB-1)/1024 + 1
|
||||
}
|
||||
return flowGB, flowMiB, nil
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestTrafficLimitMiBOverridesLegacyGB(t *testing.T) {
|
||||
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{
|
||||
"flow": float64(1), "flowMiB": float64(500),
|
||||
}, 100)
|
||||
if err != nil || flowGB != 1 || flowMiB != 500 {
|
||||
t.Fatalf("parseTrafficLimit = (%d, %d, %v), want (1, 500, nil)", flowGB, flowMiB, err)
|
||||
}
|
||||
limit := flowLimitBytes(flowGB, flowMiB)
|
||||
if limit != 500*bytesPerMiB {
|
||||
t.Fatalf("limit = %d, want %d", limit, 500*bytesPerMiB)
|
||||
}
|
||||
policy := &userTunnelPolicy{Flow: flowGB, FlowMiB: flowMiB, InFlow: limit - 1, Status: 1}
|
||||
if shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
|
||||
t.Fatal("policy paused before reaching 500 MiB")
|
||||
}
|
||||
policy.InFlow = limit
|
||||
if !shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
|
||||
t.Fatal("policy did not pause at 500 MiB")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrafficLimitLegacyAndInvalidValues(t *testing.T) {
|
||||
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{"flow": float64(2)}, 100)
|
||||
if err != nil || flowGB != 2 || flowMiB != 0 || flowLimitBytes(flowGB, flowMiB) != 2*bytesPerGB {
|
||||
t.Fatalf("legacy GB limit changed: (%d, %d, %v)", flowGB, flowMiB, err)
|
||||
}
|
||||
for _, value := range []interface{}{"1.5", -1, "999999999999999999999"} {
|
||||
if _, _, err := parseTrafficLimit(map[string]interface{}{"flowMiB": value}, 100); err == nil {
|
||||
t.Fatalf("accepted invalid flowMiB %v", value)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,440 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
tunnelStrategyBest = "best"
|
||||
bestExitRuntimeStrategy = "fifo"
|
||||
bestExitPublicTargetHost = "www.bing.com"
|
||||
bestExitPublicTargetPort = 443
|
||||
bestExitLossPenaltyMsPerPercent = 100.0
|
||||
bestExitConfirmationRounds = 3
|
||||
bestExitSwitchCooldown = 30 * time.Second
|
||||
bestExitApplyRetryCooldown = bestExitSwitchCooldown
|
||||
bestExitMinLatencyAdvantageMs = 20.0
|
||||
bestExitMinScoreAdvantageRatio = 0.15
|
||||
)
|
||||
|
||||
type bestExitOwnerKey struct {
|
||||
TunnelID int64
|
||||
OwnerNodeID int64
|
||||
}
|
||||
|
||||
type bestExitCandidateScore struct {
|
||||
OwnerNodeID int64
|
||||
ExitNodeID int64
|
||||
ExitName string
|
||||
|
||||
OwnerToExitLatency float64
|
||||
ExitToBingLatency float64
|
||||
OwnerToExitLoss float64
|
||||
ExitToBingLoss float64
|
||||
TotalLatency float64
|
||||
TotalLoss float64
|
||||
Score float64
|
||||
Success bool
|
||||
ErrorMessage string
|
||||
}
|
||||
|
||||
type bestExitSwitchDecision struct {
|
||||
Switch bool
|
||||
ExitNodeID int64
|
||||
Reason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitProbeFunc func(nodeID int64, ip string, port int, options diagnosisExecOptions) (latency float64, loss float64, err error)
|
||||
|
||||
type bestExitProbeResult struct {
|
||||
latency float64
|
||||
loss float64
|
||||
err error
|
||||
}
|
||||
|
||||
type bestExitProbeCacheKey struct {
|
||||
NodeID int64
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
|
||||
type bestExitDecision struct {
|
||||
AppliedExitNodeID int64
|
||||
PendingExitNodeID int64
|
||||
PendingCount int
|
||||
LastSwitchAt time.Time
|
||||
LastApplyFailureAt time.Time
|
||||
LastApplyFailureExitNodeID int64
|
||||
LastReason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitManager struct {
|
||||
mu sync.Mutex
|
||||
decisions map[bestExitOwnerKey]*bestExitDecision
|
||||
}
|
||||
|
||||
func newBestExitManager() *bestExitManager {
|
||||
return &bestExitManager{decisions: make(map[bestExitOwnerKey]*bestExitDecision)}
|
||||
}
|
||||
|
||||
func isBestTunnelStrategy(strategy string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(strategy), tunnelStrategyBest)
|
||||
}
|
||||
|
||||
func runtimeTunnelStrategy(strategy string) string {
|
||||
if isBestTunnelStrategy(strategy) {
|
||||
return bestExitRuntimeStrategy
|
||||
}
|
||||
return strategy
|
||||
}
|
||||
|
||||
func scoreBestExitCandidate(ownerNodeID int64, exit chainNodeRecord, ownerLatency, ownerLoss, publicLatency, publicLoss float64) bestExitCandidateScore {
|
||||
totalLatency := ownerLatency + publicLatency
|
||||
totalLoss := combineLossPercent(ownerLoss, publicLoss)
|
||||
return bestExitCandidateScore{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
ExitNodeID: exit.NodeID,
|
||||
ExitName: exit.NodeName,
|
||||
OwnerToExitLatency: ownerLatency,
|
||||
ExitToBingLatency: publicLatency,
|
||||
OwnerToExitLoss: ownerLoss,
|
||||
ExitToBingLoss: publicLoss,
|
||||
TotalLatency: totalLatency,
|
||||
TotalLoss: totalLoss,
|
||||
Score: totalLatency + totalLoss*bestExitLossPenaltyMsPerPercent,
|
||||
Success: true,
|
||||
}
|
||||
}
|
||||
|
||||
func failedBestExitCandidate(ownerNodeID int64, exit chainNodeRecord, message string) bestExitCandidateScore {
|
||||
return bestExitCandidateScore{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
ExitNodeID: exit.NodeID,
|
||||
ExitName: exit.NodeName,
|
||||
Success: false,
|
||||
ErrorMessage: message,
|
||||
}
|
||||
}
|
||||
|
||||
func combineLossPercent(a, b float64) float64 {
|
||||
a = clampPercent(a)
|
||||
b = clampPercent(b)
|
||||
return (1 - (1-a/100.0)*(1-b/100.0)) * 100.0
|
||||
}
|
||||
|
||||
func clampPercent(v float64) float64 {
|
||||
if v < 0 {
|
||||
return 0
|
||||
}
|
||||
if v > 100 {
|
||||
return 100
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func sortBestExitScores(scores []bestExitCandidateScore) {
|
||||
sort.SliceStable(scores, func(i, j int) bool {
|
||||
return bestExitScoreLess(scores[i], scores[j])
|
||||
})
|
||||
}
|
||||
|
||||
func evaluateBestExitOwner(owner chainNodeRecord, exits []chainNodeRecord, nodes map[int64]*nodeRecord, ipPreference string, options diagnosisExecOptions, target tunnelProbeTarget, ping bestExitProbeFunc) []bestExitCandidateScore {
|
||||
scores := make([]bestExitCandidateScore, 0, len(exits))
|
||||
if owner.NodeID <= 0 || len(exits) == 0 || ping == nil {
|
||||
return scores
|
||||
}
|
||||
ownerNode := nodes[owner.NodeID]
|
||||
for _, exit := range exits {
|
||||
exitNode := nodes[exit.NodeID]
|
||||
if !isTunnelProbeNodeOnline(ownerNode) {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, "owner node offline"))
|
||||
continue
|
||||
}
|
||||
if !isTunnelProbeNodeOnline(exitNode) {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, "exit node unavailable"))
|
||||
continue
|
||||
}
|
||||
targetIP, targetPort, resolveErr := resolveBestExitProbeTarget(ownerNode, exitNode, exit.Port, ipPreference, exit.ConnectIP)
|
||||
if resolveErr != nil {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, resolveErr.Error()))
|
||||
continue
|
||||
}
|
||||
ownerLatency, ownerLoss, ownerErr := ping(owner.NodeID, targetIP, targetPort, options)
|
||||
if ownerErr != nil {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, ownerErr.Error()))
|
||||
continue
|
||||
}
|
||||
publicLatency, publicLoss, publicErr := ping(exit.NodeID, target.Host, target.Port, options)
|
||||
if publicErr != nil {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, publicErr.Error()))
|
||||
continue
|
||||
}
|
||||
scores = append(scores, scoreBestExitCandidate(owner.NodeID, exit, ownerLatency, ownerLoss, publicLatency, publicLoss))
|
||||
}
|
||||
sortBestExitScores(scores)
|
||||
return scores
|
||||
}
|
||||
|
||||
func resolveBestExitProbeTarget(fromNode, targetNode *nodeRecord, preferredPort int, ipPreference string, connectIP string) (string, int, error) {
|
||||
if targetNode == nil {
|
||||
return "", 0, errors.New("目标节点不存在")
|
||||
}
|
||||
host, err := selectTunnelDialHost(fromNode, targetNode, ipPreference, connectIP)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
if strings.TrimSpace(host) == "" {
|
||||
return "", 0, errors.New("目标节点地址为空")
|
||||
}
|
||||
port := preferredPort
|
||||
if port <= 0 {
|
||||
port = firstPortFromRange(targetNode.PortRange)
|
||||
}
|
||||
if port <= 0 {
|
||||
port = 443
|
||||
}
|
||||
return host, port, nil
|
||||
}
|
||||
|
||||
func newBestExitRoundPinger(base bestExitProbeFunc) bestExitProbeFunc {
|
||||
cache := make(map[bestExitProbeCacheKey]bestExitProbeResult)
|
||||
return func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
key := bestExitProbeCacheKey{NodeID: nodeID, Host: ip, Port: port}
|
||||
if cached, ok := cache[key]; ok {
|
||||
return cached.latency, cached.loss, cached.err
|
||||
}
|
||||
lat, loss, err := base(nodeID, ip, port, options)
|
||||
cache[key] = bestExitProbeResult{latency: lat, loss: loss, err: err}
|
||||
return lat, loss, err
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitChainOwners(inNodes []chainNodeRecord, chainHops [][]chainNodeRecord) []chainNodeRecord {
|
||||
if len(chainHops) == 0 {
|
||||
return inNodes
|
||||
}
|
||||
return chainHops[len(chainHops)-1]
|
||||
}
|
||||
|
||||
func chainRecordsToRuntimeTargets(rows []chainNodeRecord) []tunnelRuntimeNode {
|
||||
out := make([]tunnelRuntimeNode, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
out = append(out, tunnelRuntimeNode{
|
||||
NodeID: row.NodeID,
|
||||
Protocol: row.Protocol,
|
||||
Strategy: row.Strategy,
|
||||
Inx: int(row.Inx),
|
||||
ChainType: row.ChainType,
|
||||
Port: row.Port,
|
||||
ConnectIP: row.ConnectIP,
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func orderRuntimeTargetsByNodeID(targets []tunnelRuntimeNode, orderedIDs []int64) []tunnelRuntimeNode {
|
||||
out := append([]tunnelRuntimeNode(nil), targets...)
|
||||
if len(out) <= 1 || len(orderedIDs) == 0 {
|
||||
return out
|
||||
}
|
||||
positions := make(map[int64]int, len(orderedIDs))
|
||||
for i, id := range orderedIDs {
|
||||
if _, ok := positions[id]; !ok {
|
||||
positions[id] = i
|
||||
}
|
||||
}
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
pi, iok := positions[out[i].NodeID]
|
||||
pj, jok := positions[out[j].NodeID]
|
||||
if iok != jok {
|
||||
return iok
|
||||
}
|
||||
if iok && jok && pi != pj {
|
||||
return pi < pj
|
||||
}
|
||||
return false
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
func cloneBestExitScores(scores []bestExitCandidateScore) []bestExitCandidateScore {
|
||||
return append([]bestExitCandidateScore(nil), scores...)
|
||||
}
|
||||
|
||||
func bestExitDecisionResult(switchNow bool, exitNodeID int64, reason string, scores []bestExitCandidateScore) bestExitSwitchDecision {
|
||||
return bestExitSwitchDecision{Switch: switchNow, ExitNodeID: exitNodeID, Reason: reason, Scores: cloneBestExitScores(scores)}
|
||||
}
|
||||
|
||||
func bestExitScoreLess(a, b bestExitCandidateScore) bool {
|
||||
if a.Success != b.Success {
|
||||
return a.Success
|
||||
}
|
||||
if !a.Success && !b.Success {
|
||||
return a.ExitNodeID < b.ExitNodeID
|
||||
}
|
||||
if a.Score != b.Score {
|
||||
return a.Score < b.Score
|
||||
}
|
||||
return a.ExitNodeID < b.ExitNodeID
|
||||
}
|
||||
|
||||
func bestExitHasMinimumAdvantage(candidate, current bestExitCandidateScore) bool {
|
||||
if !candidate.Success {
|
||||
return false
|
||||
}
|
||||
if !current.Success {
|
||||
return true
|
||||
}
|
||||
improvement := current.Score - candidate.Score
|
||||
threshold := current.Score * bestExitMinScoreAdvantageRatio
|
||||
if threshold < bestExitMinLatencyAdvantageMs {
|
||||
threshold = bestExitMinLatencyAdvantageMs
|
||||
}
|
||||
return improvement >= threshold
|
||||
}
|
||||
|
||||
func (m *bestExitManager) setApplied(key bestExitOwnerKey, exitNodeID int64, at time.Time) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisionLocked(key)
|
||||
d.AppliedExitNodeID = exitNodeID
|
||||
d.PendingExitNodeID = 0
|
||||
d.PendingCount = 0
|
||||
d.LastApplyFailureAt = time.Time{}
|
||||
d.LastApplyFailureExitNodeID = 0
|
||||
d.LastSwitchAt = at
|
||||
}
|
||||
|
||||
func (m *bestExitManager) recordApplyFailure(key bestExitOwnerKey, exitNodeID int64, at time.Time) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisionLocked(key)
|
||||
d.LastApplyFailureAt = at
|
||||
d.LastApplyFailureExitNodeID = exitNodeID
|
||||
d.LastReason = "apply retry cooldown"
|
||||
}
|
||||
|
||||
func (m *bestExitManager) ensureApplied(key bestExitOwnerKey, exitNodeID int64, at time.Time) {
|
||||
if m == nil || exitNodeID <= 0 {
|
||||
return
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisionLocked(key)
|
||||
if d.AppliedExitNodeID == 0 {
|
||||
d.AppliedExitNodeID = exitNodeID
|
||||
d.LastSwitchAt = at
|
||||
}
|
||||
}
|
||||
|
||||
func (m *bestExitManager) observeScores(key bestExitOwnerKey, scores []bestExitCandidateScore, now time.Time) bestExitSwitchDecision {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
ordered := append([]bestExitCandidateScore(nil), scores...)
|
||||
sortBestExitScores(ordered)
|
||||
d := m.decisionLocked(key)
|
||||
d.Scores = cloneBestExitScores(ordered)
|
||||
|
||||
if len(ordered) == 0 || !ordered[0].Success {
|
||||
d.LastReason = "all exits failed"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
candidate := ordered[0]
|
||||
if d.AppliedExitNodeID == 0 {
|
||||
d.AppliedExitNodeID = candidate.ExitNodeID
|
||||
d.LastSwitchAt = now
|
||||
d.LastReason = "initial best exit"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
if candidate.ExitNodeID == d.AppliedExitNodeID {
|
||||
d.PendingExitNodeID = 0
|
||||
d.PendingCount = 0
|
||||
d.LastApplyFailureAt = time.Time{}
|
||||
d.LastApplyFailureExitNodeID = 0
|
||||
d.LastReason = "current exit remains best"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
if candidate.ExitNodeID == d.LastApplyFailureExitNodeID && !d.LastApplyFailureAt.IsZero() && now.Sub(d.LastApplyFailureAt) < bestExitApplyRetryCooldown {
|
||||
d.LastReason = "apply retry cooldown"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
if now.Sub(d.LastSwitchAt) < bestExitSwitchCooldown {
|
||||
d.LastReason = "cooldown"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
current := findBestExitScore(ordered, d.AppliedExitNodeID)
|
||||
if !bestExitHasMinimumAdvantage(candidate, current) {
|
||||
d.PendingExitNodeID = 0
|
||||
d.PendingCount = 0
|
||||
d.LastReason = "insufficient advantage"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
if d.PendingExitNodeID != candidate.ExitNodeID {
|
||||
d.PendingExitNodeID = candidate.ExitNodeID
|
||||
d.PendingCount = 1
|
||||
d.LastReason = "candidate pending confirmation"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
d.PendingCount++
|
||||
if d.PendingCount < bestExitConfirmationRounds {
|
||||
d.LastReason = "candidate pending confirmation"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
d.LastReason = "switch confirmed"
|
||||
return bestExitDecisionResult(true, candidate.ExitNodeID, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
func findBestExitScore(scores []bestExitCandidateScore, exitNodeID int64) bestExitCandidateScore {
|
||||
for _, score := range scores {
|
||||
if score.ExitNodeID == exitNodeID {
|
||||
return score
|
||||
}
|
||||
}
|
||||
return failedBestExitCandidate(0, chainNodeRecord{NodeID: exitNodeID}, "current exit has no successful score")
|
||||
}
|
||||
|
||||
func (m *bestExitManager) decisionLocked(key bestExitOwnerKey) *bestExitDecision {
|
||||
if d := m.decisions[key]; d != nil {
|
||||
return d
|
||||
}
|
||||
d := &bestExitDecision{}
|
||||
m.decisions[key] = d
|
||||
return d
|
||||
}
|
||||
|
||||
func (m *bestExitManager) orderTargets(key bestExitOwnerKey, targets []tunnelRuntimeNode) []tunnelRuntimeNode {
|
||||
out := append([]tunnelRuntimeNode(nil), targets...)
|
||||
if m == nil || len(out) <= 1 {
|
||||
return out
|
||||
}
|
||||
m.mu.Lock()
|
||||
applied := int64(0)
|
||||
if d := m.decisions[key]; d != nil {
|
||||
applied = d.AppliedExitNodeID
|
||||
}
|
||||
m.mu.Unlock()
|
||||
if applied <= 0 {
|
||||
return out
|
||||
}
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
if out[i].NodeID == applied {
|
||||
return true
|
||||
}
|
||||
if out[j].NodeID == applied {
|
||||
return false
|
||||
}
|
||||
return false
|
||||
})
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
bestExitDisplayStatusApplied = "applied"
|
||||
bestExitDisplayStatusWaiting = "waiting"
|
||||
bestExitDisplaySummaryMulti = "多个出口"
|
||||
bestExitDisplaySummaryWait = "等待探测"
|
||||
bestExitUnknownExitName = "未知出口"
|
||||
bestExitUnknownEntryName = "未知入口"
|
||||
bestExitUnknownChainName = "未知中转"
|
||||
)
|
||||
|
||||
type bestExitDecisionSnapshot struct {
|
||||
AppliedExitNodeID int64
|
||||
UpdatedAt int64
|
||||
Reason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitDisplayState struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Summary string `json:"summary"`
|
||||
Status string `json:"status"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
Items []bestExitDisplayItem `json:"items"`
|
||||
}
|
||||
|
||||
type bestExitDisplayItem struct {
|
||||
OwnerNodeID int64 `json:"ownerNodeId"`
|
||||
OwnerNodeName string `json:"ownerNodeName"`
|
||||
OwnerRole string `json:"ownerRole"`
|
||||
ExitNodeID int64 `json:"exitNodeId,omitempty"`
|
||||
ExitNodeName string `json:"exitNodeName"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
}
|
||||
|
||||
type bestExitNodeNameLookup func(nodeID int64) (string, bool)
|
||||
|
||||
func (m *bestExitManager) snapshot(key bestExitOwnerKey) (bestExitDecisionSnapshot, bool) {
|
||||
if m == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisions[key]
|
||||
if d == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
updatedAt := int64(0)
|
||||
if !d.LastSwitchAt.IsZero() {
|
||||
updatedAt = d.LastSwitchAt.UnixMilli()
|
||||
}
|
||||
return bestExitDecisionSnapshot{
|
||||
AppliedExitNodeID: d.AppliedExitNodeID,
|
||||
UpdatedAt: updatedAt,
|
||||
Reason: d.LastReason,
|
||||
Scores: cloneBestExitScores(d.Scores),
|
||||
}, true
|
||||
}
|
||||
|
||||
func (h *Handler) attachBestExitStates(items []map[string]interface{}) {
|
||||
if h == nil || len(items) == 0 {
|
||||
return
|
||||
}
|
||||
lookup := h.bestExitNodeNameLookup()
|
||||
for _, item := range items {
|
||||
state, ok := buildBestExitDisplayState(item, h.bestExit, lookup)
|
||||
if !ok {
|
||||
delete(item, "bestExitState")
|
||||
continue
|
||||
}
|
||||
item["bestExitState"] = state
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) bestExitNodeNameLookup() bestExitNodeNameLookup {
|
||||
cache := map[int64]string{}
|
||||
return func(nodeID int64) (string, bool) {
|
||||
if nodeID <= 0 || h == nil {
|
||||
return "", false
|
||||
}
|
||||
if name, ok := cache[nodeID]; ok {
|
||||
return name, name != ""
|
||||
}
|
||||
node, err := h.getNodeRecord(nodeID)
|
||||
if err != nil || node == nil {
|
||||
cache[nodeID] = ""
|
||||
return "", false
|
||||
}
|
||||
name := strings.TrimSpace(node.Name)
|
||||
cache[nodeID] = name
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
|
||||
func buildBestExitDisplayState(tunnel map[string]interface{}, manager *bestExitManager, lookup bestExitNodeNameLookup) (*bestExitDisplayState, bool) {
|
||||
if tunnel == nil {
|
||||
return nil, false
|
||||
}
|
||||
tunnelID := asInt64(tunnel["id"], 0)
|
||||
outNodes := bestExitDisplayMapSlice(tunnel["outNodeId"])
|
||||
if tunnelID <= 0 || len(outNodes) <= 1 {
|
||||
return nil, false
|
||||
}
|
||||
if !isBestTunnelStrategy(asString(outNodes[0]["strategy"])) {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
owners, ownerRole := bestExitDisplayOwners(tunnel)
|
||||
state := &bestExitDisplayState{
|
||||
Enabled: true,
|
||||
Summary: bestExitDisplaySummaryWait,
|
||||
Status: bestExitDisplayStatusWaiting,
|
||||
Items: make([]bestExitDisplayItem, 0, len(owners)),
|
||||
}
|
||||
|
||||
exitsByID := map[int64]map[string]interface{}{}
|
||||
for _, exit := range outNodes {
|
||||
if id := asInt64(exit["nodeId"], 0); id > 0 {
|
||||
exitsByID[id] = exit
|
||||
}
|
||||
}
|
||||
appliedExitIDs := map[int64]string{}
|
||||
appliedCount := 0
|
||||
latestUpdatedAt := int64(0)
|
||||
latestReason := ""
|
||||
for _, owner := range owners {
|
||||
ownerNodeID := asInt64(owner["nodeId"], 0)
|
||||
if ownerNodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
item := bestExitDisplayItem{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
OwnerNodeName: bestExitDisplayNodeName(owner, ownerNodeID, lookup, bestExitUnknownOwnerName(ownerRole)),
|
||||
OwnerRole: ownerRole,
|
||||
ExitNodeName: bestExitDisplaySummaryWait,
|
||||
Reason: bestExitDisplayStatusWaiting,
|
||||
}
|
||||
if snapshot, ok := manager.snapshot(bestExitOwnerKey{TunnelID: tunnelID, OwnerNodeID: ownerNodeID}); ok && snapshot.AppliedExitNodeID > 0 {
|
||||
exit, ok := exitsByID[snapshot.AppliedExitNodeID]
|
||||
if !ok {
|
||||
state.Items = append(state.Items, item)
|
||||
continue
|
||||
}
|
||||
item.ExitNodeID = snapshot.AppliedExitNodeID
|
||||
item.ExitNodeName = bestExitDisplayNodeName(exit, snapshot.AppliedExitNodeID, lookup, bestExitUnknownExitName)
|
||||
item.UpdatedAt = snapshot.UpdatedAt
|
||||
item.Reason = snapshot.Reason
|
||||
appliedExitIDs[item.ExitNodeID] = item.ExitNodeName
|
||||
appliedCount++
|
||||
if snapshot.UpdatedAt > latestUpdatedAt {
|
||||
latestUpdatedAt = snapshot.UpdatedAt
|
||||
latestReason = snapshot.Reason
|
||||
}
|
||||
}
|
||||
state.Items = append(state.Items, item)
|
||||
}
|
||||
|
||||
if appliedCount == 0 {
|
||||
return state, true
|
||||
}
|
||||
if appliedCount < len(state.Items) {
|
||||
return state, true
|
||||
}
|
||||
state.Status = bestExitDisplayStatusApplied
|
||||
state.UpdatedAt = latestUpdatedAt
|
||||
state.Reason = latestReason
|
||||
if len(appliedExitIDs) == 1 {
|
||||
for _, name := range appliedExitIDs {
|
||||
state.Summary = name
|
||||
}
|
||||
} else {
|
||||
state.Summary = bestExitDisplaySummaryMulti
|
||||
}
|
||||
return state, true
|
||||
}
|
||||
|
||||
func bestExitDisplayOwners(tunnel map[string]interface{}) ([]map[string]interface{}, string) {
|
||||
chainGroups := bestExitDisplayChainGroups(tunnel["chainNodes"])
|
||||
if len(chainGroups) > 0 {
|
||||
return chainGroups[len(chainGroups)-1], "chain"
|
||||
}
|
||||
return bestExitDisplayMapSlice(tunnel["inNodeId"]), "entry"
|
||||
}
|
||||
|
||||
func bestExitDisplayMapSlice(v interface{}) []map[string]interface{} {
|
||||
switch arr := v.(type) {
|
||||
case []map[string]interface{}:
|
||||
return arr
|
||||
case []interface{}:
|
||||
out := make([]map[string]interface{}, 0, len(arr))
|
||||
for _, item := range arr {
|
||||
if m, ok := item.(map[string]interface{}); ok {
|
||||
out = append(out, m)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayChainGroups(v interface{}) [][]map[string]interface{} {
|
||||
switch groups := v.(type) {
|
||||
case [][]map[string]interface{}:
|
||||
return groups
|
||||
case []interface{}:
|
||||
out := make([][]map[string]interface{}, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
items := bestExitDisplayMapSlice(group)
|
||||
if len(items) > 0 {
|
||||
out = append(out, items)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayNodeName(source map[string]interface{}, nodeID int64, lookup bestExitNodeNameLookup, fallback string) string {
|
||||
if source != nil {
|
||||
for _, key := range []string{"nodeName", "name"} {
|
||||
if name := strings.TrimSpace(asString(source[key])); name != "" {
|
||||
return name
|
||||
}
|
||||
}
|
||||
}
|
||||
if lookup != nil {
|
||||
if name, ok := lookup(nodeID); ok && strings.TrimSpace(name) != "" {
|
||||
return strings.TrimSpace(name)
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func bestExitUnknownOwnerName(role string) string {
|
||||
if role == "chain" {
|
||||
return bestExitUnknownChainName
|
||||
}
|
||||
return bestExitUnknownEntryName
|
||||
}
|
||||
@@ -0,0 +1,383 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBestExitDecisionSnapshotIsDefensiveCopy(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
score := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30, NodeName: "exit-a"}, 10, 0, 20, 0)
|
||||
|
||||
m.observeScores(key, []bestExitCandidateScore{score}, now)
|
||||
snapshot, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected snapshot")
|
||||
}
|
||||
if snapshot.AppliedExitNodeID != 30 || snapshot.UpdatedAt != now.UnixMilli() {
|
||||
t.Fatalf("unexpected snapshot: %+v", snapshot)
|
||||
}
|
||||
if len(snapshot.Scores) != 1 {
|
||||
t.Fatalf("expected one score in snapshot, got %+v", snapshot.Scores)
|
||||
}
|
||||
snapshot.Scores[0].ExitNodeID = 99
|
||||
|
||||
again, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected second snapshot")
|
||||
}
|
||||
if again.Scores[0].ExitNodeID != 30 {
|
||||
t.Fatalf("snapshot score mutation leaked into manager state: %+v", again.Scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForDirectMultiEntryOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(100, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 11}, 31, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
{"nodeId": int64(11)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 11: "入口 B", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if !state.Enabled || state.Summary != "多个出口" || state.Status != "applied" {
|
||||
t.Fatalf("unexpected state summary: %+v", state)
|
||||
}
|
||||
if state.UpdatedAt != now.Add(time.Second).UnixMilli() {
|
||||
t.Fatalf("expected latest updatedAt, got %d", state.UpdatedAt)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "entry" || state.Items[0].OwnerNodeName != "入口 A" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "entry" || state.Items[1].OwnerNodeName != "入口 B" || state.Items[1].ExitNodeName != "日本节点" {
|
||||
t.Fatalf("unexpected second item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForFinalChainHopOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(200, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 20}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 21}, 30, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(88),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{
|
||||
{{"nodeId": int64(15), "inx": int64(0)}},
|
||||
{{"nodeId": int64(20), "inx": int64(1)}, {"nodeId": int64(21), "inx": int64(1)}},
|
||||
},
|
||||
}
|
||||
names := map[int64]string{20: "中转 M1", 21: "中转 M2", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Summary != "香港节点" || state.Status != "applied" {
|
||||
t.Fatalf("expected single-exit summary, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two final-hop owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "chain" || state.Items[0].OwnerNodeName != "中转 M1" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first chain owner item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "chain" || state.Items[1].OwnerNodeName != "中转 M2" || state.Items[1].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected second chain owner item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateWaitingWhenNoAppliedDecisionExists(t *testing.T) {
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, newBestExitManager(), testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected waiting best exit state")
|
||||
}
|
||||
if state.Summary != "等待探测" || state.Status != "waiting" {
|
||||
t.Fatalf("expected waiting state, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 1 || state.Items[0].ExitNodeID != 0 || state.Items[0].ExitNodeName != "等待探测" {
|
||||
t.Fatalf("unexpected waiting item: %+v", state.Items)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateKeepsTopLevelWaitingWhenSomeOwnersPending(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(400, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
{"nodeId": int64(11)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 11: "入口 B", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Status != bestExitDisplayStatusWaiting || state.Summary != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected top-level waiting for partial owner state, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].ExitNodeID != 30 || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("expected first owner applied details to remain visible, got %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].ExitNodeID != 0 || state.Items[1].ExitNodeName != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected second owner waiting details, got %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateIgnoresAppliedExitRemovedFromTunnel(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(500, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 99, now)
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 30: "香港节点", 31: "日本节点", 99: "已删除节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Status != bestExitDisplayStatusWaiting || state.Summary != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected waiting state for stale applied exit, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 1 {
|
||||
t.Fatalf("expected one item, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].ExitNodeID != 0 || state.Items[0].ExitNodeName != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected stale exit to be ignored, got %+v", state.Items[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateSkipsNonBestAndSingleExitTunnels(t *testing.T) {
|
||||
nonBest := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": "round"},
|
||||
{"nodeId": int64(31), "strategy": "round"},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(nonBest, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected non-best tunnel to skip state, got %+v", state)
|
||||
}
|
||||
|
||||
singleExit := map[string]interface{}{
|
||||
"id": int64(78),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(singleExit, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected single-exit tunnel to skip state, got %+v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelListAttachesBestExitStateOnlyForEligibleTunnels(t *testing.T) {
|
||||
h := setupBestExitTunnelHandler(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelList(res, req)
|
||||
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Data []map[string]any `json:"data"`
|
||||
}
|
||||
decodeBestExitTunnelResponse(t, res, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success response, got code %d", payload.Code)
|
||||
}
|
||||
|
||||
bestTunnel := findTunnelResponseItem(t, payload.Data, 77)
|
||||
if _, ok := bestTunnel["bestExitState"]; !ok {
|
||||
t.Fatalf("expected eligible best multi-exit tunnel to include bestExitState: %+v", bestTunnel)
|
||||
}
|
||||
|
||||
singleExitTunnel := findTunnelResponseItem(t, payload.Data, 78)
|
||||
if _, ok := singleExitTunnel["bestExitState"]; ok {
|
||||
t.Fatalf("expected single-exit tunnel to omit bestExitState: %+v", singleExitTunnel)
|
||||
}
|
||||
|
||||
nonBestTunnel := findTunnelResponseItem(t, payload.Data, 79)
|
||||
if _, ok := nonBestTunnel["bestExitState"]; ok {
|
||||
t.Fatalf("expected non-best tunnel to omit bestExitState: %+v", nonBestTunnel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelGetAttachesBestExitStateToSelectedTunnel(t *testing.T) {
|
||||
h := setupBestExitTunnelHandler(t)
|
||||
|
||||
body := bytes.NewReader([]byte(`{"id":77}`))
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/get", body)
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelGet(res, req)
|
||||
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Data map[string]any `json:"data"`
|
||||
}
|
||||
decodeBestExitTunnelResponse(t, res, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success response, got code %d", payload.Code)
|
||||
}
|
||||
if _, ok := payload.Data["bestExitState"]; !ok {
|
||||
t.Fatalf("expected selected best multi-exit tunnel to include bestExitState: %+v", payload.Data)
|
||||
}
|
||||
}
|
||||
|
||||
func setupBestExitTunnelHandler(t *testing.T) *Handler {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
h := New(r, "secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
insertNode := func(id int64, name string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, id, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
insertNode(10, "entry-a")
|
||||
insertNode(30, "exit-a")
|
||||
insertNode(31, "exit-b")
|
||||
insertNode(32, "exit-c")
|
||||
|
||||
insertTunnel := func(id int64, name string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference)
|
||||
VALUES(?, ?, 1, 1, 'tls', 1, ?, ?, 1, ?, '')
|
||||
`, id, name, now, now, id).Error; err != nil {
|
||||
t.Fatalf("insert tunnel %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
insertTunnel(77, "best-multi")
|
||||
insertTunnel(78, "best-single")
|
||||
insertTunnel(79, "round-multi")
|
||||
|
||||
insertChain := func(tunnelID int64, chainType string, nodeID int64, strategy string, inx int64) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, ?, ?, 30001, ?, ?, 'tls')
|
||||
`, tunnelID, chainType, nodeID, strategy, inx).Error; err != nil {
|
||||
t.Fatalf("insert chain tunnel %d/%s/%d: %v", tunnelID, chainType, nodeID, err)
|
||||
}
|
||||
}
|
||||
insertChain(77, "1", 10, "round", 1)
|
||||
insertChain(77, "3", 30, tunnelStrategyBest, 1)
|
||||
insertChain(77, "3", 31, tunnelStrategyBest, 2)
|
||||
insertChain(78, "1", 10, "round", 1)
|
||||
insertChain(78, "3", 30, tunnelStrategyBest, 1)
|
||||
insertChain(79, "1", 10, "round", 1)
|
||||
insertChain(79, "3", 31, "round", 1)
|
||||
insertChain(79, "3", 32, "round", 2)
|
||||
|
||||
h.bestExit.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, time.UnixMilli(now))
|
||||
return h
|
||||
}
|
||||
|
||||
func decodeBestExitTunnelResponse(t *testing.T, res *httptest.ResponseRecorder, v any) {
|
||||
t.Helper()
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected HTTP %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
if err := json.NewDecoder(res.Body).Decode(v); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func findTunnelResponseItem(t *testing.T, items []map[string]any, id float64) map[string]any {
|
||||
t.Helper()
|
||||
for _, item := range items {
|
||||
if item["id"] == id {
|
||||
return item
|
||||
}
|
||||
}
|
||||
t.Fatalf("tunnel %.0f not found in response: %+v", id, items)
|
||||
return nil
|
||||
}
|
||||
|
||||
func testBestExitNameLookup(names map[int64]string) bestExitNodeNameLookup {
|
||||
return func(nodeID int64) (string, bool) {
|
||||
name := names[nodeID]
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,469 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
var errBestExitProbeForTest = errors.New("probe failed")
|
||||
|
||||
func TestBestExitScoreCombinesLatencyAndLoss(t *testing.T) {
|
||||
exit := chainNodeRecord{NodeID: 30, NodeName: "exit-a"}
|
||||
score := scoreBestExitCandidate(10, exit, 25, 2, 80, 3)
|
||||
|
||||
if !score.Success {
|
||||
t.Fatalf("expected successful score")
|
||||
}
|
||||
if score.OwnerNodeID != 10 || score.ExitNodeID != 30 {
|
||||
t.Fatalf("unexpected owner/exit ids: %+v", score)
|
||||
}
|
||||
if score.TotalLatency != 105 {
|
||||
t.Fatalf("expected total latency 105, got %v", score.TotalLatency)
|
||||
}
|
||||
if score.TotalLoss < 4.9 || score.TotalLoss > 5.0 {
|
||||
t.Fatalf("expected combined loss about 4.94, got %v", score.TotalLoss)
|
||||
}
|
||||
if score.Score < 599 || score.Score > 600 {
|
||||
t.Fatalf("expected score about 599, got %v", score.Score)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitScorePenalizesLoss(t *testing.T) {
|
||||
stable := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 80, 0, 80, 0)
|
||||
lowLatencyLossy := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 10, 5, 10, 5)
|
||||
|
||||
if !bestExitScoreLess(stable, lowLatencyLossy) {
|
||||
t.Fatalf("expected stable exit to beat low-latency lossy exit: stable=%+v lossy=%+v", stable, lowLatencyLossy)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitFailedCandidateSortsLast(t *testing.T) {
|
||||
failed := failedBestExitCandidate(10, chainNodeRecord{NodeID: 30}, "dial timeout")
|
||||
good := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 100, 0, 100, 0)
|
||||
|
||||
scores := []bestExitCandidateScore{failed, good}
|
||||
sortBestExitScores(scores)
|
||||
|
||||
if scores[0].ExitNodeID != 31 || scores[1].ExitNodeID != 30 {
|
||||
t.Fatalf("expected good score first and failed score last, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitInitialObservationAppliesWithoutSwitch(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate}, now)
|
||||
if decision.Switch {
|
||||
t.Fatalf("initial observation should not return switch: %+v", decision)
|
||||
}
|
||||
if m.decisions[key].AppliedExitNodeID != 31 {
|
||||
t.Fatalf("expected applied exit 31, got %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionRequiresMinimumAdvantage(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 90, 0, 90, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds+1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate below minimum advantage should not switch after repeated observations: %+v", decision)
|
||||
}
|
||||
}
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("expected applied exit to remain 30, got %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionSwitchesWithMinimumAdvantage(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds-1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate should wait for confirmations before switching: %+v", decision)
|
||||
}
|
||||
}
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds-1)*time.Second))
|
||||
if !decision.Switch || decision.ExitNodeID != 31 {
|
||||
t.Fatalf("candidate with enough advantage should switch after confirmations: %+v", decision)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitConfirmedSwitchDoesNotMarkAppliedUntilSetApplied(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds-1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate should wait for confirmations before switching: %+v", decision)
|
||||
}
|
||||
}
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds-1)*time.Second))
|
||||
if !decision.Switch || decision.ExitNodeID != 31 {
|
||||
t.Fatalf("candidate with enough advantage should switch after confirmations: %+v", decision)
|
||||
}
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("confirmed switch should not mark applied before runtime update: %+v", m.decisions[key])
|
||||
}
|
||||
|
||||
m.setApplied(key, decision.ExitNodeID, now.Add(time.Second))
|
||||
if m.decisions[key].AppliedExitNodeID != 31 {
|
||||
t.Fatalf("setApplied should commit confirmed switch: %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitApplyFailureStartsRetryCooldownWithoutChangingAppliedExit(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds-1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate should wait for confirmations before switching: %+v", decision)
|
||||
}
|
||||
}
|
||||
confirmed := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds-1)*time.Second))
|
||||
if !confirmed.Switch || confirmed.ExitNodeID != 31 {
|
||||
t.Fatalf("expected confirmed switch before apply failure: %+v", confirmed)
|
||||
}
|
||||
|
||||
m.recordApplyFailure(key, confirmed.ExitNodeID, now.Add(bestExitConfirmationRounds*time.Second))
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("apply failure should leave applied exit unchanged: %+v", m.decisions[key])
|
||||
}
|
||||
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds+1)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("apply retry cooldown should suppress immediate retry: %+v", decision)
|
||||
}
|
||||
if decision.Reason != "apply retry cooldown" {
|
||||
t.Fatalf("expected apply retry cooldown reason, got %q", decision.Reason)
|
||||
}
|
||||
|
||||
retry := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(bestExitConfirmationRounds*time.Second+bestExitApplyRetryCooldown))
|
||||
if !retry.Switch || retry.ExitNodeID != 31 {
|
||||
t.Fatalf("expected retry after apply cooldown: %+v", retry)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitEnsureAppliedDoesNotOverrideExistingAppliedExit(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
|
||||
m.ensureApplied(key, 30, now)
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("expected initial applied exit 30, got %+v", m.decisions[key])
|
||||
}
|
||||
if !m.decisions[key].LastSwitchAt.Equal(now) {
|
||||
t.Fatalf("expected initial applied timestamp, got %+v", m.decisions[key])
|
||||
}
|
||||
|
||||
m.ensureApplied(key, 31, now.Add(time.Minute))
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("ensureApplied should not override existing applied exit: %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitRoundPingerCachesByNodeHostAndPort(t *testing.T) {
|
||||
publicCalls := 0
|
||||
ownerCalls := 0
|
||||
pinger := newBestExitRoundPinger(func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
if ip == bestExitPublicTargetHost && port == bestExitPublicTargetPort {
|
||||
publicCalls++
|
||||
return float64(nodeID), 0, nil
|
||||
}
|
||||
ownerCalls++
|
||||
return float64(ownerCalls), 0, nil
|
||||
})
|
||||
|
||||
if lat, _, err := pinger(30, bestExitPublicTargetHost, bestExitPublicTargetPort, diagnosisExecOptions{}); err != nil || lat != 30 {
|
||||
t.Fatalf("unexpected first public ping result lat=%v err=%v", lat, err)
|
||||
}
|
||||
if lat, _, err := pinger(30, bestExitPublicTargetHost, bestExitPublicTargetPort, diagnosisExecOptions{}); err != nil || lat != 30 {
|
||||
t.Fatalf("unexpected cached public ping result lat=%v err=%v", lat, err)
|
||||
}
|
||||
if _, _, err := pinger(31, bestExitPublicTargetHost, bestExitPublicTargetPort, diagnosisExecOptions{}); err != nil {
|
||||
t.Fatalf("unexpected second exit public ping err=%v", err)
|
||||
}
|
||||
if publicCalls != 2 {
|
||||
t.Fatalf("expected public probes cached per exit node, got %d calls", publicCalls)
|
||||
}
|
||||
|
||||
if _, _, err := pinger(10, "10.0.0.30", 30030, diagnosisExecOptions{}); err != nil {
|
||||
t.Fatalf("unexpected owner ping err=%v", err)
|
||||
}
|
||||
if _, _, err := pinger(10, "10.0.0.30", 30030, diagnosisExecOptions{}); err != nil {
|
||||
t.Fatalf("unexpected repeated owner ping err=%v", err)
|
||||
}
|
||||
if ownerCalls != 1 {
|
||||
t.Fatalf("expected owner-to-exit probes cached by target, got %d calls", ownerCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionScoresAreDefensiveCopies(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now)
|
||||
decision.Scores[0].ExitNodeID = 99
|
||||
|
||||
if m.decisions[key].Scores[0].ExitNodeID != 31 {
|
||||
t.Fatalf("decision scores mutation leaked into manager state: %+v", m.decisions[key].Scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionRequiresConfirmationsAndCooldown(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
|
||||
if decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now); decision.Switch {
|
||||
t.Fatalf("first observation should not switch: %+v", decision)
|
||||
}
|
||||
if decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Second)); decision.Switch {
|
||||
t.Fatalf("second observation should not switch: %+v", decision)
|
||||
}
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(2*time.Second))
|
||||
if !decision.Switch || decision.ExitNodeID != 31 {
|
||||
t.Fatalf("third confirmed observation should switch to 31: %+v", decision)
|
||||
}
|
||||
|
||||
betterAgain := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 20, 0, 20, 0)
|
||||
if decision := m.observeScores(key, []bestExitCandidateScore{betterAgain, candidate}, now.Add(3*time.Second)); decision.Switch {
|
||||
t.Fatalf("cooldown should block immediate switch back: %+v", decision)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitOrderingUsesAppliedDecision(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
m.setApplied(key, 31, time.Unix(100, 0))
|
||||
targets := []tunnelRuntimeNode{
|
||||
{NodeID: 30, Strategy: tunnelStrategyBest},
|
||||
{NodeID: 31, Strategy: tunnelStrategyBest},
|
||||
{NodeID: 32, Strategy: tunnelStrategyBest},
|
||||
}
|
||||
|
||||
ordered := m.orderTargets(key, targets)
|
||||
if ordered[0].NodeID != 31 || ordered[1].NodeID != 30 || ordered[2].NodeID != 32 {
|
||||
t.Fatalf("unexpected order: %+v", ordered)
|
||||
}
|
||||
if targets[0].NodeID != 30 {
|
||||
t.Fatalf("orderTargets mutated input: %+v", targets)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTunnelChainConfigMapsBestStrategyToFIFO(t *testing.T) {
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
31: {ID: 31, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
|
||||
}
|
||||
targets := []tunnelRuntimeNode{
|
||||
{NodeID: 30, Port: 30030, Protocol: "tls", Strategy: tunnelStrategyBest, ChainType: 3},
|
||||
{NodeID: 31, Port: 30031, Protocol: "tls", Strategy: tunnelStrategyBest, ChainType: 3},
|
||||
}
|
||||
|
||||
chainData, err := buildTunnelChainConfig(77, 10, targets, nodes, "")
|
||||
if err != nil {
|
||||
t.Fatalf("build chain: %v", err)
|
||||
}
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
selector := hops[0]["selector"].(map[string]interface{})
|
||||
if selector["strategy"] != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best to render as fifo, got %v", selector["strategy"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerOrdersBestExitTargetsForOwner(t *testing.T) {
|
||||
h := &Handler{bestExit: newBestExitManager()}
|
||||
key := bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}
|
||||
h.bestExit.setApplied(key, 31, time.Unix(100, 0))
|
||||
targets := []tunnelRuntimeNode{
|
||||
{NodeID: 30, Port: 30030, Strategy: tunnelStrategyBest},
|
||||
{NodeID: 31, Port: 30031, Strategy: tunnelStrategyBest},
|
||||
}
|
||||
|
||||
ordered := h.orderBestExitTargets(77, 10, targets)
|
||||
if ordered[0].NodeID != 31 || ordered[1].NodeID != 30 {
|
||||
t.Fatalf("unexpected ordered targets: %+v", ordered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeStrategyForTargetsMapsBestTargetStrategyToFIFO(t *testing.T) {
|
||||
owner := tunnelRuntimeNode{Strategy: "round"}
|
||||
targets := []tunnelRuntimeNode{{Strategy: tunnelStrategyBest}}
|
||||
|
||||
if got := runtimeStrategyForTargets(owner, targets); got != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best target strategy to map to fifo, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeStrategyForTargetsPreservesNonBestTargetStrategy(t *testing.T) {
|
||||
owner := tunnelRuntimeNode{Strategy: tunnelStrategyBest}
|
||||
targets := []tunnelRuntimeNode{{Strategy: "round"}}
|
||||
|
||||
if got := runtimeStrategyForTargets(owner, targets); got != "round" {
|
||||
t.Fatalf("expected target strategy round to remain unchanged, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeStrategyForTargetsMapsBestOwnerStrategyWhenTargetsEmpty(t *testing.T) {
|
||||
owner := tunnelRuntimeNode{Strategy: tunnelStrategyBest}
|
||||
|
||||
if got := runtimeStrategyForTargets(owner, nil); got != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best owner fallback strategy to map to fifo, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{
|
||||
{NodeID: 30, NodeName: "exit-a", Port: 30030},
|
||||
{NodeID: 31, NodeName: "exit-b", Port: 30031},
|
||||
}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
31: {ID: 31, Status: 1, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
switch {
|
||||
case nodeID == 10 && port == 30030:
|
||||
return 60, 0, nil
|
||||
case nodeID == 10 && port == 30031:
|
||||
return 20, 0, nil
|
||||
case nodeID == 30 && ip == bestExitPublicTargetHost:
|
||||
return 60, 0, nil
|
||||
case nodeID == 31 && ip == bestExitPublicTargetHost:
|
||||
return 20, 0, nil
|
||||
default:
|
||||
t.Fatalf("unexpected ping node=%d ip=%s port=%d", nodeID, ip, port)
|
||||
return 0, 100, nil
|
||||
}
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, defaultTunnelProbeTarget(), pinger)
|
||||
if len(scores) != 2 {
|
||||
t.Fatalf("expected two scores, got %+v", scores)
|
||||
}
|
||||
if scores[0].ExitNodeID != 31 {
|
||||
t.Fatalf("expected exit-b first, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerSkipsOfflineCandidate(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
|
||||
30: {ID: 30, Status: 0, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
|
||||
}
|
||||
ping := func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
t.Fatalf("offline best-exit candidate should not be probed: node=%d target=%s:%d", nodeID, ip, port)
|
||||
return 0, 100, nil
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, defaultTunnelProbeTarget(), ping)
|
||||
if len(scores) != 1 || scores[0].Success {
|
||||
t.Fatalf("expected one failed offline candidate, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerUsesConfiguredPublicProbeTarget(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry-a"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30001}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Name: "entry-a", Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
|
||||
30: {ID: 30, Name: "exit-a", Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
|
||||
}
|
||||
target := tunnelProbeTarget{Host: "speed.example.com", Port: 8443}
|
||||
var calls []string
|
||||
ping := func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
|
||||
return 10, 0, nil
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, target, ping)
|
||||
if len(scores) != 1 || !scores[0].Success {
|
||||
t.Fatalf("expected successful score, got %+v", scores)
|
||||
}
|
||||
if !slices.Contains(calls, "30|speed.example.com|8443") {
|
||||
t.Fatalf("expected exit public probe to use configured target, calls=%+v", calls)
|
||||
}
|
||||
for _, call := range calls {
|
||||
if strings.Contains(call, defaultTunnelProbeTargetHost) {
|
||||
t.Fatalf("did not expect default target call when custom target configured: %+v", calls)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerMarksCandidateFailedWhenOwnerToExitFails(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
return 0, 100, errBestExitProbeForTest
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, defaultTunnelProbeTarget(), pinger)
|
||||
if len(scores) != 1 || scores[0].Success {
|
||||
t.Fatalf("expected failed candidate, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerMarksCandidateFailedWhenTargetResolutionFails(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-v6", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Name: "entry", Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Name: "exit-v6", Status: 1, ServerIP: "2001:db8::30", ServerIPv6: "2001:db8::30", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
t.Fatalf("ping should not be called when target resolution fails: node=%d ip=%s port=%d", nodeID, ip, port)
|
||||
return 0, 100, nil
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "v4", diagnosisExecOptions{}, defaultTunnelProbeTarget(), pinger)
|
||||
if len(scores) != 1 || scores[0].Success {
|
||||
t.Fatalf("expected failed candidate, got %+v", scores)
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type tunnelTrafficDelta struct {
|
||||
@@ -21,75 +22,42 @@ func unixMilliBucketMinute(nowMs int64) int64 {
|
||||
return nowMs - (nowMs % minuteMs)
|
||||
}
|
||||
|
||||
func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowItem, nowMs int64) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
if nodeID <= 0 || len(items) == 0 {
|
||||
return
|
||||
func collectFlowUploadForwardIDs(items []flowItem) []int64 {
|
||||
ids := make([]int64, 0, len(items))
|
||||
seen := make(map[int64]struct{}, len(items))
|
||||
for _, item := range items {
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(strings.TrimSpace(item.N))
|
||||
if !ok || forwardID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[forwardID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[forwardID] = struct{}{}
|
||||
ids = append(ids, forwardID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func (h *Handler) recordTunnelMetricsFromForwardBatch(nodeID int64, forwardDeltas map[int64]tunnelTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta, nowMs int64) {
|
||||
if h == nil || h.repo == nil || nodeID <= 0 || len(forwardDeltas) == 0 {
|
||||
return
|
||||
}
|
||||
bucketTs := unixMilliBucketMinute(nowMs)
|
||||
if bucketTs <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
forwardDeltas := make(map[int64]tunnelTrafficDelta)
|
||||
var skippedParse, skippedZero int
|
||||
for _, item := range items {
|
||||
name := strings.TrimSpace(item.N)
|
||||
if name == "" || name == "web_api" {
|
||||
continue
|
||||
}
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(name)
|
||||
if !ok {
|
||||
skippedParse++
|
||||
continue
|
||||
}
|
||||
if item.D == 0 && item.U == 0 {
|
||||
skippedZero++
|
||||
continue
|
||||
}
|
||||
d := forwardDeltas[forwardID]
|
||||
d.bytesIn += item.D
|
||||
d.bytesOut += item.U
|
||||
forwardDeltas[forwardID] = d
|
||||
}
|
||||
if len(forwardDeltas) == 0 {
|
||||
if len(items) > 0 {
|
||||
log.Printf("monitoring debug op=tunnel_metric.no_forward_deltas node_id=%d items=%d skipped_parse=%d skipped_zero=%d", nodeID, len(items), skippedParse, skippedZero)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
forwardIDs := make([]int64, 0, len(forwardDeltas))
|
||||
for id := range forwardDeltas {
|
||||
forwardIDs = append(forwardIDs, id)
|
||||
}
|
||||
|
||||
forwardTunnelMap, err := h.repo.MapForwardIDsToTunnelIDs(forwardIDs)
|
||||
if err != nil {
|
||||
log.Printf("monitoring write skipped op=tunnel_metric.map_forward_to_tunnel node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
if len(forwardTunnelMap) == 0 {
|
||||
log.Printf("monitoring debug op=tunnel_metric.no_tunnel_map node_id=%d forward_ids=%v", nodeID, forwardIDs)
|
||||
return
|
||||
}
|
||||
|
||||
tunnelAgg := make(map[int64]tunnelTrafficDelta)
|
||||
for forwardID, delta := range forwardDeltas {
|
||||
tunnelID := forwardTunnelMap[forwardID]
|
||||
if tunnelID <= 0 {
|
||||
meta, ok := metas[forwardID]
|
||||
if !ok || meta.TunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
a := tunnelAgg[tunnelID]
|
||||
a.bytesIn += delta.bytesIn
|
||||
a.bytesOut += delta.bytesOut
|
||||
tunnelAgg[tunnelID] = a
|
||||
}
|
||||
if len(tunnelAgg) == 0 {
|
||||
return
|
||||
current := tunnelAgg[meta.TunnelID]
|
||||
current.bytesIn += delta.bytesIn
|
||||
current.bytesOut += delta.bytesOut
|
||||
tunnelAgg[meta.TunnelID] = current
|
||||
}
|
||||
|
||||
metrics := make([]*model.TunnelMetric, 0, len(tunnelAgg))
|
||||
@@ -98,14 +66,11 @@ func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowIte
|
||||
continue
|
||||
}
|
||||
metrics = append(metrics, &model.TunnelMetric{
|
||||
TunnelID: tunnelID,
|
||||
NodeID: nodeID,
|
||||
Timestamp: bucketTs,
|
||||
BytesIn: delta.bytesIn,
|
||||
BytesOut: delta.bytesOut,
|
||||
Connections: 0,
|
||||
Errors: 0,
|
||||
AvgLatencyMs: 0,
|
||||
TunnelID: tunnelID,
|
||||
NodeID: nodeID,
|
||||
Timestamp: bucketTs,
|
||||
BytesIn: delta.bytesIn,
|
||||
BytesOut: delta.bytesOut,
|
||||
})
|
||||
}
|
||||
if len(metrics) == 0 {
|
||||
@@ -114,7 +79,7 @@ func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowIte
|
||||
|
||||
if err := h.repo.UpsertTunnelMetricBuckets(metrics); err != nil {
|
||||
log.Printf("monitoring write failed op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d err=%v", nodeID, bucketTs, len(metrics), err)
|
||||
} else {
|
||||
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
|
||||
return
|
||||
}
|
||||
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTunnelProbeTargetHost = "www.bing.com"
|
||||
defaultTunnelProbeTargetPort = 443
|
||||
)
|
||||
|
||||
type tunnelProbeTarget struct {
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
|
||||
func defaultTunnelProbeTarget() tunnelProbeTarget {
|
||||
return tunnelProbeTarget{Host: defaultTunnelProbeTargetHost, Port: defaultTunnelProbeTargetPort}
|
||||
}
|
||||
|
||||
func normalizeTunnelProbeTarget(host string, port int) (tunnelProbeTarget, bool, error) {
|
||||
host = strings.TrimSpace(host)
|
||||
if host == "" && port == 0 {
|
||||
return defaultTunnelProbeTarget(), false, nil
|
||||
}
|
||||
if host == "" {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标 Host 不能为空")
|
||||
}
|
||||
if port <= 0 || port > 65535 {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标端口必须是 1-65535")
|
||||
}
|
||||
if strings.Contains(host, "://") || strings.ContainsAny(host, "/?#") || strings.ContainsAny(host, " \t\r\n") || isTunnelProbeTargetSchemeLikeHost(host) {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标 Host 不能包含协议或路径")
|
||||
}
|
||||
if normalized, ok := normalizeTunnelProbeTargetHost(host); ok {
|
||||
host = normalized
|
||||
} else {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标 Host 格式无效")
|
||||
}
|
||||
|
||||
return tunnelProbeTarget{Host: host, Port: port}, true, nil
|
||||
}
|
||||
|
||||
func normalizeTunnelProbeTargetHost(host string) (string, bool) {
|
||||
if strings.HasPrefix(host, "[") || strings.HasSuffix(host, "]") {
|
||||
if !strings.HasPrefix(host, "[") || !strings.HasSuffix(host, "]") {
|
||||
return "", false
|
||||
}
|
||||
inner := strings.TrimPrefix(strings.TrimSuffix(host, "]"), "[")
|
||||
addr, err := netip.ParseAddr(inner)
|
||||
if err != nil || !addr.Is6() {
|
||||
return "", false
|
||||
}
|
||||
return inner, true
|
||||
}
|
||||
|
||||
if addr, err := netip.ParseAddr(host); err == nil {
|
||||
return addr.String(), true
|
||||
}
|
||||
if strings.Contains(host, ":") || isTunnelProbeTargetIPv4Like(host) {
|
||||
return "", false
|
||||
}
|
||||
if !isValidTunnelProbeTargetHost(host) {
|
||||
return "", false
|
||||
}
|
||||
return host, true
|
||||
}
|
||||
|
||||
func isValidTunnelProbeTargetHost(host string) bool {
|
||||
if host == "" || len(host) > 253 {
|
||||
return false
|
||||
}
|
||||
for _, label := range strings.Split(host, ".") {
|
||||
if len(label) == 0 || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return false
|
||||
}
|
||||
for _, r := range label {
|
||||
if !isASCIILetter(r) && !isASCIIDigit(r) && r != '-' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func isTunnelProbeTargetIPv4Like(host string) bool {
|
||||
if host == "" {
|
||||
return false
|
||||
}
|
||||
for _, r := range host {
|
||||
if !isASCIIDigit(r) && r != '.' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return strings.Contains(host, ".")
|
||||
}
|
||||
|
||||
func isTunnelProbeTargetSchemeLikeHost(host string) bool {
|
||||
if _, err := netip.ParseAddr(host); err == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
colon := strings.IndexByte(host, ':')
|
||||
if colon <= 0 {
|
||||
return false
|
||||
}
|
||||
for i, r := range host[:colon] {
|
||||
if i == 0 {
|
||||
if !isASCIILetter(r) {
|
||||
return false
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !isASCIILetter(r) && !isASCIIDigit(r) && r != '+' && r != '-' && r != '.' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func isASCIILetter(r rune) bool {
|
||||
return (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z')
|
||||
}
|
||||
|
||||
func isASCIIDigit(r rune) bool {
|
||||
return r >= '0' && r <= '9'
|
||||
}
|
||||
|
||||
func parseTunnelProbeTargetFromRequest(req map[string]interface{}) (tunnelProbeTarget, bool, error) {
|
||||
if req == nil {
|
||||
return defaultTunnelProbeTarget(), false, nil
|
||||
}
|
||||
rawHost, hasHost := req["probeTargetHost"]
|
||||
rawPort, hasPort := req["probeTargetPort"]
|
||||
if !hasHost && !hasPort {
|
||||
return defaultTunnelProbeTarget(), false, nil
|
||||
}
|
||||
host, err := parseTunnelProbeTargetHostValue(rawHost)
|
||||
if err != nil {
|
||||
return tunnelProbeTarget{}, false, err
|
||||
}
|
||||
port, err := parseTunnelProbeTargetPortValue(rawPort)
|
||||
if err != nil {
|
||||
return tunnelProbeTarget{}, false, err
|
||||
}
|
||||
return normalizeTunnelProbeTarget(host, port)
|
||||
}
|
||||
|
||||
func parseTunnelProbeTargetHostValue(raw interface{}) (string, error) {
|
||||
if raw == nil {
|
||||
return "", nil
|
||||
}
|
||||
host, ok := raw.(string)
|
||||
if !ok {
|
||||
return "", errors.New("测试目标 Host 格式无效")
|
||||
}
|
||||
if host != strings.TrimSpace(host) {
|
||||
return "", errors.New("测试目标 Host 不能包含协议或路径")
|
||||
}
|
||||
return host, nil
|
||||
}
|
||||
|
||||
func parseTunnelProbeTargetPortValue(raw interface{}) (int, error) {
|
||||
if raw == nil {
|
||||
return 0, nil
|
||||
}
|
||||
switch v := raw.(type) {
|
||||
case float64:
|
||||
if v != float64(int64(v)) {
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
return int(v), nil
|
||||
case string:
|
||||
if v == "" {
|
||||
return 0, nil
|
||||
}
|
||||
if v != strings.TrimSpace(v) {
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
port, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
return port, nil
|
||||
case int:
|
||||
return v, nil
|
||||
case int32:
|
||||
return int(v), nil
|
||||
case int64:
|
||||
return int(v), nil
|
||||
default:
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
}
|
||||
|
||||
func effectiveTunnelProbeTarget(tunnel *model.Tunnel) tunnelProbeTarget {
|
||||
if tunnel == nil {
|
||||
return defaultTunnelProbeTarget()
|
||||
}
|
||||
return effectiveTunnelProbeTargetValues(tunnel.ProbeTargetHost, tunnel.ProbeTargetPort)
|
||||
}
|
||||
|
||||
func effectiveTunnelProbeTargetValues(host string, port int) tunnelProbeTarget {
|
||||
target, configured, err := normalizeTunnelProbeTarget(host, port)
|
||||
if err != nil || !configured {
|
||||
return defaultTunnelProbeTarget()
|
||||
}
|
||||
return target
|
||||
}
|
||||
|
||||
func formatTunnelProbeTarget(target tunnelProbeTarget) string {
|
||||
if addr, err := netip.ParseAddr(target.Host); err == nil && addr.Is6() {
|
||||
return fmt.Sprintf("[%s]:%d", target.Host, target.Port)
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", target.Host, target.Port)
|
||||
}
|
||||
@@ -0,0 +1,305 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestTunnelCreatePersistsProbeTargetAndListReturnsConfiguredValue(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"name":"custom-target",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"speed.example.com",
|
||||
"probeTargetPort":8443
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelCreate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", body))
|
||||
assertProbeTargetSuccess(t, res)
|
||||
|
||||
listRes := httptest.NewRecorder()
|
||||
h.tunnelList(listRes, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Data []map[string]any `json:"data"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, listRes, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success, got code %d", payload.Code)
|
||||
}
|
||||
item := payload.Data[0]
|
||||
if item["probeTargetHost"] != "speed.example.com" || item["probeTargetPort"] != float64(8443) {
|
||||
t.Fatalf("unexpected probe target in list response: %+v", item)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdatePersistsDefaultProbeTargetAsEmpty(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 77, "existing", "old.example.com", 9443)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":77,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"",
|
||||
"probeTargetPort":0
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
assertProbeTargetSuccess(t, res)
|
||||
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
t.Fatalf("list tunnels: %v", err)
|
||||
}
|
||||
item := findProbeTargetTunnelItem(t, items, 77)
|
||||
if item["probeTargetHost"] != "" || item["probeTargetPort"] != 0 {
|
||||
t.Fatalf("expected default target to round-trip as empty/0, got %+v", item)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateWithoutProbeTargetFieldsPreservesExistingTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 79, "existing", "old.example.com", 9443)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":79,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}]
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
assertProbeTargetSuccess(t, res)
|
||||
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
t.Fatalf("list tunnels: %v", err)
|
||||
}
|
||||
item := findProbeTargetTunnelItem(t, items, 79)
|
||||
if item["probeTargetHost"] != "old.example.com" || item["probeTargetPort"] != 9443 {
|
||||
t.Fatalf("expected omitted probe target fields to preserve existing target, got %+v", item)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateRejectsInvalidProbeTargetWithoutClearingExistingTarget(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
probeFields string
|
||||
}{
|
||||
{name: "non numeric port", probeFields: `,"probeTargetPort":"abc"`},
|
||||
{name: "fractional port", probeFields: `,"probeTargetPort":443.5`},
|
||||
{name: "whitespace host", probeFields: `,"probeTargetHost":" "`},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 80, "existing", "old.example.com", 9443)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":80,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}]
|
||||
` + tt.probeFields + `}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code == 0 || payload.Msg == "" {
|
||||
t.Fatalf("expected validation failure, got %+v", payload)
|
||||
}
|
||||
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
t.Fatalf("list tunnels: %v", err)
|
||||
}
|
||||
item := findProbeTargetTunnelItem(t, items, 80)
|
||||
if item["probeTargetHost"] != "old.example.com" || item["probeTargetPort"] != 9443 {
|
||||
t.Fatalf("expected invalid probe target to preserve existing target, got %+v", item)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelCreateRejectsInvalidProbeTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"name":"bad-target",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"https://example.com",
|
||||
"probeTargetPort":443
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelCreate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", body))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code == 0 || payload.Msg == "" {
|
||||
t.Fatalf("expected validation failure, got %+v", payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateInvalidProbeTargetDoesNotCleanFederationBindings(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 88, "existing", "old.example.com", 9443)
|
||||
seedProbeTargetFederationBinding(t, h, 88)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":88,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"https://example.com",
|
||||
"probeTargetPort":443
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code == 0 || payload.Msg == "" {
|
||||
t.Fatalf("expected validation failure, got %+v", payload)
|
||||
}
|
||||
|
||||
bindings, err := h.repo.ListActiveFederationTunnelBindingsByTunnel(88)
|
||||
if err != nil {
|
||||
t.Fatalf("list federation bindings: %v", err)
|
||||
}
|
||||
if len(bindings) != 1 {
|
||||
t.Fatalf("expected federation binding to remain after invalid update, got %d", len(bindings))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelDiagnosisUsesConfiguredProbeTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 90, "diagnosis-target", "speed.example.com", 8443)
|
||||
|
||||
_, _, workItems, err := h.prepareTunnelDiagnosis(90)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare tunnel diagnosis: %v", err)
|
||||
}
|
||||
if len(workItems) != 1 {
|
||||
t.Fatalf("expected one diagnosis item, got %d", len(workItems))
|
||||
}
|
||||
if workItems[0].targetIP != "speed.example.com" || workItems[0].targetPort != 8443 {
|
||||
t.Fatalf("expected custom diagnosis target speed.example.com:8443, got %s:%d", workItems[0].targetIP, workItems[0].targetPort)
|
||||
}
|
||||
}
|
||||
|
||||
func setupProbeTargetTunnelHandler(t *testing.T) *Handler {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
h := New(r, "secret")
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(10, 'entry-a', 'entry-secret', '10.0.0.1', '10.0.0.1', '', '30000-30010', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
func seedProbeTargetTunnel(t *testing.T, h *Handler, id int64, name string, host string, port int) {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference, probe_target_host, probe_target_port)
|
||||
VALUES(?, ?, 1, 1, 'tls', 1, ?, ?, 1, ?, '', ?, ?)
|
||||
`, id, name, now, now, id, host, port).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, '1', 10, 30001, 'round', 1, 'tls')
|
||||
`, id).Error; err != nil {
|
||||
t.Fatalf("insert chain: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func seedProbeTargetFederationBinding(t *testing.T, h *Handler, tunnelID int64) {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time)
|
||||
VALUES(?, 10, 1, 0, 'http://peer.example', ?, 'remote-binding', 30001, 1, ?, ?)
|
||||
`, tunnelID, "probe-target-test-binding", now, now).Error; err != nil {
|
||||
t.Fatalf("insert federation binding: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertProbeTargetSuccess(t *testing.T, res *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success, got %+v", payload)
|
||||
}
|
||||
}
|
||||
|
||||
func decodeProbeTargetResponse(t *testing.T, res *httptest.ResponseRecorder, v any) {
|
||||
t.Helper()
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected HTTP %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
if err := json.NewDecoder(res.Body).Decode(v); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func findProbeTargetTunnelItem(t *testing.T, items []map[string]interface{}, id int64) map[string]interface{} {
|
||||
t.Helper()
|
||||
for _, item := range items {
|
||||
if asInt64(item["id"], 0) == id {
|
||||
return item
|
||||
}
|
||||
}
|
||||
t.Fatalf("tunnel %d not found: %+v", id, items)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package handler
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestNormalizeTunnelProbeTargetDefaultsWhenEmpty(t *testing.T) {
|
||||
target, configured, err := normalizeTunnelProbeTarget("", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if configured {
|
||||
t.Fatalf("expected empty input to be default, not configured")
|
||||
}
|
||||
if target.Host != defaultTunnelProbeTargetHost || target.Port != defaultTunnelProbeTargetPort {
|
||||
t.Fatalf("unexpected default target: %+v", target)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetAcceptsHostPortAndIPv6(t *testing.T) {
|
||||
target, configured, err := normalizeTunnelProbeTarget(" [2001:db8::1] ", 8443)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !configured {
|
||||
t.Fatalf("expected explicit target")
|
||||
}
|
||||
if target.Host != "2001:db8::1" || target.Port != 8443 {
|
||||
t.Fatalf("unexpected normalized target: %+v", target)
|
||||
}
|
||||
if got := formatTunnelProbeTarget(target); got != "[2001:db8::1]:8443" {
|
||||
t.Fatalf("unexpected formatted target: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetRejectsPartialAndInvalidInputs(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
host string
|
||||
port int
|
||||
}{
|
||||
{name: "missing host", host: "", port: 443},
|
||||
{name: "missing port", host: "example.com", port: 0},
|
||||
{name: "port too high", host: "example.com", port: 70000},
|
||||
{name: "scheme", host: "https://example.com", port: 443},
|
||||
{name: "path", host: "example.com/ping", port: 443},
|
||||
{name: "space", host: "example .com", port: 443},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if _, _, err := normalizeTunnelProbeTarget(tt.host, tt.port); err == nil {
|
||||
t.Fatalf("expected validation error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetRejectsSchemePrefixButAllowsIPv6(t *testing.T) {
|
||||
for _, host := range []string{"https:example.com", "mailto:ops@example.com"} {
|
||||
if _, _, err := normalizeTunnelProbeTarget(host, 443); err == nil {
|
||||
t.Fatalf("expected scheme-like host %q to be rejected", host)
|
||||
}
|
||||
}
|
||||
|
||||
for _, host := range []string{"2001:db8::1", "[2001:db8::1]"} {
|
||||
target, configured, err := normalizeTunnelProbeTarget(host, 443)
|
||||
if err != nil {
|
||||
t.Fatalf("expected IPv6 host %q to be accepted: %v", host, err)
|
||||
}
|
||||
if !configured || target.Host != "2001:db8::1" {
|
||||
t.Fatalf("unexpected IPv6 normalization for %q: %+v configured=%v", host, target, configured)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetValidatesHostShape(t *testing.T) {
|
||||
validHosts := []string{
|
||||
"example.com",
|
||||
"localhost",
|
||||
"api-1.example.co.uk",
|
||||
"192.0.2.10",
|
||||
"2001:db8::1",
|
||||
"[2001:db8::1]",
|
||||
}
|
||||
for _, host := range validHosts {
|
||||
if _, _, err := normalizeTunnelProbeTarget(host, 443); err != nil {
|
||||
t.Fatalf("expected valid host %q: %v", host, err)
|
||||
}
|
||||
}
|
||||
|
||||
invalidHosts := []string{
|
||||
"1:2:3",
|
||||
"[2001:db8::1",
|
||||
"2001:db8::1]",
|
||||
"[example.com]",
|
||||
"example..com",
|
||||
"-example.com",
|
||||
"example-.com",
|
||||
"exa_mple.com",
|
||||
"999.1.1.1",
|
||||
}
|
||||
for _, host := range invalidHosts {
|
||||
if _, _, err := normalizeTunnelProbeTarget(host, 443); err == nil {
|
||||
t.Fatalf("expected invalid host %q to be rejected", host)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseTunnelProbeTargetFromRequest(t *testing.T) {
|
||||
req := map[string]interface{}{
|
||||
"probeTargetHost": "speed.example.com",
|
||||
"probeTargetPort": float64(1443),
|
||||
}
|
||||
target, configured, err := parseTunnelProbeTargetFromRequest(req)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !configured || target.Host != "speed.example.com" || target.Port != 1443 {
|
||||
t.Fatalf("unexpected request target: %+v configured=%v", target, configured)
|
||||
}
|
||||
}
|
||||
@@ -3,19 +3,20 @@ package handler
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
const (
|
||||
tunnelQualityProbeInterval = 1 * time.Second
|
||||
tunnelQualityProbeTimeout = 8 * time.Second
|
||||
tunnelQualityPingTimeoutMs = 5000
|
||||
tunnelQualityRetention = 24 * time.Hour // keep 24h of history
|
||||
tunnelQualityPruneInterval = 10 * time.Minute
|
||||
tunnelQualityReportInterval = 30 * time.Second // DB save interval
|
||||
)
|
||||
@@ -31,6 +32,26 @@ type TunnelQualityHop struct {
|
||||
TargetPort int `json:"targetPort,omitempty"`
|
||||
}
|
||||
|
||||
type TunnelQualityCandidateHop struct {
|
||||
TunnelQualityHop
|
||||
FromRole string `json:"fromRole"`
|
||||
ToRole string `json:"toRole"`
|
||||
HopIndex int `json:"hopIndex"`
|
||||
Selected bool `json:"selected"`
|
||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||
}
|
||||
|
||||
type tunnelQualityChainDetails struct {
|
||||
PrimaryPath []TunnelQualityHop `json:"primaryPath,omitempty"`
|
||||
CandidateHops []TunnelQualityCandidateHop `json:"candidateHops,omitempty"`
|
||||
}
|
||||
|
||||
type tunnelQualityCandidateGroup struct {
|
||||
role string
|
||||
roleIndex int
|
||||
nodes []chainNodeRecord
|
||||
}
|
||||
|
||||
// tunnelQualitySnapshot is the in-memory latest probe result for a tunnel.
|
||||
type tunnelQualitySnapshot struct {
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
@@ -42,6 +63,8 @@ type tunnelQualitySnapshot struct {
|
||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||
Timestamp int64 `json:"timestamp"`
|
||||
ChainDetails string `json:"chainDetails,omitempty"`
|
||||
ProbeTargetHost string `json:"probeTargetHost,omitempty"`
|
||||
ProbeTargetPort int `json:"probeTargetPort,omitempty"`
|
||||
|
||||
// internal fields for db reporting
|
||||
lastDBWrite int64 `json:"-"`
|
||||
@@ -54,16 +77,17 @@ type tunnelQualityProber struct {
|
||||
cache sync.Map // tunnelID (int64) → *tunnelQualitySnapshot
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
interval time.Duration
|
||||
wake chan struct{}
|
||||
lastPrune int64
|
||||
probing int32 // atomic flag: 1 = probeAll running, 0 = idle
|
||||
probeNode bestExitProbeFunc
|
||||
}
|
||||
|
||||
// newTunnelQualityProber creates a new prober (not yet running).
|
||||
func newTunnelQualityProber(h *Handler) *tunnelQualityProber {
|
||||
return &tunnelQualityProber{
|
||||
handler: h,
|
||||
interval: tunnelQualityProbeInterval,
|
||||
handler: h,
|
||||
wake: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -83,6 +107,16 @@ func (p *tunnelQualityProber) Stop() {
|
||||
p.cancel()
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) NotifyConfigChanged() {
|
||||
if p == nil || p.wake == nil {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case p.wake <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// GetAll returns all cached quality snapshots (latest per tunnel).
|
||||
func (p *tunnelQualityProber) GetAll() []tunnelQualitySnapshot {
|
||||
var items []tunnelQualitySnapshot
|
||||
@@ -106,20 +140,44 @@ func (p *tunnelQualityProber) loop() {
|
||||
// Run once immediately
|
||||
p.probeAll()
|
||||
|
||||
ticker := time.NewTicker(p.interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
timer := time.NewTimer(p.probeInterval())
|
||||
select {
|
||||
case <-p.ctx.Done():
|
||||
stopAndDrainTunnelQualityTimer(timer)
|
||||
return
|
||||
case <-ticker.C:
|
||||
case <-p.wake:
|
||||
stopAndDrainTunnelQualityTimer(timer)
|
||||
continue
|
||||
case <-timer.C:
|
||||
p.probeAll()
|
||||
p.maybePrune()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func stopAndDrainTunnelQualityTimer(timer *time.Timer) {
|
||||
if timer == nil || timer.Stop() {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case <-timer.C:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeInterval() time.Duration {
|
||||
if p == nil || p.handler == nil || p.handler.repo == nil {
|
||||
return time.Duration(monitoring.DefaultTunnelQualityProbeIntervalSec) * time.Second
|
||||
}
|
||||
cfg, err := p.handler.repo.GetConfigsByNames([]string{monitoring.ConfigTunnelQualityProbeIntervalSec})
|
||||
if err != nil {
|
||||
return time.Duration(monitoring.DefaultTunnelQualityProbeIntervalSec) * time.Second
|
||||
}
|
||||
seconds := monitoring.TunnelQualityProbeIntervalSecondsFromConfigMap(cfg)
|
||||
return time.Duration(seconds) * time.Second
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) isEnabled() bool {
|
||||
if p == nil || p.handler == nil {
|
||||
return true
|
||||
@@ -128,12 +186,19 @@ func (p *tunnelQualityProber) isEnabled() bool {
|
||||
return p.handler.isTunnelQualityMonitoringEnabled()
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) retentionDays() int {
|
||||
if p == nil || p.handler == nil || p.handler.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := p.handler.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
// maybePrune deletes old quality rows periodically (mirrors PruneServiceMonitorResults).
|
||||
func (p *tunnelQualityProber) maybePrune() {
|
||||
if !p.isEnabled() {
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if p.lastPrune > 0 && now-p.lastPrune < int64(tunnelQualityPruneInterval/time.Millisecond) {
|
||||
return
|
||||
@@ -145,7 +210,7 @@ func (p *tunnelQualityProber) maybePrune() {
|
||||
return
|
||||
}
|
||||
|
||||
cutoff := now - int64(tunnelQualityRetention/time.Millisecond)
|
||||
cutoff := now - int64(time.Duration(p.retentionDays())*24*time.Hour/time.Millisecond)
|
||||
if err := h.repo.PruneTunnelQualityResults(cutoff); err != nil {
|
||||
log.Printf("tunnel_quality_prober: prune err=%v", err)
|
||||
}
|
||||
@@ -219,6 +284,9 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
p.storeResult(snap)
|
||||
return
|
||||
}
|
||||
probeTarget := effectiveTunnelProbeTargetValues(tunnel.ProbeTargetHost, tunnel.ProbeTargetPort)
|
||||
snap.ProbeTargetHost = probeTarget.Host
|
||||
snap.ProbeTargetPort = probeTarget.Port
|
||||
|
||||
chainRows, err := h.listChainNodesForTunnel(tunnelID)
|
||||
if err != nil || len(chainRows) == 0 {
|
||||
@@ -233,14 +301,28 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
options := diagnosisExecOptions{
|
||||
commandTimeout: tunnelQualityProbeTimeout,
|
||||
pingTimeoutMS: tunnelQualityPingTimeoutMs,
|
||||
pingCount: 1,
|
||||
timeoutMessage: "探测超时",
|
||||
}
|
||||
roundPinger := newBestExitRoundPinger(p.pingNode)
|
||||
p.probeBestExitOwners(tunnelID, inNodes, midNodesGrouped, outNodes, ipPreference, options, probeTarget, roundPinger)
|
||||
|
||||
entry, _, entryOnline := p.firstOnlineChainNode(inNodes)
|
||||
exit, _, exitOnline := p.firstOnlineChainNode(outNodes)
|
||||
selectedNodeIDs := make(map[string]int64, 2+len(midNodesGrouped))
|
||||
if entryOnline {
|
||||
selectedNodeIDs[tunnelQualityGroupKey("entry", 0)] = entry.NodeID
|
||||
}
|
||||
if exitOnline {
|
||||
selectedNodeIDs[tunnelQualityGroupKey("exit", 0)] = exit.NodeID
|
||||
}
|
||||
var primaryHops []TunnelQualityHop
|
||||
|
||||
switch tunnel.Type {
|
||||
case 1:
|
||||
// Port forwarding: entry → Bing only
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(inNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
// Port forwarding: entry → public probe target only.
|
||||
if entryOnline {
|
||||
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -248,24 +330,42 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
} else {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
} else {
|
||||
snap.ErrorMessage = "入口节点均不在线"
|
||||
}
|
||||
case 2:
|
||||
// Tunnel forwarding: entry → exit + exit → Bing
|
||||
probeOK := true
|
||||
|
||||
if len(inNodes) > 0 && len(outNodes) > 0 {
|
||||
var hops []TunnelQualityHop
|
||||
if !entryOnline {
|
||||
probeOK = false
|
||||
snap.ErrorMessage = "入口节点均不在线"
|
||||
snap.EntryToExitLatency = -1
|
||||
snap.EntryToExitLoss = 100
|
||||
} else if !exitOnline {
|
||||
probeOK = false
|
||||
snap.ErrorMessage = "出口节点均不在线"
|
||||
snap.EntryToExitLatency = -1
|
||||
snap.EntryToExitLoss = 100
|
||||
} else {
|
||||
var totalLat float64
|
||||
remainingSuccessProb := 1.0
|
||||
|
||||
nodesInPath := make([]chainNodeRecord, 0, 2+len(midNodesGrouped))
|
||||
nodesInPath = append(nodesInPath, inNodes[0])
|
||||
for _, midGroup := range midNodesGrouped {
|
||||
if len(midGroup) > 0 {
|
||||
nodesInPath = append(nodesInPath, midGroup[0])
|
||||
nodesInPath = append(nodesInPath, entry)
|
||||
for midIndex, midGroup := range midNodesGrouped {
|
||||
mid, _, online := p.firstOnlineChainNode(midGroup)
|
||||
if !online {
|
||||
probeOK = false
|
||||
snap.ErrorMessage = "中间节点组均不在线"
|
||||
break
|
||||
}
|
||||
nodesInPath = append(nodesInPath, mid)
|
||||
selectedNodeIDs[tunnelQualityGroupKey("middle", midIndex)] = mid.NodeID
|
||||
}
|
||||
if probeOK {
|
||||
nodesInPath = append(nodesInPath, exit)
|
||||
}
|
||||
nodesInPath = append(nodesInPath, outNodes[0])
|
||||
|
||||
for i := 0; i < len(nodesInPath)-1; i++ {
|
||||
source := nodesInPath[i]
|
||||
@@ -279,15 +379,15 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
}
|
||||
|
||||
targetNode, nodeErr := h.getNodeRecord(target.NodeID)
|
||||
if nodeErr != nil || targetNode == nil {
|
||||
if nodeErr != nil || !isTunnelProbeNodeOnline(targetNode) {
|
||||
snap.ErrorMessage = "节点 " + target.NodeName + " 不可用"
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
break
|
||||
}
|
||||
|
||||
|
||||
fromNode, _ := h.getNodeRecord(source.NodeID)
|
||||
targetIP, targetPort, resolveErr := resolveChainProbeTarget(fromNode, targetNode, target.Port, ipPreference, target.ConnectIP)
|
||||
if resolveErr != nil {
|
||||
@@ -295,25 +395,25 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
break
|
||||
}
|
||||
|
||||
hop.TargetIP = targetIP
|
||||
hop.TargetPort = targetPort
|
||||
|
||||
lat, loss, err := p.tcpPingNode(source.NodeID, targetIP, targetPort, options)
|
||||
lat, loss, err := roundPinger(source.NodeID, targetIP, targetPort, options)
|
||||
if err == nil {
|
||||
hop.Latency = lat
|
||||
hop.Loss = loss
|
||||
totalLat += lat
|
||||
remainingSuccessProb *= (1.0 - loss/100.0)
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
} else {
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
if snap.ErrorMessage == "" {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
@@ -328,17 +428,11 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
snap.EntryToExitLatency = -1
|
||||
snap.EntryToExitLoss = 100
|
||||
}
|
||||
|
||||
if len(hops) > 0 {
|
||||
if b, err := json.Marshal(hops); err == nil {
|
||||
snap.ChainDetails = string(b)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Exit → Bing
|
||||
if len(outNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(outNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
if exitOnline {
|
||||
lat, loss, err := roundPinger(exit.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -352,9 +446,9 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
|
||||
snap.Success = probeOK
|
||||
default:
|
||||
// Unknown type: entry → Bing
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(inNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
// Unknown type: entry → public probe target.
|
||||
if entryOnline {
|
||||
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -362,12 +456,263 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
} else {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
} else {
|
||||
snap.ErrorMessage = "入口节点均不在线"
|
||||
}
|
||||
}
|
||||
|
||||
candidateHops := p.probeTunnelCandidateHops(
|
||||
tunnel.Type,
|
||||
inNodes,
|
||||
midNodesGrouped,
|
||||
outNodes,
|
||||
selectedNodeIDs,
|
||||
ipPreference,
|
||||
options,
|
||||
probeTarget,
|
||||
roundPinger,
|
||||
)
|
||||
if len(primaryHops) > 0 || len(candidateHops) > 0 {
|
||||
details := tunnelQualityChainDetails{
|
||||
PrimaryPath: primaryHops,
|
||||
CandidateHops: candidateHops,
|
||||
}
|
||||
if b, err := json.Marshal(details); err == nil {
|
||||
snap.ChainDetails = string(b)
|
||||
}
|
||||
}
|
||||
|
||||
p.storeResult(snap)
|
||||
}
|
||||
|
||||
func tunnelQualityGroupKey(role string, index int) string {
|
||||
return fmt.Sprintf("%s:%d", role, index)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeTunnelCandidateHops(
|
||||
tunnelType int,
|
||||
inNodes []chainNodeRecord,
|
||||
chainHops [][]chainNodeRecord,
|
||||
outNodes []chainNodeRecord,
|
||||
selectedNodeIDs map[string]int64,
|
||||
ipPreference string,
|
||||
options diagnosisExecOptions,
|
||||
probeTarget tunnelProbeTarget,
|
||||
ping bestExitProbeFunc,
|
||||
) []TunnelQualityCandidateHop {
|
||||
if p == nil || p.handler == nil || ping == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if tunnelType != 2 {
|
||||
return p.probePublicTargetCandidates("entry", 0, inNodes, selectedNodeIDs, options, probeTarget, ping)
|
||||
}
|
||||
|
||||
groups := make([]tunnelQualityCandidateGroup, 0, 2+len(chainHops))
|
||||
groups = append(groups, tunnelQualityCandidateGroup{role: "entry", roleIndex: 0, nodes: inNodes})
|
||||
for i, hop := range chainHops {
|
||||
groups = append(groups, tunnelQualityCandidateGroup{role: "middle", roleIndex: i, nodes: hop})
|
||||
}
|
||||
groups = append(groups, tunnelQualityCandidateGroup{role: "exit", roleIndex: 0, nodes: outNodes})
|
||||
|
||||
var items []TunnelQualityCandidateHop
|
||||
for i := 0; i < len(groups)-1; i++ {
|
||||
items = append(items, p.probeCandidateGroupLinks(
|
||||
groups[i],
|
||||
groups[i+1],
|
||||
i,
|
||||
selectedNodeIDs,
|
||||
ipPreference,
|
||||
options,
|
||||
ping,
|
||||
)...)
|
||||
}
|
||||
items = append(items, p.probePublicTargetCandidates(
|
||||
"exit",
|
||||
0,
|
||||
outNodes,
|
||||
selectedNodeIDs,
|
||||
options,
|
||||
probeTarget,
|
||||
ping,
|
||||
)...)
|
||||
return items
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeCandidateGroupLinks(
|
||||
fromGroup tunnelQualityCandidateGroup,
|
||||
toGroup tunnelQualityCandidateGroup,
|
||||
hopIndex int,
|
||||
selectedNodeIDs map[string]int64,
|
||||
ipPreference string,
|
||||
options diagnosisExecOptions,
|
||||
ping bestExitProbeFunc,
|
||||
) []TunnelQualityCandidateHop {
|
||||
items := make([]TunnelQualityCandidateHop, 0, len(fromGroup.nodes)*len(toGroup.nodes))
|
||||
for _, source := range fromGroup.nodes {
|
||||
for _, target := range toGroup.nodes {
|
||||
item := TunnelQualityCandidateHop{
|
||||
TunnelQualityHop: TunnelQualityHop{
|
||||
FromNodeID: source.NodeID,
|
||||
FromNodeName: source.NodeName,
|
||||
ToNodeID: target.NodeID,
|
||||
ToNodeName: target.NodeName,
|
||||
Latency: -1,
|
||||
Loss: 100,
|
||||
},
|
||||
FromRole: fromGroup.role,
|
||||
ToRole: toGroup.role,
|
||||
HopIndex: hopIndex,
|
||||
Selected: selectedNodeIDs[tunnelQualityGroupKey(fromGroup.role, fromGroup.roleIndex)] == source.NodeID &&
|
||||
selectedNodeIDs[tunnelQualityGroupKey(toGroup.role, toGroup.roleIndex)] == target.NodeID,
|
||||
}
|
||||
|
||||
sourceNode, sourceErr := p.handler.getNodeRecord(source.NodeID)
|
||||
if sourceErr != nil || !isTunnelProbeNodeOnline(sourceNode) {
|
||||
item.ErrorMessage = "来源节点不在线"
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
targetNode, targetErr := p.handler.getNodeRecord(target.NodeID)
|
||||
if targetErr != nil || !isTunnelProbeNodeOnline(targetNode) {
|
||||
item.ErrorMessage = "目标节点不在线"
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
|
||||
targetIP, targetPort, resolveErr := resolveChainProbeTarget(sourceNode, targetNode, target.Port, ipPreference, target.ConnectIP)
|
||||
if resolveErr != nil {
|
||||
item.ErrorMessage = resolveErr.Error()
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
item.TargetIP = targetIP
|
||||
item.TargetPort = targetPort
|
||||
latency, loss, probeErr := ping(source.NodeID, targetIP, targetPort, options)
|
||||
if probeErr != nil {
|
||||
item.ErrorMessage = probeErr.Error()
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
item.Latency = latency
|
||||
item.Loss = loss
|
||||
items = append(items, item)
|
||||
}
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probePublicTargetCandidates(
|
||||
fromRole string,
|
||||
fromIndex int,
|
||||
nodes []chainNodeRecord,
|
||||
selectedNodeIDs map[string]int64,
|
||||
options diagnosisExecOptions,
|
||||
probeTarget tunnelProbeTarget,
|
||||
ping bestExitProbeFunc,
|
||||
) []TunnelQualityCandidateHop {
|
||||
items := make([]TunnelQualityCandidateHop, 0, len(nodes))
|
||||
for _, source := range nodes {
|
||||
item := TunnelQualityCandidateHop{
|
||||
TunnelQualityHop: TunnelQualityHop{
|
||||
FromNodeID: source.NodeID,
|
||||
FromNodeName: source.NodeName,
|
||||
ToNodeName: formatTunnelProbeTarget(probeTarget),
|
||||
Latency: -1,
|
||||
Loss: 100,
|
||||
TargetIP: probeTarget.Host,
|
||||
TargetPort: probeTarget.Port,
|
||||
},
|
||||
FromRole: fromRole,
|
||||
ToRole: "target",
|
||||
HopIndex: fromIndex,
|
||||
Selected: selectedNodeIDs[tunnelQualityGroupKey(fromRole, fromIndex)] == source.NodeID,
|
||||
}
|
||||
sourceNode, sourceErr := p.handler.getNodeRecord(source.NodeID)
|
||||
if sourceErr != nil || !isTunnelProbeNodeOnline(sourceNode) {
|
||||
item.ErrorMessage = "来源节点不在线"
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
latency, loss, probeErr := ping(source.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if probeErr != nil {
|
||||
item.ErrorMessage = probeErr.Error()
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
item.Latency = latency
|
||||
item.Loss = loss
|
||||
items = append(items, item)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func isTunnelProbeNodeOnline(node *nodeRecord) bool {
|
||||
return node != nil && (node.IsRemote == 1 || node.Status == 1)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) firstOnlineChainNode(nodes []chainNodeRecord) (chainNodeRecord, *nodeRecord, bool) {
|
||||
if p == nil || p.handler == nil {
|
||||
return chainNodeRecord{}, nil, false
|
||||
}
|
||||
for _, candidate := range nodes {
|
||||
node, err := p.handler.getNodeRecord(candidate.NodeID)
|
||||
if err == nil && isTunnelProbeNodeOnline(node) {
|
||||
return candidate, node, true
|
||||
}
|
||||
}
|
||||
return chainNodeRecord{}, nil, false
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chainNodeRecord, chainHops [][]chainNodeRecord, outNodes []chainNodeRecord, ipPreference string, options diagnosisExecOptions, probeTarget tunnelProbeTarget, roundPinger bestExitProbeFunc) {
|
||||
if p == nil || p.handler == nil || p.handler.bestExit == nil || len(outNodes) <= 1 {
|
||||
return
|
||||
}
|
||||
if !isBestTunnelStrategy(outNodes[0].Strategy) {
|
||||
return
|
||||
}
|
||||
owners := bestExitChainOwners(inNodes, chainHops)
|
||||
if len(owners) == 0 {
|
||||
return
|
||||
}
|
||||
nodeMap := make(map[int64]*nodeRecord, len(owners)+len(outNodes))
|
||||
for _, owner := range owners {
|
||||
if node, err := p.handler.getNodeRecord(owner.NodeID); err == nil && node != nil {
|
||||
nodeMap[owner.NodeID] = node
|
||||
}
|
||||
}
|
||||
for _, exit := range outNodes {
|
||||
if node, err := p.handler.getNodeRecord(exit.NodeID); err == nil && node != nil {
|
||||
nodeMap[exit.NodeID] = node
|
||||
}
|
||||
}
|
||||
for _, owner := range owners {
|
||||
if nodeMap[owner.NodeID] == nil {
|
||||
continue
|
||||
}
|
||||
key := bestExitOwnerKey{TunnelID: tunnelID, OwnerNodeID: owner.NodeID}
|
||||
p.handler.bestExit.ensureApplied(key, outNodes[0].NodeID, time.Now())
|
||||
scores := evaluateBestExitOwner(owner, outNodes, nodeMap, ipPreference, options, probeTarget, roundPinger)
|
||||
decision := p.handler.bestExit.observeScores(key, scores, time.Now())
|
||||
if decision.Switch {
|
||||
now := time.Now()
|
||||
if err := p.handler.applyBestExitChainOrder(tunnelID, owner.NodeID, outNodes, decision.Scores, ipPreference); err != nil {
|
||||
log.Printf("best_exit: switch apply failed tunnel=%d owner=%d exit=%d err=%v", tunnelID, owner.NodeID, decision.ExitNodeID, err)
|
||||
p.handler.bestExit.recordApplyFailure(key, decision.ExitNodeID, now)
|
||||
continue
|
||||
}
|
||||
p.handler.bestExit.setApplied(key, decision.ExitNodeID, time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) pingNode(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
if p != nil && p.probeNode != nil {
|
||||
return p.probeNode(nodeID, ip, port, options)
|
||||
}
|
||||
return p.tcpPingNode(nodeID, ip, port, options)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) tcpPingNode(nodeID int64, ip string, port int, options diagnosisExecOptions) (latency float64, loss float64, err error) {
|
||||
h := p.handler
|
||||
if h == nil {
|
||||
@@ -378,6 +723,9 @@ func (p *tunnelQualityProber) tcpPingNode(nodeID int64, ip string, port int, opt
|
||||
if nodeErr != nil {
|
||||
return 0, 100, nodeErr
|
||||
}
|
||||
if !isTunnelProbeNodeOnline(node) {
|
||||
return 0, 100, errors.New("节点不在线")
|
||||
}
|
||||
|
||||
var pingData map[string]interface{}
|
||||
var pingErr error
|
||||
|
||||
@@ -0,0 +1,247 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 77, "quality-target", "speed.example.com", 8443)
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(30, 'exit-a', 'exit-secret', '10.0.0.30', '10.0.0.30', '', '30000-30010', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
|
||||
`, time.Now().UnixMilli(), time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("insert exit node: %v", err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(77, '3', 30, 30001, 'round', 1, 'tls')
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert exit chain: %v", err)
|
||||
}
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
var calls []string
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
if options.pingCount != 1 {
|
||||
t.Fatalf("expected real-time quality probe count 1, got %d", options.pingCount)
|
||||
}
|
||||
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
|
||||
return 10, 0, nil
|
||||
}
|
||||
p.probeTunnel(77)
|
||||
|
||||
if !slices.Contains(calls, "10|speed.example.com|8443") {
|
||||
t.Fatalf("expected type 1 public probe from entry to configured target, calls=%+v", calls)
|
||||
}
|
||||
if slices.Contains(calls, "30|speed.example.com|8443") {
|
||||
t.Fatalf("did not expect type 1 public probe from exit node, calls=%+v", calls)
|
||||
}
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].ProbeTargetHost != "speed.example.com" || snaps[0].ProbeTargetPort != 8443 {
|
||||
t.Fatalf("unexpected snapshot target metadata: %+v", snaps[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberSkipsAllOfflineExits(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedQualityForwardTunnel(t, h, 81, []int{0, 0, 0})
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
probeCalls := 0
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
probeCalls++
|
||||
return 0, 100, fmt.Errorf("unexpected probe node=%d target=%s:%d", nodeID, ip, port)
|
||||
}
|
||||
p.probeTunnel(81)
|
||||
|
||||
if probeCalls != 0 {
|
||||
t.Fatalf("expected no TCP probes when all exits are offline, got %d", probeCalls)
|
||||
}
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].Success || snaps[0].ErrorMessage != "出口节点均不在线" {
|
||||
t.Fatalf("expected offline exit snapshot, got %+v", snaps[0])
|
||||
}
|
||||
if snaps[0].EntryToExitLoss != 100 {
|
||||
t.Fatalf("expected 100%% entry-to-exit loss, got %+v", snaps[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberUsesOnlineBackupExit(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedQualityForwardTunnel(t, h, 82, []int{0, 1})
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
var calls []string
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
if options.pingCount != 1 {
|
||||
t.Fatalf("expected real-time quality probe count 1, got %d", options.pingCount)
|
||||
}
|
||||
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
|
||||
return 10, 0, nil
|
||||
}
|
||||
p.probeTunnel(82)
|
||||
|
||||
if slices.Contains(calls, "10|10.0.0.30|30030") {
|
||||
t.Fatalf("did not expect probe to offline primary exit, calls=%+v", calls)
|
||||
}
|
||||
if !slices.Contains(calls, "10|10.0.0.31|30031") {
|
||||
t.Fatalf("expected entry probe to online backup exit, calls=%+v", calls)
|
||||
}
|
||||
if !slices.Contains(calls, "31|www.bing.com|443") {
|
||||
t.Fatalf("expected public probe from online backup exit, calls=%+v", calls)
|
||||
}
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 || !snaps[0].Success {
|
||||
t.Fatalf("expected successful backup exit snapshot, got %+v", snaps)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberReportsAllExitCandidateLatencies(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedQualityForwardTunnel(t, h, 83, []int{1, 1})
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
switch fmt.Sprintf("%d|%s|%d", nodeID, ip, port) {
|
||||
case "10|10.0.0.30|30030":
|
||||
return 20, 0, nil
|
||||
case "10|10.0.0.31|30031":
|
||||
return 35, 0, nil
|
||||
case "30|www.bing.com|443":
|
||||
return 50, 0, nil
|
||||
case "31|www.bing.com|443":
|
||||
return 65, 0, nil
|
||||
default:
|
||||
return 0, 100, fmt.Errorf("unexpected probe node=%d target=%s:%d", nodeID, ip, port)
|
||||
}
|
||||
}
|
||||
p.probeTunnel(83)
|
||||
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].EntryToExitLatency != 20 || snaps[0].ExitToBingLatency != 50 {
|
||||
t.Fatalf("expected primary path metrics to remain unchanged, got %+v", snaps[0])
|
||||
}
|
||||
|
||||
var details tunnelQualityChainDetails
|
||||
if err := json.Unmarshal([]byte(snaps[0].ChainDetails), &details); err != nil {
|
||||
t.Fatalf("decode chain details: %v", err)
|
||||
}
|
||||
assertCandidateHop := func(fromID, toID int64, latency float64, selected bool) {
|
||||
t.Helper()
|
||||
for _, hop := range details.CandidateHops {
|
||||
if hop.FromNodeID == fromID && hop.ToNodeID == toID {
|
||||
if hop.Latency != latency || hop.Selected != selected || hop.ErrorMessage != "" {
|
||||
t.Fatalf("unexpected candidate hop: %+v", hop)
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatalf("candidate hop %d -> %d not found in %+v", fromID, toID, details.CandidateHops)
|
||||
}
|
||||
assertCandidateHop(10, 30, 20, true)
|
||||
assertCandidateHop(10, 31, 35, false)
|
||||
assertCandidateHop(30, 0, 50, true)
|
||||
assertCandidateHop(31, 0, 65, false)
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberStoresProbeTargetWhenChainIncomplete(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 78, "quality-target-incomplete", "speed.example.com", 8443)
|
||||
if err := h.repo.DB().Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, 78).Error; err != nil {
|
||||
t.Fatalf("delete chain rows: %v", err)
|
||||
}
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
p.probeTunnel(78)
|
||||
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].ErrorMessage == "" {
|
||||
t.Fatalf("expected incomplete chain error, got %+v", snaps[0])
|
||||
}
|
||||
if snaps[0].ProbeTargetHost != "speed.example.com" || snaps[0].ProbeTargetPort != 8443 {
|
||||
t.Fatalf("unexpected snapshot target metadata: %+v", snaps[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberUsesConfiguredInterval(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
if err := h.repo.UpsertConfig("monitor_tunnel_quality_interval_sec", "15", time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("upsert interval config: %v", err)
|
||||
}
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
if got := p.probeInterval(); got != 15*time.Second {
|
||||
t.Fatalf("probe interval = %s, want 15s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberConfigNotificationIsCoalesced(t *testing.T) {
|
||||
p := newTunnelQualityProber(nil)
|
||||
p.NotifyConfigChanged()
|
||||
p.NotifyConfigChanged()
|
||||
|
||||
if got := len(p.wake); got != 1 {
|
||||
t.Fatalf("wake notifications = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelQualityProbeIntervalConfigValue(t *testing.T) {
|
||||
got, err := normalizeAndValidateConfigValue("monitor_tunnel_quality_interval_sec", " 15 ")
|
||||
if err != nil || got != "15" {
|
||||
t.Fatalf("normalize interval = %q, %v", got, err)
|
||||
}
|
||||
if _, err := normalizeAndValidateConfigValue("monitor_tunnel_quality_interval_sec", "0"); err == nil {
|
||||
t.Fatalf("expected invalid interval to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func seedQualityForwardTunnel(t *testing.T, h *Handler, tunnelID int64, exitStatuses []int) {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference, probe_target_host, probe_target_port)
|
||||
VALUES(?, ?, 1, 2, 'tls', 1, ?, ?, 1, ?, '', '', 0)
|
||||
`, tunnelID, fmt.Sprintf("quality-forward-%d", tunnelID), now, now, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert forwarding tunnel: %v", err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, '1', 10, 30001, 'fifo', 1, 'tls')
|
||||
`, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert entry chain: %v", err)
|
||||
}
|
||||
for i, status := range exitStatuses {
|
||||
nodeID := int64(30 + i)
|
||||
port := 30030 + i
|
||||
ip := fmt.Sprintf("10.0.0.%d", nodeID)
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, '', '30000-30100', '', 'v1', 1, 1, 1, ?, ?, ?, '[::]', '[::]', 0)
|
||||
`, nodeID, fmt.Sprintf("exit-%d", i+1), fmt.Sprintf("exit-secret-%d", i+1), ip, ip, now, now, status).Error; err != nil {
|
||||
t.Fatalf("insert exit node %d: %v", nodeID, err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, '3', ?, ?, 'fifo', ?, 'tls')
|
||||
`, tunnelID, nodeID, port, i+1).Error; err != nil {
|
||||
t.Fatalf("insert exit chain %d: %v", nodeID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -13,6 +13,13 @@ import (
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
// failedForward tracks a forward that failed redeployment, for retry.
|
||||
type failedForward struct {
|
||||
id int64
|
||||
forward *forwardRecord
|
||||
err error
|
||||
}
|
||||
|
||||
const (
|
||||
githubRepo = "Sagit-chu/flvx"
|
||||
githubAPIBase = "https://api.github.com"
|
||||
@@ -32,6 +39,8 @@ var (
|
||||
testKeywordPattern = regexp.MustCompile(`(?i)(alpha|beta|rc)`)
|
||||
)
|
||||
|
||||
const nodeOnlineRedeployCooldown = 30 * time.Second
|
||||
|
||||
type githubRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Name string `json:"name"`
|
||||
@@ -389,22 +398,123 @@ func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool {
|
||||
}
|
||||
|
||||
func (h *Handler) onNodeOnline(nodeID int64) {
|
||||
h.consumeNodePendingUpgradeRedeploy(nodeID)
|
||||
h.redeployNodeRuntimeAfterUpgrade(nodeID)
|
||||
if !h.startNodeOnlineRedeploy(nodeID, time.Now()) {
|
||||
return
|
||||
}
|
||||
defer h.finishNodeOnlineRedeploy(nodeID)
|
||||
|
||||
// Reconcile node runtime on the first reconnect, but suppress rapid flapping
|
||||
// so websocket churn does not trigger repeated full redeploy storms.
|
||||
if !h.redeployNodeRuntimeAfterUpgrade(nodeID) {
|
||||
h.markNodePendingUpgradeRedeploy(nodeID)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
func (h *Handler) startNodeOnlineRedeploy(nodeID int64, now time.Time) bool {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return false
|
||||
}
|
||||
if now.IsZero() {
|
||||
now = time.Now()
|
||||
}
|
||||
|
||||
h.upgradeMu.Lock()
|
||||
defer h.upgradeMu.Unlock()
|
||||
if h.pendingUpgradeRedeploy == nil {
|
||||
h.pendingUpgradeRedeploy = make(map[int64]struct{})
|
||||
}
|
||||
if h.nodeOnlineRedeployAt == nil {
|
||||
h.nodeOnlineRedeployAt = make(map[int64]time.Time)
|
||||
}
|
||||
if h.nodeOnlineRedeployQueued == nil {
|
||||
h.nodeOnlineRedeployQueued = make(map[int64]struct{})
|
||||
}
|
||||
if h.nodeOnlineRedeploying == nil {
|
||||
h.nodeOnlineRedeploying = make(map[int64]struct{})
|
||||
}
|
||||
|
||||
_, pendingUpgrade := h.pendingUpgradeRedeploy[nodeID]
|
||||
lastRedeployAt := h.nodeOnlineRedeployAt[nodeID]
|
||||
_, inFlight := h.nodeOnlineRedeploying[nodeID]
|
||||
if fireAt, start := nextNodeOnlineRedeployFireAt(lastRedeployAt, now, pendingUpgrade, inFlight); !start {
|
||||
h.queueNodeOnlineRedeployLocked(nodeID, fireAt)
|
||||
return false
|
||||
}
|
||||
|
||||
delete(h.pendingUpgradeRedeploy, nodeID)
|
||||
h.nodeOnlineRedeployAt[nodeID] = now
|
||||
h.nodeOnlineRedeploying[nodeID] = struct{}{}
|
||||
return true
|
||||
}
|
||||
|
||||
func nextNodeOnlineRedeployFireAt(lastRedeployAt, now time.Time, pendingUpgrade bool, inFlight bool) (time.Time, bool) {
|
||||
if now.IsZero() {
|
||||
now = time.Now()
|
||||
}
|
||||
if inFlight {
|
||||
fireAt := now.Add(nodeOnlineRedeployCooldown)
|
||||
if !lastRedeployAt.IsZero() {
|
||||
cooldownAt := lastRedeployAt.Add(nodeOnlineRedeployCooldown)
|
||||
if cooldownAt.After(now) {
|
||||
fireAt = cooldownAt
|
||||
}
|
||||
}
|
||||
return fireAt, false
|
||||
}
|
||||
if !pendingUpgrade && !lastRedeployAt.IsZero() && now.Sub(lastRedeployAt) < nodeOnlineRedeployCooldown {
|
||||
return lastRedeployAt.Add(nodeOnlineRedeployCooldown), false
|
||||
}
|
||||
return time.Time{}, true
|
||||
}
|
||||
|
||||
func (h *Handler) queueNodeOnlineRedeployLocked(nodeID int64, fireAt time.Time) {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return
|
||||
}
|
||||
if h.nodeOnlineRedeployQueued == nil {
|
||||
h.nodeOnlineRedeployQueued = make(map[int64]struct{})
|
||||
}
|
||||
if _, queued := h.nodeOnlineRedeployQueued[nodeID]; queued {
|
||||
return
|
||||
}
|
||||
if fireAt.IsZero() {
|
||||
fireAt = time.Now().Add(nodeOnlineRedeployCooldown)
|
||||
}
|
||||
delay := time.Until(fireAt)
|
||||
if delay < 0 {
|
||||
delay = 0
|
||||
}
|
||||
h.nodeOnlineRedeployQueued[nodeID] = struct{}{}
|
||||
time.AfterFunc(delay, func() {
|
||||
h.upgradeMu.Lock()
|
||||
delete(h.nodeOnlineRedeployQueued, nodeID)
|
||||
h.upgradeMu.Unlock()
|
||||
h.onNodeOnline(nodeID)
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) finishNodeOnlineRedeploy(nodeID int64) {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return
|
||||
}
|
||||
h.upgradeMu.Lock()
|
||||
delete(h.nodeOnlineRedeploying, nodeID)
|
||||
h.upgradeMu.Unlock()
|
||||
}
|
||||
|
||||
func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) bool {
|
||||
tunnelIDs, err := h.repo.ListActiveTunnelIDsByNode(nodeID)
|
||||
if err != nil {
|
||||
fmt.Printf("post-upgrade redeploy: list tunnels for node %d failed: %v\n", nodeID, err)
|
||||
return
|
||||
return false
|
||||
}
|
||||
forwardIDs, err := h.repo.ListForwardIDsByNode(nodeID)
|
||||
if err != nil {
|
||||
fmt.Printf("post-upgrade redeploy: list forwards for node %d failed: %v\n", nodeID, err)
|
||||
return
|
||||
return false
|
||||
}
|
||||
|
||||
// First pass: deploy everything
|
||||
tunnelFailed := make(map[int64]struct{})
|
||||
for _, tunnelID := range tunnelIDs {
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
@@ -413,6 +523,9 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
}
|
||||
}
|
||||
|
||||
// Collect forwards that failed independently (not skipped due to tunnel failure)
|
||||
var failedForwards []failedForward
|
||||
|
||||
for _, forwardID := range forwardIDs {
|
||||
forward, getErr := h.getForwardRecord(forwardID)
|
||||
if getErr != nil || forward == nil {
|
||||
@@ -422,7 +535,87 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
failedForwards = append(failedForwards, failedForward{id: forwardID, forward: forward, err: err})
|
||||
fmt.Printf("post-upgrade redeploy: forward %d failed on node %d: %v\n", forwardID, nodeID, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Retry failed items with exponential backoff (max 3 attempts)
|
||||
return h.retryFailedRedeploys(nodeID, tunnelFailed, failedForwards)
|
||||
}
|
||||
|
||||
// isRetryableError returns true if the error looks transient and worth retrying.
|
||||
func isRetryableError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
// Skip non-retryable errors: not-found, already-exists, validation errors
|
||||
if strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") {
|
||||
return false
|
||||
}
|
||||
if strings.Contains(msg, "already exists") || strings.Contains(msg, "已存在") {
|
||||
return false
|
||||
}
|
||||
// Everything else (timeout, connection lost, port in use, etc.) is retryable
|
||||
return true
|
||||
}
|
||||
|
||||
// retryFailedRedeploys retries failed tunnels and forwards with exponential backoff.
|
||||
func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]struct{}, failedForwards []failedForward) bool {
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
const maxRetries = 3
|
||||
baseDelay := time.Second
|
||||
|
||||
for attempt := 1; attempt <= maxRetries; attempt++ {
|
||||
delay := baseDelay * time.Duration(1<<uint(attempt-1)) // 1s, 2s, 4s
|
||||
time.Sleep(delay)
|
||||
|
||||
// Retry failed tunnels
|
||||
for tunnelID := range tunnelFailed {
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err == nil {
|
||||
delete(tunnelFailed, tunnelID)
|
||||
fmt.Printf("post-upgrade redeploy retry: tunnel %d succeeded on node %d (attempt %d)\n", tunnelID, nodeID, attempt)
|
||||
} else if !isRetryableError(err) {
|
||||
delete(tunnelFailed, tunnelID) // Non-retryable, don't retry again
|
||||
} else {
|
||||
fmt.Printf("post-upgrade redeploy retry: tunnel %d still failing on node %d (attempt %d): %v\n", tunnelID, nodeID, attempt, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Retry failed forwards
|
||||
var stillFailed []failedForward
|
||||
for _, ff := range failedForwards {
|
||||
if _, skipped := tunnelFailed[ff.forward.TunnelID]; skipped {
|
||||
stillFailed = append(stillFailed, ff) // Tunnel still failed, skip forward
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(ff.forward, "UpdateService", true); err == nil {
|
||||
fmt.Printf("post-upgrade redeploy retry: forward %d succeeded on node %d (attempt %d)\n", ff.id, nodeID, attempt)
|
||||
} else if !isRetryableError(err) {
|
||||
// Non-retryable, drop it
|
||||
} else {
|
||||
stillFailed = append(stillFailed, ff)
|
||||
fmt.Printf("post-upgrade redeploy retry: forward %d still failing on node %d (attempt %d): %v\n", ff.id, nodeID, attempt, err)
|
||||
}
|
||||
}
|
||||
failedForwards = stillFailed
|
||||
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
fmt.Printf("post-upgrade redeploy retry: all items recovered on node %d\n", nodeID)
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// Final summary
|
||||
for tunnelID := range tunnelFailed {
|
||||
fmt.Printf("post-upgrade redeploy: tunnel %d permanently failed on node %d after retries\n", tunnelID, nodeID)
|
||||
}
|
||||
for _, ff := range failedForwards {
|
||||
fmt.Printf("post-upgrade redeploy: forward %d permanently failed on node %d after retries\n", ff.id, nodeID)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestStartNodeOnlineRedeploySkipsRecentReconnects(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
|
||||
if h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
|
||||
t.Fatalf("expected recent reconnect to skip redeploy")
|
||||
}
|
||||
if h.consumeNodePendingUpgradeRedeploy(54) {
|
||||
t.Fatalf("did not expect pending upgrade marker to be consumed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartNodeOnlineRedeployAllowsPendingUpgradeDuringCooldown(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
h.markNodePendingUpgradeRedeploy(54)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
|
||||
t.Fatalf("expected pending upgrade reconnect to bypass cooldown")
|
||||
}
|
||||
if h.consumeNodePendingUpgradeRedeploy(54) {
|
||||
t.Fatalf("expected pending upgrade marker to be consumed during redeploy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartNodeOnlineRedeployQueuesCooldownReconnect(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
|
||||
if h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
|
||||
t.Fatalf("expected cooldown reconnect to skip immediate redeploy")
|
||||
}
|
||||
if _, queued := h.nodeOnlineRedeployQueued[54]; !queued {
|
||||
t.Fatalf("expected cooldown reconnect to queue a follow-up redeploy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartNodeOnlineRedeployKeepsPendingUpgradeWhileInFlight(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.markNodePendingUpgradeRedeploy(54)
|
||||
|
||||
if h.startNodeOnlineRedeploy(54, now.Add(time.Second)) {
|
||||
t.Fatalf("expected in-flight redeploy to suppress parallel restart")
|
||||
}
|
||||
if !h.consumeNodePendingUpgradeRedeploy(54) {
|
||||
t.Fatalf("expected pending upgrade marker to remain for the next retry")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
}
|
||||
|
||||
func TestNextNodeOnlineRedeployFireAtDefersExpiredInFlightReconnect(t *testing.T) {
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
last := now.Add(-nodeOnlineRedeployCooldown - 5*time.Second)
|
||||
|
||||
fireAt, start := nextNodeOnlineRedeployFireAt(last, now, false, true)
|
||||
if start {
|
||||
t.Fatalf("expected in-flight reconnect to queue instead of starting immediately")
|
||||
}
|
||||
|
||||
want := now.Add(nodeOnlineRedeployCooldown)
|
||||
if !fireAt.Equal(want) {
|
||||
t.Fatalf("expected queued reconnect at %s, got %s", want, fireAt)
|
||||
}
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package middleware
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
@@ -14,7 +15,8 @@ type contextKey string
|
||||
const ClaimsContextKey contextKey = "claims"
|
||||
|
||||
type AuthOptions struct {
|
||||
JWTSecret string
|
||||
JWTSecret string
|
||||
GetUserAuthState func(userID int64) (*auth.UserAuthState, error)
|
||||
}
|
||||
|
||||
func JWT(opts AuthOptions) func(http.Handler) http.Handler {
|
||||
@@ -32,16 +34,45 @@ func JWT(opts AuthOptions) func(http.Handler) http.Handler {
|
||||
|
||||
token := strings.TrimSpace(r.Header.Get("Authorization"))
|
||||
if token == "" {
|
||||
if allowsOptionalAuth(r.URL.Path) {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(401, "未登录或token已过期"))
|
||||
return
|
||||
}
|
||||
|
||||
claims, ok := auth.ValidateToken(token, opts.JWTSecret)
|
||||
if !ok {
|
||||
if allowsOptionalAuth(r.URL.Path) {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
|
||||
if opts.GetUserAuthState != nil {
|
||||
userID, err := strconv.ParseInt(claims.Sub, 10, 64)
|
||||
if err != nil {
|
||||
if allowsOptionalAuth(r.URL.Path) {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
state, err := opts.GetUserAuthState(userID)
|
||||
if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt {
|
||||
if allowsOptionalAuth(r.URL.Path) {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if requiresAdmin(r.URL.Path) && claims.RoleID != 0 {
|
||||
response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作"))
|
||||
return
|
||||
@@ -69,6 +100,10 @@ func RequireAdmin(next http.Handler) http.Handler {
|
||||
})
|
||||
}
|
||||
|
||||
func allowsOptionalAuth(path string) bool {
|
||||
return path == "/api/v1/config/get"
|
||||
}
|
||||
|
||||
func shouldSkip(path string) bool {
|
||||
switch {
|
||||
case strings.HasPrefix(path, "/flow/"):
|
||||
@@ -78,9 +113,11 @@ func shouldSkip(path string) bool {
|
||||
case strings.HasPrefix(path, "/api/v1/captcha/"):
|
||||
return true
|
||||
case path == "/api/v1/config/get":
|
||||
return true
|
||||
return false
|
||||
case path == "/api/v1/user/login":
|
||||
return true
|
||||
case path == "/api/v1/public/config/get":
|
||||
return true
|
||||
case path == "/api/v1/federation/connect":
|
||||
return true
|
||||
case path == "/api/v1/federation/tunnel/create":
|
||||
@@ -105,6 +142,10 @@ func requiresAdmin(path string) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/system/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/group/") {
|
||||
return true
|
||||
}
|
||||
@@ -141,6 +182,8 @@ func requiresAdmin(path string) bool {
|
||||
return true
|
||||
case "/api/v1/config/update", "/api/v1/config/update-single":
|
||||
return true
|
||||
case "/api/v1/license/activate":
|
||||
return true
|
||||
case "/api/v1/announcement/update":
|
||||
return true
|
||||
default:
|
||||
|
||||
@@ -0,0 +1,276 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestJWTRejectsPasswordChangedToken(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
claims, err := auth.ParseClaims(token, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("parse claims: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs + 1}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTAcceptsCurrentUserState(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
claims, err := auth.ParseClaims(token, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("parse claims: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs - 1}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
}
|
||||
|
||||
func TestJWTRejectsDisabledUserToken(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 0, PasswordChangedAt: time.Now().Unix()}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsPasswordChangedAtSameMillisecond(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
claims, err := auth.ParseClaims(token, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("parse claims: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsRoleMismatch(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: 0}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsMissingUserState(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return nil, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsAuthStateLookupError(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return nil, errors.New("boom")
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestShouldSkipDoesNotBypassConfigGet(t *testing.T) {
|
||||
if shouldSkip("/api/v1/config/get") {
|
||||
t.Fatal("expected /api/v1/config/get to require auth")
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldSkipBypassesPublicConfigGet(t *testing.T) {
|
||||
if !shouldSkip("/api/v1/public/config/get") {
|
||||
t.Fatal("expected /api/v1/public/config/get to remain public")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseActivateRequiresAdmin(t *testing.T) {
|
||||
if !requiresAdmin("/api/v1/license/activate") {
|
||||
t.Fatal("expected license activation to require admin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTRejectsNonAdminLicenseActivation(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(2, "regular_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
claims, err := auth.ParseClaims(token, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("parse claims: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: claims.IatMs - 1}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCodeMsg(t, res, 403, "权限不足,仅管理员可操作")
|
||||
}
|
||||
|
||||
func TestJWTExpiresAfterSevenDays(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
claims, err := auth.ParseClaims(token, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("parse claims: %v", err)
|
||||
}
|
||||
if got := claims.Exp - claims.Iat; got != int64(7*24*time.Hour/time.Second) {
|
||||
t.Fatalf("expected 7 day token lifetime, got %d seconds", got)
|
||||
}
|
||||
if claims.IatMs <= 0 {
|
||||
t.Fatalf("expected millisecond issuance time to be populated, got %d", claims.IatMs)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expected {
|
||||
t.Fatalf("expected code %d, got %d", expected, out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func assertAuthDenied(t *testing.T, rec *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
assertCodeMsg(t, rec, 401, "无效的token或token已过期")
|
||||
}
|
||||
|
||||
func assertCodeMsg(t *testing.T, rec *httptest.ResponseRecorder, expectedCode int, expectedMsg string) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expectedCode || out.Msg != expectedMsg {
|
||||
t.Fatalf("expected (%d,%q), got (%d,%q)", expectedCode, expectedMsg, out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
@@ -13,7 +13,7 @@ func NewRouter(h *handler.Handler, jwtSecret string) http.Handler {
|
||||
mux.Handle("/system-info", h.WebSocketHandler())
|
||||
|
||||
wrapped := middleware.Recover(mux)
|
||||
wrapped = middleware.JWT(middleware.AuthOptions{JWTSecret: jwtSecret})(wrapped)
|
||||
wrapped = middleware.JWT(middleware.AuthOptions{JWTSecret: jwtSecret, GetUserAuthState: h.GetUserAuthState})(wrapped)
|
||||
wrapped = middleware.RequestLog(wrapped)
|
||||
wrapped = middleware.CORS(wrapped)
|
||||
return wrapped
|
||||
|
||||
@@ -6,24 +6,53 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
var AccountID string
|
||||
|
||||
type KeygenClient struct {
|
||||
AccountID string
|
||||
Token string
|
||||
BaseURL string
|
||||
HTTPClient *http.Client
|
||||
}
|
||||
|
||||
type APIError struct {
|
||||
Operation string
|
||||
StatusCode int
|
||||
Body string
|
||||
}
|
||||
|
||||
func (e *APIError) Error() string {
|
||||
return fmt.Sprintf("keygen %s failed: status %d, response: %s", e.Operation, e.StatusCode, e.Body)
|
||||
}
|
||||
|
||||
func (e *APIError) HasCode(code string) bool {
|
||||
return e != nil && hasKeygenErrorCode([]byte(e.Body), code)
|
||||
}
|
||||
|
||||
const defaultAPIBaseURL = "https://api.keygen.sh/v1"
|
||||
|
||||
func NewKeygenClient(accountID, token string) *KeygenClient {
|
||||
return &KeygenClient{
|
||||
AccountID: accountID,
|
||||
Token: token,
|
||||
AccountID: accountID,
|
||||
Token: token,
|
||||
BaseURL: defaultAPIBaseURL,
|
||||
HTTPClient: &http.Client{Timeout: 10 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *KeygenClient) apiURL(path string) string {
|
||||
baseURL := strings.TrimRight(c.BaseURL, "/")
|
||||
if baseURL == "" {
|
||||
baseURL = defaultAPIBaseURL
|
||||
}
|
||||
return fmt.Sprintf("%s/accounts/%s/%s", baseURL, c.AccountID, strings.TrimLeft(path, "/"))
|
||||
}
|
||||
|
||||
type ValidateResponse struct {
|
||||
Meta struct {
|
||||
Valid bool `json:"valid"`
|
||||
@@ -35,6 +64,7 @@ type ValidateResponse struct {
|
||||
Expiry string `json:"expiry"`
|
||||
} `json:"attributes"`
|
||||
} `json:"data"`
|
||||
MachineID string `json:"-"`
|
||||
}
|
||||
|
||||
type ActivateMachineRequest struct {
|
||||
@@ -54,8 +84,37 @@ type ActivateMachineRequest struct {
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
type keygenErrorResponse struct {
|
||||
Errors []struct {
|
||||
Code string `json:"code"`
|
||||
} `json:"errors"`
|
||||
}
|
||||
|
||||
type MachineResponse struct {
|
||||
Data struct {
|
||||
ID string `json:"id"`
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
func hasKeygenErrorCode(body []byte, code string) bool {
|
||||
var resp keygenErrorResponse
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
return false
|
||||
}
|
||||
for _, item := range resp.Errors {
|
||||
if strings.EqualFold(strings.TrimSpace(item.Code), code) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string) (*ValidateResponse, error) {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
|
||||
return c.ValidateKeyWithMachine(key, fingerprint, "")
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ValidateKeyWithMachine(key, fingerprint, machineID string) (*ValidateResponse, error) {
|
||||
url := c.apiURL("licenses/actions/validate-key")
|
||||
|
||||
meta := map[string]interface{}{
|
||||
"key": key,
|
||||
@@ -66,6 +125,14 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
|
||||
"fingerprint": fingerprint,
|
||||
}
|
||||
}
|
||||
if machineID != "" {
|
||||
scope, _ := meta["scope"].(map[string]interface{})
|
||||
if scope == nil {
|
||||
scope = make(map[string]interface{})
|
||||
meta["scope"] = scope
|
||||
}
|
||||
scope["machine"] = machineID
|
||||
}
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"meta": meta,
|
||||
@@ -91,7 +158,8 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
var valResp ValidateResponse
|
||||
@@ -102,8 +170,44 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
|
||||
return &valResp, nil
|
||||
}
|
||||
|
||||
func (c *KeygenClient) GetMachineID(fingerprint string) (string, error) {
|
||||
machineURL := c.apiURL("machines/" + url.PathEscape(fingerprint))
|
||||
req, err := http.NewRequest(http.MethodGet, machineURL, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.api+json")
|
||||
if c.Token != "" {
|
||||
if !strings.HasPrefix(c.Token, "Bearer ") && !strings.HasPrefix(c.Token, "License ") {
|
||||
req.Header.Set("Authorization", "License "+c.Token)
|
||||
} else {
|
||||
req.Header.Set("Authorization", c.Token)
|
||||
}
|
||||
}
|
||||
|
||||
resp, err := c.HTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return "", &APIError{Operation: "retrieve machine", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
var machineResp MachineResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&machineResp); err != nil {
|
||||
return "", err
|
||||
}
|
||||
machineID := strings.TrimSpace(machineResp.Data.ID)
|
||||
if machineID == "" {
|
||||
return "", fmt.Errorf("failed to retrieve machine: empty machine id")
|
||||
}
|
||||
return machineID, nil
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
|
||||
url := c.apiURL("licenses/actions/validate-key")
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"meta": map[string]string{
|
||||
@@ -130,7 +234,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
var valResp ValidateResponse
|
||||
@@ -141,8 +246,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
|
||||
return &valResp, nil
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/machines", c.AccountID)
|
||||
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) (string, error) {
|
||||
url := c.apiURL("machines")
|
||||
|
||||
var reqBody ActivateMachineRequest
|
||||
reqBody.Data.Type = "machines"
|
||||
@@ -165,23 +270,24 @@ func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
|
||||
|
||||
resp, err := c.HTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
|
||||
if resp.StatusCode == http.StatusCreated || resp.StatusCode == http.StatusOK {
|
||||
return nil
|
||||
}
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
|
||||
if resp.StatusCode == http.StatusConflict || resp.StatusCode == http.StatusUnprocessableEntity {
|
||||
if strings.Contains(string(body), "FINGERPRINT_TAKEN") || strings.Contains(string(body), "MACHINE_LIMIT_EXCEEDED") {
|
||||
// Machine already registered to this license or limit reached because it's already us.
|
||||
// The subsequent ValidateKey check will determine if the existing machine is actually us.
|
||||
return nil
|
||||
var machineResp MachineResponse
|
||||
if json.Unmarshal(body, &machineResp) == nil {
|
||||
return strings.TrimSpace(machineResp.Data.ID), nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to activate machine: status %d, response: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
if resp.StatusCode == http.StatusUnprocessableEntity && hasKeygenErrorCode(body, "FINGERPRINT_TAKEN") {
|
||||
// Machine activation is idempotent. Keygen scopes fingerprint uniqueness
|
||||
// to the target license, so this means the same machine is already bound.
|
||||
return "", nil
|
||||
}
|
||||
|
||||
return "", &APIError{Operation: "activate machine", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
package license
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateKeyWithMachineSendsFingerprintAndMachineScopes(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var body struct {
|
||||
Meta struct {
|
||||
Scope map[string]string `json:"scope"`
|
||||
} `json:"meta"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||
t.Fatalf("decode request: %v", err)
|
||||
}
|
||||
if body.Meta.Scope["fingerprint"] != "fingerprint" {
|
||||
t.Fatalf("fingerprint scope = %q", body.Meta.Scope["fingerprint"])
|
||||
}
|
||||
if body.Meta.Scope["machine"] != "machine-id" {
|
||||
t.Fatalf("machine scope = %q", body.Meta.Scope["machine"])
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "")
|
||||
client.BaseURL = server.URL
|
||||
validation, err := client.ValidateKeyWithMachine("license-key", "fingerprint", "machine-id")
|
||||
if err != nil || !validation.Meta.Valid {
|
||||
t.Fatalf("ValidateKeyWithMachine() validation=%+v err=%v", validation, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetMachineIDRetrievesMachineByFingerprint(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet || !strings.HasSuffix(r.URL.Path, "/machines/fingerprint") {
|
||||
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
|
||||
}
|
||||
if r.Header.Get("Authorization") != "License license-key" {
|
||||
t.Fatalf("authorization = %q", r.Header.Get("Authorization"))
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
machineID, err := client.GetMachineID("fingerprint")
|
||||
if err != nil || machineID != "machine-id" {
|
||||
t.Fatalf("GetMachineID() = %q, %v", machineID, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateMachineTreatsFingerprintTakenAsIdempotent(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
|
||||
if machineID, err := client.ActivateMachine("license-id", "fingerprint"); err != nil || machineID != "" {
|
||||
t.Fatalf("ActivateMachine() error = %v, want idempotent success", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateMachineReturnsCreatedMachineID(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
machineID, err := client.ActivateMachine("license-id", "fingerprint")
|
||||
if err != nil || machineID != "machine-id" {
|
||||
t.Fatalf("ActivateMachine() = %q, %v", machineID, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateMachineRejectsMachineLimitWithoutFingerprintTaken(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
|
||||
_, err := client.ActivateMachine("license-id", "fingerprint")
|
||||
if err == nil || !strings.Contains(err.Error(), "MACHINE_LIMIT_EXCEEDED") {
|
||||
t.Fatalf("ActivateMachine() error = %v, want machine limit failure", err)
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
@@ -31,7 +32,6 @@ type IngestionService struct {
|
||||
nodeBuffer []*model.NodeMetric
|
||||
nodeBufferMu sync.Mutex
|
||||
flushInterval time.Duration
|
||||
retentionDays int
|
||||
}
|
||||
|
||||
func NewIngestionService(repo *repo.Repository) *IngestionService {
|
||||
@@ -39,7 +39,6 @@ func NewIngestionService(repo *repo.Repository) *IngestionService {
|
||||
repo: repo,
|
||||
nodeBuffer: make([]*model.NodeMetric, 0, 500),
|
||||
flushInterval: 30 * time.Second,
|
||||
retentionDays: 7,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,7 +110,22 @@ func (s *IngestionService) flushNodeMetrics() {
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetrics() {
|
||||
cutoff := time.Now().Add(-time.Duration(s.retentionDays) * 24 * time.Hour).UnixMilli()
|
||||
s.pruneMetricsAt(time.Now())
|
||||
}
|
||||
|
||||
func (s *IngestionService) retentionDaysFromConfig() int {
|
||||
if s == nil || s.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := s.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetricsAt(now time.Time) {
|
||||
cutoff := now.Add(-time.Duration(s.retentionDaysFromConfig()) * 24 * time.Hour).UnixMilli()
|
||||
if s.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
@@ -215,7 +216,6 @@ func TestPruneMetrics(t *testing.T) {
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.retentionDays = 1
|
||||
|
||||
info := SystemInfo{CPUUsage: 50.0, MemoryUsage: 60.0, DiskUsage: 30.0}
|
||||
|
||||
@@ -233,6 +233,39 @@ func TestPruneMetrics(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPruneMetricsUsesConfiguredRetentionDays(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.UpsertConfig("monitor_retention_days", "2", now); err != nil {
|
||||
t.Fatalf("upsert retention config: %v", err)
|
||||
}
|
||||
|
||||
oldMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(3*24*time.Hour/time.Millisecond), CPUUsage: 10}
|
||||
newMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(1*24*time.Hour/time.Millisecond), CPUUsage: 20}
|
||||
if err := r.InsertNodeMetric(oldMetric); err != nil {
|
||||
t.Fatalf("insert old metric: %v", err)
|
||||
}
|
||||
if err := r.InsertNodeMetric(newMetric); err != nil {
|
||||
t.Fatalf("insert new metric: %v", err)
|
||||
}
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.pruneMetricsAt(time.UnixMilli(now))
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, now-int64(4*24*time.Hour/time.Millisecond), now+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get node metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 || metrics[0].CPUUsage != 20 {
|
||||
t.Fatalf("expected only newer metric to remain, got %#v", metrics)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultipleNodes(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
package monitoring
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
ConfigMonitorRetentionDays = "monitor_retention_days"
|
||||
DefaultMonitorRetentionDays = 7
|
||||
MinMonitorRetentionDays = 1
|
||||
MaxMonitorRetentionDays = 3650
|
||||
)
|
||||
|
||||
func MonitoringRetentionDaysFromConfigMap(cfg map[string]string) int {
|
||||
if cfg == nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
days, err := parseMonitoringRetentionDays(cfg[ConfigMonitorRetentionDays])
|
||||
if err != nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
return days
|
||||
}
|
||||
|
||||
func NormalizeMonitoringRetentionDays(value string) (string, error) {
|
||||
days, err := parseMonitoringRetentionDays(value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strconv.Itoa(days), nil
|
||||
}
|
||||
|
||||
func parseMonitoringRetentionDays(value string) (int, error) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return 0, fmt.Errorf("监控数据保留天数不能为空")
|
||||
}
|
||||
days, err := strconv.Atoi(trimmed)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须是整数")
|
||||
}
|
||||
if days < MinMonitorRetentionDays || days > MaxMonitorRetentionDays {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须在 %d 到 %d 之间", MinMonitorRetentionDays, MaxMonitorRetentionDays)
|
||||
}
|
||||
return days, nil
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package monitoring
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestMonitoringRetentionDaysFromConfigMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg map[string]string
|
||||
want int
|
||||
}{
|
||||
{"missing uses default", nil, 7},
|
||||
{"valid custom", map[string]string{ConfigMonitorRetentionDays: "3"}, 3},
|
||||
{"trimmed custom", map[string]string{ConfigMonitorRetentionDays: " 30 "}, 30},
|
||||
{"invalid uses default", map[string]string{ConfigMonitorRetentionDays: "abc"}, 7},
|
||||
{"too small uses default", map[string]string{ConfigMonitorRetentionDays: "0"}, 7},
|
||||
{"too large uses default", map[string]string{ConfigMonitorRetentionDays: "3651"}, 7},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := MonitoringRetentionDaysFromConfigMap(tc.cfg); got != tc.want {
|
||||
t.Fatalf("expected %d, got %d", tc.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeMonitoringRetentionDays(t *testing.T) {
|
||||
for _, value := range []string{"1", "7", "3650", " 30 "} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err != nil || got == "" {
|
||||
t.Fatalf("expected %q valid, got value=%q err=%v", value, got, err)
|
||||
}
|
||||
}
|
||||
|
||||
for _, value := range []string{"", "0", "-1", "3651", "abc", "1.5"} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err == nil {
|
||||
t.Fatalf("expected %q invalid, got value=%q", value, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package monitoring
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
ConfigTunnelQualityProbeIntervalSec = "monitor_tunnel_quality_interval_sec"
|
||||
DefaultTunnelQualityProbeIntervalSec = 1
|
||||
MinTunnelQualityProbeIntervalSec = 1
|
||||
MaxTunnelQualityProbeIntervalSec = 3600
|
||||
)
|
||||
|
||||
func TunnelQualityProbeIntervalSecondsFromConfigMap(cfg map[string]string) int {
|
||||
if cfg == nil {
|
||||
return DefaultTunnelQualityProbeIntervalSec
|
||||
}
|
||||
seconds, err := parseTunnelQualityProbeIntervalSeconds(cfg[ConfigTunnelQualityProbeIntervalSec])
|
||||
if err != nil {
|
||||
return DefaultTunnelQualityProbeIntervalSec
|
||||
}
|
||||
return seconds
|
||||
}
|
||||
|
||||
func NormalizeTunnelQualityProbeIntervalSeconds(value string) (string, error) {
|
||||
seconds, err := parseTunnelQualityProbeIntervalSeconds(value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strconv.Itoa(seconds), nil
|
||||
}
|
||||
|
||||
func parseTunnelQualityProbeIntervalSeconds(value string) (int, error) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return 0, fmt.Errorf("隧道质量探测间隔不能为空")
|
||||
}
|
||||
seconds, err := strconv.Atoi(trimmed)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("隧道质量探测间隔必须是整数")
|
||||
}
|
||||
if seconds < MinTunnelQualityProbeIntervalSec || seconds > MaxTunnelQualityProbeIntervalSec {
|
||||
return 0, fmt.Errorf(
|
||||
"隧道质量探测间隔必须在 %d 到 %d 秒之间",
|
||||
MinTunnelQualityProbeIntervalSec,
|
||||
MaxTunnelQualityProbeIntervalSec,
|
||||
)
|
||||
}
|
||||
return seconds, nil
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package monitoring
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestTunnelQualityProbeIntervalSecondsFromConfigMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg map[string]string
|
||||
want int
|
||||
}{
|
||||
{name: "missing config", cfg: nil, want: DefaultTunnelQualityProbeIntervalSec},
|
||||
{name: "configured", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "15"}, want: 15},
|
||||
{name: "invalid", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "0"}, want: DefaultTunnelQualityProbeIntervalSec},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := TunnelQualityProbeIntervalSecondsFromConfigMap(tt.cfg); got != tt.want {
|
||||
t.Fatalf("interval = %d, want %d", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelQualityProbeIntervalSeconds(t *testing.T) {
|
||||
for _, value := range []string{"1", "15", "3600"} {
|
||||
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err != nil || got != value {
|
||||
t.Fatalf("normalize %q = %q, %v", value, got, err)
|
||||
}
|
||||
}
|
||||
|
||||
for _, value := range []string{"", "0", "3601", "1.5", "abc"} {
|
||||
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err == nil {
|
||||
t.Fatalf("normalize %q unexpectedly succeeded with %q", value, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,148 @@
|
||||
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(" meta l4proto %s ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
|
||||
protocol,
|
||||
rule.InPort,
|
||||
family,
|
||||
targetHost,
|
||||
protocol,
|
||||
rule.TargetPort,
|
||||
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
|
||||
))
|
||||
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
|
||||
protocol,
|
||||
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"`,
|
||||
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`meta l4proto tcp ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
|
||||
`meta l4proto udp ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
|
||||
`meta l4proto 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"`,
|
||||
`meta l4proto tcp ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
|
||||
`meta l4proto 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{
|
||||
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`meta l4proto 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,241 @@
|
||||
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 {
|
||||
nft := nftBinary(cfg)
|
||||
tableName := fmt.Sprintf("flvx_capability_%d", time.Now().UnixNano())
|
||||
return r.run(ctx, cfg, buildCapabilityCheckCommand(nft, tableName))
|
||||
}
|
||||
|
||||
func buildCapabilityCheckCommand(nft, tableName string) string {
|
||||
script := RenderTable(NodePlan{
|
||||
Rules: []Rule{{
|
||||
ForwardID: 1,
|
||||
InPort: 12345,
|
||||
TargetHost: "192.0.2.1",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}},
|
||||
})
|
||||
script = strings.Replace(script, "table inet flvx {", "table inet "+tableName+" {", 1)
|
||||
return "set -eu\n" +
|
||||
"command -v nft >/dev/null 2>&1\n" +
|
||||
nft + " --version >/dev/null 2>&1\n" +
|
||||
"tmp=$(mktemp /tmp/flvx-nft-capability-XXXXXX.nft)\n" +
|
||||
"trap 'rm -f \"$tmp\"' EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
"if ! " + nft + " -c -f \"$tmp\"; then\n" +
|
||||
" echo 'nftables cannot validate the generated FLVX rules' >&2\n" +
|
||||
" exit 1\n" +
|
||||
"fi"
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
command := buildApplyCommand(nftBinary(cfg), script)
|
||||
return r.run(ctx, cfg, command)
|
||||
}
|
||||
|
||||
func buildApplyCommand(nft, script string) string {
|
||||
return "set -eu\n" +
|
||||
"tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft)\n" +
|
||||
"batch=$(mktemp /tmp/flvx-nft-batch-XXXXXX.nft) || { rm -f \"$tmp\"; exit 1; }\n" +
|
||||
"cleanup() {\n" +
|
||||
" rm -f \"$tmp\" \"$batch\"\n" +
|
||||
"}\n" +
|
||||
"trap cleanup EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
|
||||
" { printf '%s\\n' 'delete table inet flvx'; cat \"$tmp\"; } > \"$batch\"\n" +
|
||||
"else\n" +
|
||||
" cp \"$tmp\" \"$batch\"\n" +
|
||||
"fi\n" +
|
||||
"if ! " + nft + " -c -f \"$batch\"; then\n" +
|
||||
" echo 'nftables rule validation failed; active rules were preserved' >&2\n" +
|
||||
" exit 1\n" +
|
||||
"fi\n" +
|
||||
nft + " -f \"$batch\""
|
||||
}
|
||||
|
||||
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,145 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestBuildApplyCommandStopsAfterValidationFailure(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
logPath := filepath.Join(dir, "calls.log")
|
||||
applyMarker := filepath.Join(dir, "applied")
|
||||
nftPath := filepath.Join(dir, "nft")
|
||||
fake := `#!/bin/sh
|
||||
echo "$*" >> "` + logPath + `"
|
||||
if [ "$1" = "list" ]; then
|
||||
exit 0
|
||||
fi
|
||||
if [ "$1" = "-c" ]; then
|
||||
exit 1
|
||||
fi
|
||||
touch "` + applyMarker + `"
|
||||
`
|
||||
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
|
||||
t.Fatalf("write fake nft: %v", err)
|
||||
}
|
||||
|
||||
command := buildApplyCommand(nftPath, "table inet flvx { }")
|
||||
result := exec.Command("sh", "-c", command)
|
||||
if err := result.Run(); err == nil {
|
||||
t.Fatal("expected validation failure")
|
||||
}
|
||||
if _, err := os.Stat(applyMarker); !os.IsNotExist(err) {
|
||||
t.Fatalf("apply ran after validation failure, stat err=%v", err)
|
||||
}
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read fake nft calls: %v", err)
|
||||
}
|
||||
if strings.Count(string(calls), "-f ") != 1 {
|
||||
t.Fatalf("expected validation only, got calls:\n%s", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildCapabilityCheckCommandValidatesRenderedRulesWithoutApplying(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
logPath := filepath.Join(dir, "calls.log")
|
||||
nftPath := filepath.Join(dir, "nft")
|
||||
fake := `#!/bin/sh
|
||||
echo "$*" >> "` + logPath + `"
|
||||
if [ "$1" = "--version" ] || [ "$1" = "-c" ]; then
|
||||
exit 0
|
||||
fi
|
||||
exit 1
|
||||
`
|
||||
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
|
||||
t.Fatalf("write fake nft: %v", err)
|
||||
}
|
||||
|
||||
command := strings.Replace(buildCapabilityCheckCommand(nftPath, "flvx_capability_test"), "command -v nft", "command -v "+nftPath, 1)
|
||||
result := exec.Command("sh", "-c", command)
|
||||
if output, err := result.CombinedOutput(); err != nil {
|
||||
t.Fatalf("capability command failed: %v: %s", err, output)
|
||||
}
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read fake nft calls: %v", err)
|
||||
}
|
||||
if strings.Count(string(calls), "-c -f ") != 1 || strings.Contains(string(calls), "\n-f ") {
|
||||
t.Fatalf("expected one check-only invocation, got calls:\n%s", calls)
|
||||
}
|
||||
if !strings.Contains(command, "table inet flvx_capability_test") || !strings.Contains(command, "meta l4proto tcp ct original proto-dst") {
|
||||
t.Fatalf("capability check does not contain representative rendered rules:\n%s", command)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildApplyCommandUsesAtomicReplacementBatch(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
batchPath := filepath.Join(dir, "batch.nft")
|
||||
nftPath := filepath.Join(dir, "nft")
|
||||
fake := `#!/bin/sh
|
||||
if [ "$1" = "list" ]; then
|
||||
exit 0
|
||||
fi
|
||||
if [ "$1" = "-f" ]; then
|
||||
cp "$2" "` + batchPath + `"
|
||||
fi
|
||||
exit 0
|
||||
`
|
||||
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
|
||||
t.Fatalf("write fake nft: %v", err)
|
||||
}
|
||||
|
||||
script := "table inet flvx {\n chain forward { }\n}"
|
||||
result := exec.Command("sh", "-c", buildApplyCommand(nftPath, script))
|
||||
if output, err := result.CombinedOutput(); err != nil {
|
||||
t.Fatalf("apply command failed: %v: %s", err, output)
|
||||
}
|
||||
batch, err := os.ReadFile(batchPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read applied batch: %v", err)
|
||||
}
|
||||
want := "delete table inet flvx\n" + script + "\n"
|
||||
if string(batch) != want {
|
||||
t.Fatalf("atomic batch = %q, want %q", batch, want)
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user