mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 15:46:38 +08:00
Compare commits
68 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| cbec9a63da | |||
| 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 |
@@ -1,6 +1,6 @@
|
||||
# FLVX
|
||||
|
||||
> **联系我们**: [Telegram群组](https://t.me/flvxpanel)
|
||||
> **联系我们**: [Telegram群组](https://t.me/flvxchannel)
|
||||
|
||||
|
||||
## 特性
|
||||
|
||||
+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,并使用 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 才能启用此功能。
|
||||
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,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`
|
||||
|
||||
不需要数据库迁移或新增依赖。
|
||||
@@ -9,10 +9,14 @@ ARG TARGETOS
|
||||
ARG TARGETARCH
|
||||
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -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
|
||||
|
||||
+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,44 @@
|
||||
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
|
||||
}
|
||||
|
||||
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))
|
||||
}
|
||||
@@ -0,0 +1,277 @@
|
||||
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, "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)
|
||||
}
|
||||
|
||||
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 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 TestConfigUpdateAllowsLicenseKeyWrite(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)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
|
||||
cfg, err := r.GetConfigByName("license_key")
|
||||
if err != nil {
|
||||
t.Fatalf("get config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "license-secret" {
|
||||
t.Fatalf("expected license_key to be updated, got %#v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateSingleAllowsLicenseKeyWrite(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)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
|
||||
cfg, err := r.GetConfigByName("license_key")
|
||||
if err != nil {
|
||||
t.Fatalf("get config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "license-secret" {
|
||||
t.Fatalf("expected license_key to be updated, got %#v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
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"
|
||||
)
|
||||
@@ -58,6 +59,7 @@ type diagnosisWorkItem struct {
|
||||
type diagnosisExecOptions struct {
|
||||
commandTimeout time.Duration
|
||||
pingTimeoutMS int
|
||||
pingCount int
|
||||
timeoutMessage string
|
||||
}
|
||||
|
||||
@@ -240,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
|
||||
@@ -759,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
|
||||
@@ -774,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
|
||||
@@ -978,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)
|
||||
|
||||
@@ -987,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{}{
|
||||
@@ -1079,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{}{
|
||||
@@ -1093,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{}{
|
||||
@@ -1473,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 {
|
||||
@@ -1503,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",
|
||||
})
|
||||
@@ -1743,6 +1875,7 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
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"
|
||||
}
|
||||
@@ -1787,12 +1920,16 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
if runtimeLimiters.TrafficLimiter != "" {
|
||||
service["limiter"] = runtimeLimiters.TrafficLimiter
|
||||
}
|
||||
if forward.ProxyProtocol > 0 {
|
||||
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"] = forward.ProxyProtocol
|
||||
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = proxyProtocolSend
|
||||
}
|
||||
if protocol == "udp" {
|
||||
listenerMetadata := map[string]interface{}{
|
||||
@@ -1805,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) != "" {
|
||||
if service["metadata"] == nil {
|
||||
service["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
|
||||
serviceMetadata := ensureServiceMetadata(service)
|
||||
serviceMetadata["interface"] = node.InterfaceName
|
||||
}
|
||||
services = append(services, service)
|
||||
}
|
||||
@@ -1827,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 == "" {
|
||||
|
||||
@@ -9,14 +9,15 @@ import (
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigsAppliesProxyProtocolReceiveAndSendIndependently(t *testing.T) {
|
||||
forward := &forwardRecord{
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocol: 2,
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocolReceive: 1,
|
||||
ProxyProtocolSend: 2,
|
||||
}
|
||||
tunnel := &tunnelRecord{Type: 1}
|
||||
node := &nodeRecord{
|
||||
@@ -38,8 +39,8 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
|
||||
if serviceMetadata["interface"] != "eth0" {
|
||||
t.Fatalf("expected interface metadata eth0, got %v", serviceMetadata["interface"])
|
||||
}
|
||||
if _, ok := serviceMetadata["proxyProtocol"]; ok {
|
||||
t.Fatalf("proxyProtocol should not be listener metadata: %v", serviceMetadata)
|
||||
if serviceMetadata["proxyProtocol"] != 1 {
|
||||
t.Fatalf("expected service proxyProtocol 1 for receive mode, got %v", serviceMetadata["proxyProtocol"])
|
||||
}
|
||||
|
||||
handlerConfig, ok := service["handler"].(map[string]interface{})
|
||||
@@ -56,6 +57,46 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
|
||||
}
|
||||
}
|
||||
|
||||
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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -23,6 +23,7 @@ import (
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/metrics"
|
||||
"go-backend/internal/monitoring"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/repo"
|
||||
"go-backend/internal/ws"
|
||||
@@ -31,11 +32,12 @@ import (
|
||||
)
|
||||
|
||||
type Handler struct {
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
metrics *metrics.IngestionService
|
||||
healthCheck *health.Checker
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
metrics *metrics.IngestionService
|
||||
healthCheck *health.Checker
|
||||
nftablesManager nftablesRuntimeManager
|
||||
|
||||
captchaMu sync.Mutex
|
||||
captchaTokens map[string]int64
|
||||
@@ -46,6 +48,7 @@ type Handler struct {
|
||||
jobsWG sync.WaitGroup
|
||||
|
||||
upgradeMu sync.Mutex
|
||||
systemUpgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
nodeOnlineRedeployAt map[int64]time.Time
|
||||
nodeOnlineRedeployQueued map[int64]struct{}
|
||||
@@ -107,6 +110,7 @@ func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
metrics: metrics.NewIngestionService(repo),
|
||||
healthCheck: nil,
|
||||
nftablesManager: runtimenft.NewManager(nil),
|
||||
captchaTokens: make(map[string]int64),
|
||||
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||
nodeOnlineRedeployAt: make(map[int64]time.Time),
|
||||
@@ -135,6 +139,7 @@ func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
}
|
||||
h.metrics.RecordNodeMetric(nodeID, metricInfo)
|
||||
})
|
||||
h.wsServer.SetUserAuthStateLookup(h.GetUserAuthState)
|
||||
return h
|
||||
}
|
||||
|
||||
@@ -142,6 +147,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)
|
||||
@@ -151,11 +160,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)
|
||||
@@ -180,6 +193,9 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/node/batch-upgrade", h.nodeBatchUpgrade)
|
||||
mux.HandleFunc("/api/v1/node/rollback", h.nodeRollback)
|
||||
mux.HandleFunc("/api/v1/node/releases", h.listReleases)
|
||||
mux.HandleFunc("/api/v1/node/nftables/test", h.nodeNftablesTest)
|
||||
mux.HandleFunc("/api/v1/node/nftables/reconcile", h.nodeNftablesReconcile)
|
||||
mux.HandleFunc("/api/v1/node/nftables/clear", h.nodeNftablesClear)
|
||||
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
|
||||
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
|
||||
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||
@@ -205,6 +221,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)
|
||||
@@ -330,7 +347,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
|
||||
}
|
||||
@@ -338,8 +356,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
|
||||
@@ -370,13 +400,17 @@ 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 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
|
||||
@@ -553,10 +587,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 := ""
|
||||
@@ -944,6 +995,10 @@ 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 protectedKeys[key] && isCommercial != "true" {
|
||||
response.WriteJSON(w, response.ErrDefault("需要商业版授权"))
|
||||
@@ -960,6 +1015,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())
|
||||
@@ -981,6 +1037,10 @@ 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
|
||||
}
|
||||
|
||||
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
|
||||
if (name == "app_name" || name == "app_logo" || name == "app_favicon" || name == "hide_footer_brand") && isCommercial != "true" {
|
||||
@@ -1003,10 +1063,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":
|
||||
@@ -1043,11 +1112,23 @@ func normalizeAndValidateConfigValue(key, value string) (string, error) {
|
||||
}
|
||||
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
|
||||
@@ -1251,7 +1332,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
|
||||
}
|
||||
@@ -1266,7 +1348,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,14 @@ 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 +23,7 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
h.jobsCancel = cancel
|
||||
h.jobsStarted = true
|
||||
h.jobsWG.Add(7)
|
||||
h.jobsWG.Add(8)
|
||||
h.jobsMu.Unlock()
|
||||
|
||||
go h.runHourlyStatsLoop(ctx)
|
||||
@@ -30,6 +33,7 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
go h.runHealthChecks(ctx)
|
||||
go h.runTunnelQualityProber(ctx)
|
||||
go h.runValidateLicenseJob(ctx)
|
||||
go h.runNftablesTrafficCollectLoop(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) runValidateLicenseJob(ctx context.Context) {
|
||||
@@ -64,7 +68,7 @@ func (h *Handler) validateLicenseJob() {
|
||||
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
|
||||
|
||||
if err != nil {
|
||||
// Network error or timeout. Grace period by not revoking immediately here.
|
||||
return
|
||||
@@ -128,6 +132,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()
|
||||
|
||||
|
||||
@@ -71,7 +71,12 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
now := time.Now().UnixMilli()
|
||||
maxConn := asInt(req["maxConn"], 0)
|
||||
|
||||
userID, err := h.repo.CreateUser(username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
|
||||
hashedPassword, err := security.HashPassword(pwd)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -176,7 +181,12 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
} else {
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
hashedPassword, err := security.HashPassword(pwd)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, hashedPassword, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -331,6 +341,11 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
inx := h.repo.NextIndex("node")
|
||||
forwardMode := defaultNodeForwardMode(asString(req["forwardMode"]))
|
||||
status := 0
|
||||
if forwardMode == "nftables" {
|
||||
status = 1
|
||||
}
|
||||
if err := h.repo.CreateNode(
|
||||
name,
|
||||
randomToken(16),
|
||||
@@ -347,7 +362,7 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
asInt(req["tls"], 0),
|
||||
asInt(req["socks"], 0),
|
||||
now,
|
||||
0,
|
||||
status,
|
||||
defaultString(asString(req["tcpListenAddr"]), "[::]"),
|
||||
defaultString(asString(req["udpListenAddr"]), "[::]"),
|
||||
inx,
|
||||
@@ -356,10 +371,20 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(asString(req["remoteToken"])),
|
||||
nullableText(asString(req["remoteConfig"])),
|
||||
nullableText(asString(req["extraIPs"])),
|
||||
forwardMode,
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
nodeID, err := h.findCreatedNodeID(name, serverIP)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.persistNodeSSHConfig(nodeID, req, forwardMode, now, false); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -392,6 +417,12 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
newHTTP := asInt(req["http"], currentHTTP)
|
||||
newTLS := asInt(req["tls"], currentTLS)
|
||||
newSocks := asInt(req["socks"], currentSocks)
|
||||
forwardMode := defaultNodeForwardMode(strings.TrimSpace(asString(req["forwardMode"])))
|
||||
currentForwardMode, err := h.repo.GetNodeForwardMode(id)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
serverIP := asString(req["serverIp"])
|
||||
if serverIP != "" {
|
||||
if err := IsValidNodeAddress(serverIP); err != nil {
|
||||
@@ -399,7 +430,8 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
}
|
||||
if currentStatus == 1 && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
|
||||
usesNftablesRuntime := forwardMode == "nftables" || defaultNodeForwardMode(currentForwardMode) == "nftables"
|
||||
if currentStatus == 1 && !usesNftablesRuntime && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
|
||||
if err := h.applyNodeProtocolChange(id, newHTTP, newTLS, newSocks); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
@@ -418,6 +450,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(strings.TrimSpace(asString(req["remark"]))),
|
||||
nullableUnixMilli(asInt64(req["expiryTime"], 0)),
|
||||
nullableText(normalizeNodeRenewalCycle(asString(req["renewalCycle"]))),
|
||||
forwardMode,
|
||||
newHTTP,
|
||||
newTLS,
|
||||
newSocks,
|
||||
@@ -428,9 +461,142 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.persistNodeSSHConfig(id, req, forwardMode, now, true); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if forwardMode == "nftables" && currentStatus != 1 {
|
||||
if err := h.repo.UpdateNodeStatus(id, 1); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) findCreatedNodeID(name, serverIP string) (int64, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return 0, errors.New("handler not initialized")
|
||||
}
|
||||
nodes, err := h.repo.ListNodes()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
for i := len(nodes) - 1; i >= 0; i-- {
|
||||
item := nodes[i]
|
||||
if asString(item["name"]) != name {
|
||||
continue
|
||||
}
|
||||
if asString(item["serverIp"]) != serverIP {
|
||||
continue
|
||||
}
|
||||
if nodeID := asInt64(item["id"], 0); nodeID > 0 {
|
||||
return nodeID, nil
|
||||
}
|
||||
}
|
||||
return 0, errors.New("节点创建成功,但未能查询到节点记录")
|
||||
}
|
||||
|
||||
func (h *Handler) persistNodeSSHConfig(nodeID int64, req map[string]interface{}, forwardMode string, now int64, preserveSecrets bool) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
if nodeID <= 0 {
|
||||
return errors.New("节点ID不能为空")
|
||||
}
|
||||
if forwardMode != "nftables" {
|
||||
return h.repo.DeleteNodeSSHConfig(nodeID)
|
||||
}
|
||||
cfgMap := asMap(req["sshConfig"])
|
||||
host := strings.TrimSpace(asString(cfgMap["host"]))
|
||||
if host == "" {
|
||||
host = strings.TrimSpace(asString(req["serverIp"]))
|
||||
}
|
||||
port := asInt(cfgMap["port"], 22)
|
||||
username := strings.TrimSpace(asString(cfgMap["username"]))
|
||||
authType := strings.TrimSpace(asString(cfgMap["authType"]))
|
||||
password := asString(cfgMap["password"])
|
||||
privateKey := asString(cfgMap["privateKey"])
|
||||
passphrase := asString(cfgMap["passphrase"])
|
||||
sudoMode := strings.TrimSpace(asString(cfgMap["sudoMode"]))
|
||||
|
||||
if preserveSecrets {
|
||||
existing, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return err
|
||||
}
|
||||
if existing != nil {
|
||||
if host == "" {
|
||||
host = strings.TrimSpace(existing.Host)
|
||||
}
|
||||
if port <= 0 {
|
||||
port = existing.Port
|
||||
}
|
||||
if username == "" {
|
||||
username = strings.TrimSpace(existing.Username)
|
||||
}
|
||||
if authType == "" {
|
||||
authType = strings.TrimSpace(existing.AuthType)
|
||||
}
|
||||
if strings.TrimSpace(password) == "" && existing.Password.Valid {
|
||||
password = existing.Password.String
|
||||
}
|
||||
if strings.TrimSpace(privateKey) == "" && existing.PrivateKey.Valid {
|
||||
privateKey = existing.PrivateKey.String
|
||||
}
|
||||
if strings.TrimSpace(passphrase) == "" && existing.Passphrase.Valid {
|
||||
passphrase = existing.Passphrase.String
|
||||
}
|
||||
if sudoMode == "" {
|
||||
sudoMode = strings.TrimSpace(existing.SudoMode)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if host == "" || username == "" {
|
||||
return errors.New("nftables 节点 SSH 配置不完整")
|
||||
}
|
||||
if port <= 0 || port > 65535 {
|
||||
return errors.New("nftables 节点 SSH 端口无效")
|
||||
}
|
||||
|
||||
authType = strings.ToLower(authType)
|
||||
switch authType {
|
||||
case "password":
|
||||
if strings.TrimSpace(password) == "" {
|
||||
return errors.New("nftables 节点 SSH 密码不能为空")
|
||||
}
|
||||
privateKey = ""
|
||||
case "private_key", "":
|
||||
authType = "private_key"
|
||||
if strings.TrimSpace(privateKey) == "" {
|
||||
return errors.New("nftables 节点 SSH 私钥不能为空")
|
||||
}
|
||||
password = ""
|
||||
default:
|
||||
return errors.New("nftables 节点 SSH 认证方式无效")
|
||||
}
|
||||
|
||||
switch strings.ToLower(sudoMode) {
|
||||
case "", "none":
|
||||
sudoMode = "none"
|
||||
case "sudo", "sudo_su":
|
||||
default:
|
||||
return errors.New("nftables 节点 sudo 模式无效")
|
||||
}
|
||||
|
||||
return h.repo.UpsertNodeSSHConfig(nodeID, repo.NftSSHConfigInput{
|
||||
Host: host,
|
||||
Port: port,
|
||||
Username: username,
|
||||
AuthType: authType,
|
||||
Password: password,
|
||||
PrivateKey: privateKey,
|
||||
Passphrase: passphrase,
|
||||
SudoMode: sudoMode,
|
||||
}, now)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -638,6 +804,16 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if strings.TrimSpace(inIP) == "" {
|
||||
inIP = buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
|
||||
}
|
||||
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
|
||||
for _, inNode := range runtimeState.InNodes {
|
||||
if inNode.NodeID > 0 {
|
||||
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateNftablesTunnelStateTx(tx, entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if len(runtimeState.InNodes) > 0 {
|
||||
firstNodeID := runtimeState.InNodes[0].NodeID
|
||||
@@ -924,6 +1100,16 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
runtimeState.TunnelID = id
|
||||
runtimeState.IPPreference = ipPreference
|
||||
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
|
||||
for _, inNode := range runtimeState.InNodes {
|
||||
if inNode.NodeID > 0 {
|
||||
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateNftablesTunnelState(entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
|
||||
|
||||
@@ -1690,6 +1876,24 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, tunnelErr)
|
||||
continue
|
||||
}
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(tunnelID); modeErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
fail++
|
||||
failures = appendBatchFailureReason(failures, tunnelID, tunnelName, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
if reconcileErr := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); reconcileErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, reconcileErr)
|
||||
continue
|
||||
}
|
||||
success++
|
||||
continue
|
||||
}
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, err)
|
||||
@@ -1899,6 +2103,17 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
port = 10000
|
||||
}
|
||||
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
|
||||
isNftTunnel, _, err := h.tunnelUsesNftables(tunnelID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if isNftTunnel {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
inIp := strings.TrimSpace(asString(req["inIp"]))
|
||||
if inIp != "" && len(entryNodes) > 1 {
|
||||
response.WriteJSON(w, response.ErrDefault("多入口隧道的转发不支持自定义监听IP"))
|
||||
@@ -1934,8 +2149,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
ipMaxConn = 0
|
||||
}
|
||||
proxyProtocol := asInt(req["proxyProtocol"], 0)
|
||||
proxyProtocolReceive := asInt(req["proxyProtocolReceive"], 0)
|
||||
proxyProtocolSend := asInt(req["proxyProtocolSend"], proxyProtocol)
|
||||
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, ipMaxConn, nullableInt(ipSpeedID), proxyProtocol)
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, ipMaxConn, nullableInt(ipSpeedID), proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -2006,6 +2223,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if remoteAddr == "" {
|
||||
remoteAddr = forward.RemoteAddr
|
||||
}
|
||||
isNftTunnel, entryNodes, err := h.tunnelUsesNftables(tunnelID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if isNftTunnel {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
if actorRole != 0 && !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
@@ -2113,8 +2341,10 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
ipMaxConn = 0
|
||||
}
|
||||
proxyProtocol := asInt(req["proxyProtocol"], forward.ProxyProtocol)
|
||||
proxyProtocolReceive := asInt(req["proxyProtocolReceive"], forward.ProxyProtocolReceive)
|
||||
proxyProtocolSend := asInt(req["proxyProtocolSend"], forward.ProxyProtocolSend)
|
||||
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, ipMaxConn, newIPSpeedID, proxyProtocol); err != nil {
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, ipMaxConn, newIPSpeedID, proxyProtocol, proxyProtocolReceive, proxyProtocolSend); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -2192,7 +2422,17 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
var nftNodeID int64
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, modeErr.Error()))
|
||||
return
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("nftables 转发缺少入口节点"))
|
||||
return
|
||||
}
|
||||
nftNodeID = entryNodeIDs[0]
|
||||
} else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -2200,6 +2440,12 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if nftNodeID > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -2208,7 +2454,7 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
_, _, _, err := h.resolveForwardAccess(r, id)
|
||||
forward, _, _, err := h.resolveForwardAccess(r, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, errForwardNotFound) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发不存在"))
|
||||
@@ -2224,6 +2470,13 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
_ = h.repo.DeleteNftRuleBindingsByForward(id)
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr == nil && nftMode && len(entryNodeIDs) > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -2249,6 +2502,26 @@ func (h *Handler) forwardPause(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
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())
|
||||
}
|
||||
|
||||
func (h *Handler) forwardResume(w http.ResponseWriter, r *http.Request) {
|
||||
id := idFromBody(r, w)
|
||||
if id <= 0 {
|
||||
@@ -2341,7 +2614,19 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
var nftNodeID int64
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
f++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
nftNodeID = entryNodeIDs[0]
|
||||
} else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
@@ -2349,9 +2634,16 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if err := h.deleteForwardByID(id); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
continue
|
||||
}
|
||||
if nftNodeID > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
s++
|
||||
}
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||
}
|
||||
@@ -2452,6 +2744,24 @@ func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
f++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
@@ -4442,7 +4752,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
|
||||
h.repo.RollbackForwardFields(
|
||||
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
||||
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
||||
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol,
|
||||
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol, oldForward.ProxyProtocolReceive, oldForward.ProxyProtocolSend,
|
||||
time.Now().UnixMilli(),
|
||||
)
|
||||
|
||||
@@ -4695,6 +5005,13 @@ func asMapSlice(v interface{}) []map[string]interface{} {
|
||||
return out
|
||||
}
|
||||
|
||||
func asMap(v interface{}) map[string]interface{} {
|
||||
if m, ok := v.(map[string]interface{}); ok && m != nil {
|
||||
return m
|
||||
}
|
||||
return map[string]interface{}{}
|
||||
}
|
||||
|
||||
func asString(v interface{}) string {
|
||||
switch t := v.(type) {
|
||||
case nil:
|
||||
@@ -4833,6 +5150,15 @@ func normalizeNodeRenewalCycle(v string) string {
|
||||
}
|
||||
}
|
||||
|
||||
func defaultNodeForwardMode(mode string) string {
|
||||
switch strings.TrimSpace(strings.ToLower(mode)) {
|
||||
case "nftables":
|
||||
return "nftables"
|
||||
default:
|
||||
return "agent"
|
||||
}
|
||||
}
|
||||
|
||||
func nullableInt(v *int64) interface{} {
|
||||
if v == nil {
|
||||
return nil
|
||||
|
||||
@@ -0,0 +1,371 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type nftablesRuntimeManager interface {
|
||||
Test(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
Reconcile(ctx context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error)
|
||||
Clear(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
CollectCounters(ctx context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error)
|
||||
}
|
||||
|
||||
func isNftablesForwardMode(mode string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(mode), runtimenft.ModeNftables)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeUsesNftables(nodeID int64) (bool, error) {
|
||||
return h.nodeUsesNftablesTx(nil, nodeID)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeUsesNftablesTx(tx *gorm.DB, nodeID int64) (bool, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return false, errors.New("handler not initialized")
|
||||
}
|
||||
var (
|
||||
mode string
|
||||
err error
|
||||
)
|
||||
if tx != nil {
|
||||
mode, err = h.repo.GetNodeForwardModeTx(tx, nodeID)
|
||||
} else {
|
||||
mode, err = h.repo.GetNodeForwardMode(nodeID)
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return isNftablesForwardMode(mode), nil
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelUsesNftables(tunnelID int64) (bool, []int64, error) {
|
||||
entryNodeIDs, err := h.tunnelEntryNodeIDs(tunnelID)
|
||||
if err != nil {
|
||||
return false, nil, err
|
||||
}
|
||||
for _, nodeID := range entryNodeIDs {
|
||||
ok, modeErr := h.nodeUsesNftables(nodeID)
|
||||
if modeErr != nil {
|
||||
return false, nil, modeErr
|
||||
}
|
||||
if ok {
|
||||
return true, entryNodeIDs, nil
|
||||
}
|
||||
}
|
||||
return false, entryNodeIDs, nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesForwardRequest(tunnel *tunnelRecord, remoteAddr string, entryNodeIDs []int64) error {
|
||||
if tunnel == nil {
|
||||
return errors.New("隧道不存在")
|
||||
}
|
||||
if tunnel.Type != 1 {
|
||||
return errors.New("nftables 节点仅支持直连隧道")
|
||||
}
|
||||
if len(entryNodeIDs) != 1 {
|
||||
return errors.New("nftables 节点仅支持单入口隧道")
|
||||
}
|
||||
target, err := runtimenft.ParseSingleTarget(remoteAddr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if net.ParseIP(strings.Trim(strings.TrimSpace(target.Host), "[]")) == nil {
|
||||
return errors.New("nftables 节点仅支持 IP 目标地址")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func sshConfigFromModel(cfg *model.NodeSSHConfig) (runtimenft.SSHConfig, error) {
|
||||
if cfg == nil {
|
||||
return runtimenft.SSHConfig{}, errors.New("节点缺少 SSH 配置")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Host) == "" || strings.TrimSpace(cfg.Username) == "" {
|
||||
return runtimenft.SSHConfig{}, errors.New("节点 SSH 配置不完整")
|
||||
}
|
||||
return runtimenft.SSHConfig{
|
||||
Host: strings.TrimSpace(cfg.Host),
|
||||
Port: cfg.Port,
|
||||
Username: strings.TrimSpace(cfg.Username),
|
||||
AuthType: strings.TrimSpace(cfg.AuthType),
|
||||
Password: cfg.Password.String,
|
||||
PrivateKey: cfg.PrivateKey.String,
|
||||
Passphrase: cfg.Passphrase.String,
|
||||
SudoMode: strings.TrimSpace(cfg.SudoMode),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesTunnelState(entryNodeIDs []int64) error {
|
||||
return h.validateNftablesTunnelStateTx(nil, entryNodeIDs)
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesTunnelStateTx(tx *gorm.DB, entryNodeIDs []int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
for _, nodeID := range entryNodeIDs {
|
||||
isNft, err := h.nodeUsesNftablesTx(tx, nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !isNft {
|
||||
continue
|
||||
}
|
||||
var cfg *model.NodeSSHConfig
|
||||
if tx != nil {
|
||||
cfg, err = h.repo.GetNodeSSHConfigTx(tx, nodeID)
|
||||
} else {
|
||||
cfg, err = h.repo.GetNodeSSHConfig(nodeID)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return errors.New("nftables 节点缺少 SSH 配置")
|
||||
}
|
||||
return err
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
|
||||
return fmt.Errorf("nftables 节点能力校验失败: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) buildNftablesNodePlan(nodeID int64) (runtimenft.NodePlan, *model.NodeSSHConfig, error) {
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
forwards, err := h.repo.ListActiveForwardsByNode(nodeID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
plan := runtimenft.NodePlan{NodeID: nodeID, Rules: make([]runtimenft.Rule, 0, len(forwards))}
|
||||
for i := range forwards {
|
||||
forward := &forwards[i]
|
||||
tunnel, err := h.getTunnelRecord(forward.TunnelID)
|
||||
if err != nil || tunnel == nil || tunnel.Status != 1 {
|
||||
continue
|
||||
}
|
||||
entryNodeIDs, err := h.tunnelEntryNodeIDs(forward.TunnelID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
if len(entryNodeIDs) != 1 || entryNodeIDs[0] != nodeID {
|
||||
continue
|
||||
}
|
||||
if err := h.validateNftablesForwardRequest(tunnel, forward.RemoteAddr, entryNodeIDs); err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
ports, err := h.listForwardPorts(forward.ID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
for _, fp := range ports {
|
||||
if fp.NodeID != nodeID {
|
||||
continue
|
||||
}
|
||||
target, err := runtimenft.ParseSingleTarget(forward.RemoteAddr)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
plan.Rules = append(plan.Rules, runtimenft.Rule{
|
||||
ForwardID: forward.ID,
|
||||
InPort: fp.Port,
|
||||
BindIP: strings.TrimSpace(fp.InIP),
|
||||
TargetHost: target.Host,
|
||||
TargetPort: target.Port,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
})
|
||||
}
|
||||
}
|
||||
return plan, cfg, nil
|
||||
}
|
||||
|
||||
func (h *Handler) syncNftablesNode(nodeID int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
plan, cfgModel, err := h.buildNftablesNodePlan(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfgModel)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
result, err := h.nftablesManager.Reconcile(context.Background(), sshCfg, plan)
|
||||
now := time.Now().UnixMilli()
|
||||
if err != nil {
|
||||
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
for _, binding := range bindings {
|
||||
_ = h.repo.MarkNftRuleBindingError(binding.ForwardID, nodeID, err.Error(), now)
|
||||
}
|
||||
return err
|
||||
}
|
||||
activeForwardIDs := make(map[int64]struct{}, len(plan.Rules))
|
||||
for _, rule := range plan.Rules {
|
||||
activeForwardIDs[rule.ForwardID] = struct{}{}
|
||||
hash := result.Hashes[rule.ForwardID]
|
||||
_ = h.repo.UpsertNftRuleBinding(modelToRuleBindingInput(nodeID, rule, hash), now)
|
||||
}
|
||||
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
for _, binding := range bindings {
|
||||
if _, ok := activeForwardIDs[binding.ForwardID]; ok {
|
||||
continue
|
||||
}
|
||||
_ = h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func modelToRuleBindingInput(nodeID int64, rule runtimenft.Rule, hash string) repo.NftRuleBindingInput {
|
||||
return repo.NftRuleBindingInput{
|
||||
ForwardID: rule.ForwardID,
|
||||
NodeID: nodeID,
|
||||
InPort: rule.InPort,
|
||||
Protocols: strings.Join(rule.Protocols, ","),
|
||||
TargetAddr: fmt.Sprintf("%s:%d", rule.TargetHost, rule.TargetPort),
|
||||
BindIP: rule.BindIP,
|
||||
RuleHash: hash,
|
||||
Status: runtimenft.StatusApplied,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) nftablesNodeIDFromRequest(r *http.Request, w http.ResponseWriter) (int64, bool) {
|
||||
nodeID := asInt64FromBodyKey(r, w, "nodeId")
|
||||
if nodeID <= 0 {
|
||||
return 0, false
|
||||
}
|
||||
return nodeID, true
|
||||
}
|
||||
|
||||
func (h *Handler) loadNftablesSSHConfig(nodeID int64) (runtimenft.SSHConfig, error) {
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return runtimenft.SSHConfig{}, errors.New("nftables 节点缺少 SSH 配置")
|
||||
}
|
||||
return runtimenft.SSHConfig{}, err
|
||||
}
|
||||
return sshConfigFromModel(cfg)
|
||||
}
|
||||
|
||||
func (h *Handler) clearNftablesNode(nodeID int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
if err := h.nftablesManager.Clear(context.Background(), sshCfg); err != nil {
|
||||
return err
|
||||
}
|
||||
bindings, listErr := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
if listErr != nil {
|
||||
return listErr
|
||||
}
|
||||
for _, binding := range bindings {
|
||||
if err := h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) reconcileNftablesNodeByRequest(nodeID int64) error {
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !usesNft {
|
||||
return errors.New("节点未启用 nftables 转发模式")
|
||||
}
|
||||
return h.syncNftablesNode(nodeID)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesTest(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if !usesNft {
|
||||
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
|
||||
return
|
||||
}
|
||||
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "nftables manager not initialized"))
|
||||
return
|
||||
}
|
||||
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesReconcile(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := h.reconcileNftablesNodeByRequest(nodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesClear(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if !usesNft {
|
||||
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
|
||||
return
|
||||
}
|
||||
if err := h.clearNftablesNode(nodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
@@ -0,0 +1,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"
|
||||
)
|
||||
|
||||
@@ -75,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,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))
|
||||
}
|
||||
}
|
||||
@@ -152,7 +152,11 @@ func evaluateBestExitOwner(owner chainNodeRecord, exits []chainNodeRecord, nodes
|
||||
ownerNode := nodes[owner.NodeID]
|
||||
for _, exit := range exits {
|
||||
exitNode := nodes[exit.NodeID]
|
||||
if exitNode == nil {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -358,9 +358,9 @@ func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
|
||||
{NodeID: 31, NodeName: "exit-b", Port: 30031},
|
||||
}
|
||||
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: "[::]"},
|
||||
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 {
|
||||
@@ -387,12 +387,30 @@ func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
|
||||
30: {ID: 30, Name: "exit-a", ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
|
||||
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
|
||||
@@ -419,8 +437,8 @@ func TestEvaluateBestExitOwnerMarksCandidateFailedWhenOwnerToExitFails(t *testin
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
|
||||
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: "[::]"},
|
||||
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
|
||||
@@ -436,8 +454,8 @@ func TestEvaluateBestExitOwnerMarksCandidateFailedWhenTargetResolutionFails(t *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", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Name: "exit-v6", ServerIP: "2001:db8::30", ServerIPv6: "2001:db8::30", TCPListenAddr: "[::]"},
|
||||
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)
|
||||
|
||||
@@ -209,6 +209,22 @@ func TestTunnelUpdateInvalidProbeTargetDoesNotCleanFederationBindings(t *testing
|
||||
}
|
||||
}
|
||||
|
||||
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"))
|
||||
|
||||
@@ -3,6 +3,8 @@ package handler
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -13,7 +15,6 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
tunnelQualityProbeInterval = 1 * time.Second
|
||||
tunnelQualityProbeTimeout = 8 * time.Second
|
||||
tunnelQualityPingTimeoutMs = 5000
|
||||
tunnelQualityPruneInterval = 10 * time.Minute
|
||||
@@ -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"`
|
||||
@@ -56,7 +77,7 @@ 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
|
||||
@@ -65,8 +86,8 @@ type tunnelQualityProber struct {
|
||||
// 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),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -86,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
|
||||
@@ -109,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
|
||||
@@ -246,15 +301,28 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
options := diagnosisExecOptions{
|
||||
commandTimeout: tunnelQualityProbeTimeout,
|
||||
pingTimeoutMS: tunnelQualityPingTimeoutMs,
|
||||
pingCount: 1,
|
||||
timeoutMessage: "探测超时",
|
||||
}
|
||||
p.probeBestExitOwners(tunnelID, inNodes, midNodesGrouped, outNodes, ipPreference, options, probeTarget)
|
||||
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 → public probe target only.
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if entryOnline {
|
||||
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -262,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]
|
||||
@@ -293,12 +379,12 @@ 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
|
||||
}
|
||||
|
||||
@@ -309,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.pingNode(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()
|
||||
}
|
||||
@@ -342,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.pingNode(outNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if exitOnline {
|
||||
lat, loss, err := roundPinger(exit.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -367,8 +447,8 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
snap.Success = probeOK
|
||||
default:
|
||||
// Unknown type: entry → public probe target.
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if entryOnline {
|
||||
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -376,13 +456,215 @@ 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 (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chainNodeRecord, chainHops [][]chainNodeRecord, outNodes []chainNodeRecord, ipPreference string, options diagnosisExecOptions, probeTarget tunnelProbeTarget) {
|
||||
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
|
||||
}
|
||||
@@ -404,9 +686,6 @@ func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chai
|
||||
nodeMap[exit.NodeID] = node
|
||||
}
|
||||
}
|
||||
// This best-exit decision cache is per decision round; the display-oriented
|
||||
// tunnel quality snapshot may still collect its own first-exit public probe.
|
||||
roundPinger := newBestExitRoundPinger(p.pingNode)
|
||||
for _, owner := range owners {
|
||||
if nodeMap[owner.NodeID] == nil {
|
||||
continue
|
||||
@@ -444,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
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"slices"
|
||||
"testing"
|
||||
@@ -26,6 +27,9 @@ func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
|
||||
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
|
||||
}
|
||||
@@ -46,6 +50,114 @@ func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
@@ -67,3 +179,69 @@ func TestTunnelQualityProberStoresProbeTargetWhenChainIncomplete(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
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 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
|
||||
|
||||
@@ -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,146 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func RenderTable(plan NodePlan) string {
|
||||
var b strings.Builder
|
||||
b.WriteString("table inet flvx {\n")
|
||||
b.WriteString(" chain prerouting {\n")
|
||||
b.WriteString(" type nat hook prerouting priority dstnat; policy accept;\n")
|
||||
for _, rule := range sortedRules(plan.Rules) {
|
||||
family := nftAddressFamily(rule.TargetHost)
|
||||
dnatFamily := ""
|
||||
if family != "" {
|
||||
dnatFamily = family + " "
|
||||
}
|
||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
||||
b.WriteString(fmt.Sprintf(" %s dport %d counter dnat %sto %s comment %q\n",
|
||||
protocol,
|
||||
rule.InPort,
|
||||
dnatFamily,
|
||||
formatDNATTarget(rule.TargetHost, rule.TargetPort),
|
||||
counterComment(rule.ForwardID, CounterDirectionDNAT, protocol),
|
||||
))
|
||||
}
|
||||
}
|
||||
b.WriteString(" }\n\n")
|
||||
b.WriteString(" chain postrouting {\n")
|
||||
b.WriteString(" type nat hook postrouting priority srcnat; policy accept;\n")
|
||||
if len(plan.Rules) > 0 {
|
||||
b.WriteString(" masquerade comment \"flvx masquerade\"\n")
|
||||
}
|
||||
b.WriteString(" }\n\n")
|
||||
b.WriteString(" chain forward {\n")
|
||||
b.WriteString(" type filter hook forward priority filter; policy accept;\n")
|
||||
for _, rule := range sortedRules(plan.Rules) {
|
||||
family := nftAddressFamily(rule.TargetHost)
|
||||
if family == "" {
|
||||
continue
|
||||
}
|
||||
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
|
||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
||||
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
|
||||
rule.InPort,
|
||||
family,
|
||||
targetHost,
|
||||
protocol,
|
||||
rule.TargetPort,
|
||||
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
|
||||
))
|
||||
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
|
||||
rule.InPort,
|
||||
family,
|
||||
targetHost,
|
||||
protocol,
|
||||
rule.TargetPort,
|
||||
counterComment(rule.ForwardID, CounterDirectionFromTarget, protocol),
|
||||
))
|
||||
}
|
||||
}
|
||||
b.WriteString(" }\n")
|
||||
b.WriteString("}\n")
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func counterComment(forwardID int64, direction, protocol string) string {
|
||||
return fmt.Sprintf("flvx forward:%d %s %s", forwardID, direction, protocol)
|
||||
}
|
||||
|
||||
func RuleHash(rule Rule) string {
|
||||
protocols := normalizedProtocols(rule.Protocols)
|
||||
sum := sha256.Sum256([]byte(fmt.Sprintf("%d|%d|%s|%d|%s",
|
||||
rule.ForwardID,
|
||||
rule.InPort,
|
||||
strings.TrimSpace(rule.TargetHost),
|
||||
rule.TargetPort,
|
||||
strings.Join(protocols, ","),
|
||||
)))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func PlanHashes(plan NodePlan) map[int64]string {
|
||||
hashes := make(map[int64]string, len(plan.Rules))
|
||||
for _, rule := range plan.Rules {
|
||||
hashes[rule.ForwardID] = RuleHash(rule)
|
||||
}
|
||||
return hashes
|
||||
}
|
||||
|
||||
func sortedRules(rules []Rule) []Rule {
|
||||
out := append([]Rule(nil), rules...)
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
if out[i].InPort == out[j].InPort {
|
||||
return out[i].ForwardID < out[j].ForwardID
|
||||
}
|
||||
return out[i].InPort < out[j].InPort
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
func normalizedProtocols(protocols []string) []string {
|
||||
seen := map[string]struct{}{}
|
||||
out := make([]string, 0, 2)
|
||||
for _, protocol := range protocols {
|
||||
p := strings.ToLower(strings.TrimSpace(protocol))
|
||||
if p != "tcp" && p != "udp" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[p]; ok {
|
||||
continue
|
||||
}
|
||||
seen[p] = struct{}{}
|
||||
out = append(out, p)
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return []string{"tcp", "udp"}
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func formatDNATTarget(host string, port int) string {
|
||||
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
|
||||
if ip := net.ParseIP(trimmed); ip != nil && ip.To4() == nil {
|
||||
return fmt.Sprintf("[%s]:%d", trimmed, port)
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", trimmed, port)
|
||||
}
|
||||
|
||||
func nftAddressFamily(host string) string {
|
||||
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
|
||||
ip := net.ParseIP(trimmed)
|
||||
if ip == nil {
|
||||
return ""
|
||||
}
|
||||
if ip.To4() == nil {
|
||||
return "ip6"
|
||||
}
|
||||
return "ip"
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRenderTableIncludesDNATAndMasquerade(t *testing.T) {
|
||||
script := RenderTable(NodePlan{
|
||||
NodeID: 10,
|
||||
Rules: []Rule{
|
||||
{
|
||||
ForwardID: 42,
|
||||
InPort: 24000,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
expectedParts := []string{
|
||||
"table inet flvx",
|
||||
"type nat hook prerouting priority dstnat; policy accept;",
|
||||
"type nat hook postrouting priority srcnat; policy accept;",
|
||||
"tcp dport 24000 counter dnat ip to 198.51.100.20:443 comment \"flvx forward:42 dnat tcp\"",
|
||||
"udp dport 24000 counter dnat ip to 198.51.100.20:443 comment \"flvx forward:42 dnat udp\"",
|
||||
"masquerade comment \"flvx masquerade\"",
|
||||
}
|
||||
for _, part := range expectedParts {
|
||||
if !strings.Contains(script, part) {
|
||||
t.Fatalf("script missing %q:\n%s", part, script)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableBracketsIPv6Target(t *testing.T) {
|
||||
script := RenderTable(NodePlan{
|
||||
NodeID: 10,
|
||||
Rules: []Rule{
|
||||
{ForwardID: 42, InPort: 24000, TargetHost: "2001:db8::1", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
},
|
||||
})
|
||||
if !strings.Contains(script, "dnat ip6 to [2001:db8::1]:443") {
|
||||
t.Fatalf("expected bracketed IPv6 dnat target, got:\n%s", script)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 42,
|
||||
InPort: 12345,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`,
|
||||
`udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"`,
|
||||
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
|
||||
`ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
|
||||
`ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 43,
|
||||
InPort: 12346,
|
||||
TargetHost: "2001:db8::20",
|
||||
TargetPort: 8443,
|
||||
Protocols: []string{"tcp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`,
|
||||
`ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
|
||||
`ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableAccountingCountersIncludeOriginalPort(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{
|
||||
{ForwardID: 42, InPort: 12345, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
{ForwardID: 43, InPort: 12346, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTablePreservesHostnameDNATAndSkipsAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 44,
|
||||
InPort: 12347,
|
||||
TargetHost: "example.com",
|
||||
TargetPort: 9443,
|
||||
Protocols: []string{"tcp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
want := `tcp dport 12347 counter dnat to example.com:9443 comment "flvx forward:44 dnat tcp"`
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
unwantedLines := []string{
|
||||
`dnat ip to example.com`,
|
||||
`ip daddr example.com`,
|
||||
`ip saddr example.com`,
|
||||
}
|
||||
for _, unwanted := range unwantedLines {
|
||||
if strings.Contains(got, unwanted) {
|
||||
t.Fatalf("RenderTable() unexpectedly contains %q\n%s", unwanted, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHashIsStable(t *testing.T) {
|
||||
rule := Rule{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}
|
||||
if RuleHash(rule) != RuleHash(rule) {
|
||||
t.Fatalf("expected stable rule hash")
|
||||
}
|
||||
if RuleHash(rule) == RuleHash(Rule{ForwardID: 42, InPort: 24001, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}) {
|
||||
t.Fatalf("expected hash to change when port changes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHashIgnoresBindIPWhenRenderingDoesNotUseIt(t *testing.T) {
|
||||
base := Rule{
|
||||
ForwardID: 42,
|
||||
InPort: 24000,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}
|
||||
withBind := base
|
||||
withBind.BindIP = "192.0.2.10"
|
||||
|
||||
if RuleHash(base) != RuleHash(withBind) {
|
||||
t.Fatalf("expected bind IP to be ignored by hash when it is not rendered")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,206 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/ssh"
|
||||
)
|
||||
|
||||
type Runner interface {
|
||||
ApplyScript(ctx context.Context, cfg SSHConfig, script string) error
|
||||
Test(ctx context.Context, cfg SSHConfig) error
|
||||
ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error)
|
||||
}
|
||||
|
||||
type SSHRunner struct {
|
||||
Timeout time.Duration
|
||||
}
|
||||
|
||||
func NewSSHRunner() *SSHRunner {
|
||||
return &SSHRunner{Timeout: 15 * time.Second}
|
||||
}
|
||||
|
||||
func (r *SSHRunner) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
return r.run(ctx, cfg, "command -v nft >/dev/null 2>&1 && nft --version >/dev/null 2>&1")
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
nft := nftBinary(cfg)
|
||||
command := "tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft) || exit 1\n" +
|
||||
"cleanup() {\n" +
|
||||
" rm -f \"$tmp\"\n" +
|
||||
"}\n" +
|
||||
"trap cleanup EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
nft + " -c -f \"$tmp\"\n" +
|
||||
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
|
||||
" " + nft + " delete table inet flvx\n" +
|
||||
"fi\n" +
|
||||
nft + " -f \"$tmp\""
|
||||
return r.run(ctx, cfg, command)
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
|
||||
return r.runOutput(ctx, cfg, nftBinary(cfg)+" -j list table inet flvx")
|
||||
}
|
||||
|
||||
func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) error {
|
||||
_, err := r.runOutput(ctx, cfg, command)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *SSHRunner) runOutput(ctx context.Context, cfg SSHConfig, command string) ([]byte, error) {
|
||||
timeout := r.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 15 * time.Second
|
||||
}
|
||||
runCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
clientConfig, err := buildSSHClientConfig(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
addr := net.JoinHostPort(strings.TrimSpace(cfg.Host), fmt.Sprintf("%d", normalizedSSHPort(cfg.Port)))
|
||||
dialer := net.Dialer{Timeout: timeout}
|
||||
conn, err := dialer.DialContext(runCtx, "tcp", addr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 连接失败: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 认证失败: %w", err)
|
||||
}
|
||||
client := ssh.NewClient(sshConn, chans, reqs)
|
||||
defer client.Close()
|
||||
|
||||
session, err := client.NewSession()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 会话创建失败: %w", err)
|
||||
}
|
||||
defer session.Close()
|
||||
|
||||
var stdout bytes.Buffer
|
||||
var stderr bytes.Buffer
|
||||
session.Stdout = &stdout
|
||||
session.Stderr = &stderr
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- session.Run(command)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-runCtx.Done():
|
||||
_ = session.Close()
|
||||
return nil, fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
message := strings.TrimSpace(stderr.String())
|
||||
if message != "" {
|
||||
return nil, fmt.Errorf("远程执行失败: %s: %w", message, err)
|
||||
}
|
||||
return nil, fmt.Errorf("远程执行失败: %w", err)
|
||||
}
|
||||
return stdout.Bytes(), nil
|
||||
}
|
||||
}
|
||||
|
||||
func buildSSHClientConfig(cfg SSHConfig) (*ssh.ClientConfig, error) {
|
||||
if strings.TrimSpace(cfg.Host) == "" {
|
||||
return nil, fmt.Errorf("SSH 主机不能为空")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Username) == "" {
|
||||
return nil, fmt.Errorf("SSH 用户名不能为空")
|
||||
}
|
||||
|
||||
auth, err := authMethods(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(auth) == 0 {
|
||||
return nil, fmt.Errorf("SSH 认证方式不能为空")
|
||||
}
|
||||
|
||||
return &ssh.ClientConfig{
|
||||
User: strings.TrimSpace(cfg.Username),
|
||||
Auth: auth,
|
||||
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
||||
Timeout: 15 * time.Second,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func authMethods(cfg SSHConfig) ([]ssh.AuthMethod, error) {
|
||||
switch strings.ToLower(strings.TrimSpace(cfg.AuthType)) {
|
||||
case "":
|
||||
if strings.TrimSpace(cfg.PrivateKey) == "" {
|
||||
return nil, fmt.Errorf("SSH 私钥不能为空")
|
||||
}
|
||||
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
|
||||
case "password":
|
||||
if cfg.Password == "" {
|
||||
return nil, fmt.Errorf("SSH 密码不能为空")
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.Password(cfg.Password)}, nil
|
||||
case "private_key":
|
||||
if strings.TrimSpace(cfg.PrivateKey) == "" {
|
||||
return nil, fmt.Errorf("SSH 私钥不能为空")
|
||||
}
|
||||
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("不支持的 SSH 认证方式: %s", cfg.AuthType)
|
||||
}
|
||||
}
|
||||
|
||||
func parsePrivateKey(privateKey, passphrase string) (ssh.Signer, error) {
|
||||
if passphrase != "" {
|
||||
signer, err := ssh.ParsePrivateKeyWithPassphrase([]byte(privateKey), []byte(passphrase))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
|
||||
}
|
||||
return signer, nil
|
||||
}
|
||||
signer, err := ssh.ParsePrivateKey([]byte(privateKey))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
|
||||
}
|
||||
return signer, nil
|
||||
}
|
||||
|
||||
func nftCommand(cfg SSHConfig, command string) string {
|
||||
return "sh -lc " + sshQuote(command)
|
||||
}
|
||||
|
||||
func nftBinary(cfg SSHConfig) string {
|
||||
if strings.EqualFold(strings.TrimSpace(cfg.SudoMode), "sudo") {
|
||||
return "sudo -n nft"
|
||||
}
|
||||
return "nft"
|
||||
}
|
||||
|
||||
func sshQuote(value string) string {
|
||||
return "'" + strings.ReplaceAll(value, "'", "'\"'\"'") + "'"
|
||||
}
|
||||
|
||||
func normalizedSSHPort(port int) int {
|
||||
if port <= 0 {
|
||||
return 22
|
||||
}
|
||||
return port
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
|
||||
privateKey := mustGeneratePrivateKey(t)
|
||||
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
|
||||
if err != nil {
|
||||
t.Fatalf("authMethods: %v", err)
|
||||
}
|
||||
if len(methods) != 1 {
|
||||
t.Fatalf("expected 1 auth method, got %d", len(methods))
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthMethodsDefaultPrivateKeyRequiresKey(t *testing.T) {
|
||||
_, err := authMethods(SSHConfig{})
|
||||
if err == nil || !strings.Contains(err.Error(), "SSH 私钥不能为空") {
|
||||
t.Fatalf("expected private key required error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustGeneratePrivateKey(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
key, err := rsa.GenerateKey(rand.Reader, 1024)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateKey: %v", err)
|
||||
}
|
||||
block := &pem.Block{
|
||||
Type: "RSA PRIVATE KEY",
|
||||
Bytes: x509.MarshalPKCS1PrivateKey(key),
|
||||
}
|
||||
return string(pem.EncodeToMemory(block))
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package nftables
|
||||
|
||||
const (
|
||||
ModeAgent = "agent"
|
||||
ModeNftables = "nftables"
|
||||
|
||||
StatusPending = "pending"
|
||||
StatusApplied = "applied"
|
||||
StatusError = "error"
|
||||
|
||||
CounterDirectionDNAT = "dnat"
|
||||
CounterDirectionToTarget = "to-target"
|
||||
CounterDirectionFromTarget = "from-target"
|
||||
)
|
||||
|
||||
type Target struct {
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
|
||||
type Rule struct {
|
||||
ForwardID int64
|
||||
InPort int
|
||||
BindIP string
|
||||
TargetHost string
|
||||
TargetPort int
|
||||
Protocols []string
|
||||
}
|
||||
|
||||
type NodePlan struct {
|
||||
NodeID int64
|
||||
Rules []Rule
|
||||
}
|
||||
|
||||
type SSHConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
AuthType string
|
||||
Password string
|
||||
PrivateKey string
|
||||
Passphrase string
|
||||
SudoMode string
|
||||
}
|
||||
|
||||
type ApplyResult struct {
|
||||
NodeID int64
|
||||
Script string
|
||||
Hashes map[int64]string
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func HashPassword(plain string) (string, error) {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(plain), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(hash), nil
|
||||
}
|
||||
|
||||
func VerifyPassword(storedHash, plain string) (bool, bool) {
|
||||
if strings.HasPrefix(storedHash, "$2") {
|
||||
return bcrypt.CompareHashAndPassword([]byte(storedHash), []byte(plain)) == nil, false
|
||||
}
|
||||
if MD5(plain) == storedHash {
|
||||
return true, true
|
||||
}
|
||||
return false, false
|
||||
}
|
||||
|
||||
func IsLegacyPasswordHash(storedHash string) bool {
|
||||
storedHash = strings.TrimSpace(storedHash)
|
||||
if len(storedHash) != 32 {
|
||||
return false
|
||||
}
|
||||
for _, r := range storedHash {
|
||||
if (r >= '0' && r <= '9') || (r >= 'a' && r <= 'f') || (r >= 'A' && r <= 'F') {
|
||||
continue
|
||||
}
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestHashPasswordProducesBcrypt(t *testing.T) {
|
||||
hash, err := HashPassword("admin_user")
|
||||
if err != nil {
|
||||
t.Fatalf("HashPassword() error = %v", err)
|
||||
}
|
||||
if len(hash) < 50 || !strings.HasPrefix(hash, "$2") {
|
||||
t.Fatalf("expected bcrypt hash, got %q", hash)
|
||||
}
|
||||
if ok, legacy := VerifyPassword(hash, "admin_user"); !ok || legacy {
|
||||
t.Fatalf("VerifyPassword() = (%v,%v), want (true,false)", ok, legacy)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyPasswordAcceptsLegacyMD5(t *testing.T) {
|
||||
if ok, legacy := VerifyPassword("3c85cdebade1c51cf64ca9f3c09d182d", "admin_user"); !ok || !legacy {
|
||||
t.Fatalf("VerifyPassword() = (%v,%v), want (true,true)", ok, legacy)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsLegacyPasswordHash(t *testing.T) {
|
||||
if !IsLegacyPasswordHash("3c85cdebade1c51cf64ca9f3c09d182d") {
|
||||
t.Fatal("expected 32-char hex MD5 hash to be legacy")
|
||||
}
|
||||
hash, err := HashPassword("admin_user")
|
||||
if err != nil {
|
||||
t.Fatalf("HashPassword() error = %v", err)
|
||||
}
|
||||
if IsLegacyPasswordHash(hash) {
|
||||
t.Fatalf("expected bcrypt hash not to be legacy: %q", hash)
|
||||
}
|
||||
}
|
||||
@@ -10,44 +10,47 @@ import "database/sql"
|
||||
// User maps to the "user" table. PostgreSQL treats "user" as a reserved
|
||||
// word, so TableName() is required for correct quoting.
|
||||
type User struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
User string `gorm:"column:user;type:varchar(100);not null"`
|
||||
Pwd string `gorm:"type:varchar(100);not null"`
|
||||
RoleID int `gorm:"column:role_id;not null"`
|
||||
ExpTime int64 `gorm:"column:exp_time;not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
Num int `gorm:"not null"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
User string `gorm:"column:user;type:varchar(100);not null"`
|
||||
Pwd string `gorm:"type:varchar(100);not null"`
|
||||
RoleID int `gorm:"column:role_id;not null"`
|
||||
ExpTime int64 `gorm:"column:exp_time;not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
Num int `gorm:"not null"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
PasswordChangedAt int64 `gorm:"column:password_changed_at;not null;default:0"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
}
|
||||
|
||||
func (User) TableName() string { return "user" }
|
||||
|
||||
// Forward maps to the "forward" table.
|
||||
type Forward struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID int64 `gorm:"column:user_id;not null"`
|
||||
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
|
||||
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
|
||||
InFlow int64 `gorm:"not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
|
||||
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
|
||||
ProxyProtocol int `gorm:"column:proxy_protocol;not null;default:0"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID int64 `gorm:"column:user_id;not null"`
|
||||
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
|
||||
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
|
||||
InFlow int64 `gorm:"not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
|
||||
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
|
||||
ProxyProtocol int `gorm:"column:proxy_protocol;not null;default:0"`
|
||||
ProxyProtocolReceive int `gorm:"column:proxy_protocol_receive;not null;default:0"`
|
||||
ProxyProtocolSend int `gorm:"column:proxy_protocol_send;not null;default:0"`
|
||||
}
|
||||
|
||||
func (Forward) TableName() string { return "forward" }
|
||||
@@ -86,6 +89,7 @@ type Node struct {
|
||||
UDPListenAddr string `gorm:"column:udp_listen_addr;type:varchar(100);not null;default:'[::]'"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
IsRemote int `gorm:"column:is_remote;default:0"`
|
||||
ForwardMode string `gorm:"column:forward_mode;type:varchar(20);not null;default:'agent'"`
|
||||
RemoteURL sql.NullString `gorm:"column:remote_url;type:text"`
|
||||
RemoteToken sql.NullString `gorm:"column:remote_token;type:text"`
|
||||
RemoteConfig sql.NullString `gorm:"column:remote_config;type:text"`
|
||||
@@ -94,6 +98,57 @@ type Node struct {
|
||||
|
||||
func (Node) TableName() string { return "node" }
|
||||
|
||||
type NodeSSHConfig struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex"`
|
||||
Host string `gorm:"type:varchar(255);not null"`
|
||||
Port int `gorm:"not null;default:22"`
|
||||
Username string `gorm:"type:varchar(100);not null"`
|
||||
AuthType string `gorm:"column:auth_type;type:varchar(20);not null"`
|
||||
Password sql.NullString `gorm:"type:text"`
|
||||
PrivateKey sql.NullString `gorm:"column:private_key;type:text"`
|
||||
Passphrase sql.NullString `gorm:"type:text"`
|
||||
SudoMode string `gorm:"column:sudo_mode;type:varchar(20);not null;default:'none'"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (NodeSSHConfig) TableName() string { return "node_ssh_config" }
|
||||
|
||||
type NftRuleBinding struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
ForwardID int64 `gorm:"column:forward_id;not null;uniqueIndex:idx_nft_rule_binding_forward_node;index"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_nft_rule_binding_forward_node;index"`
|
||||
InPort int `gorm:"column:in_port;not null"`
|
||||
Protocols string `gorm:"type:varchar(20);not null;default:'tcp,udp'"`
|
||||
TargetAddr string `gorm:"column:target_addr;type:text;not null"`
|
||||
BindIP string `gorm:"column:bind_ip;type:text;not null;default:''"`
|
||||
RuleHash string `gorm:"column:rule_hash;type:varchar(128);not null;default:''"`
|
||||
Status string `gorm:"type:varchar(20);not null;default:'pending'"`
|
||||
LastError string `gorm:"column:last_error;type:text;not null;default:''"`
|
||||
AppliedTime int64 `gorm:"column:applied_time;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (NftRuleBinding) TableName() string { return "nft_rule_binding" }
|
||||
|
||||
type NftCounterState struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||
ForwardID int64 `gorm:"column:forward_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||
Protocol string `gorm:"type:varchar(10);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||
Direction string `gorm:"type:varchar(20);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||
RuleHash string `gorm:"column:rule_hash;type:varchar(128);not null;default:''"`
|
||||
Bytes int64 `gorm:"not null;default:0"`
|
||||
Packets int64 `gorm:"not null;default:0"`
|
||||
CollectedTime int64 `gorm:"column:collected_time;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (NftCounterState) TableName() string { return "nft_counter_state" }
|
||||
|
||||
type SpeedLimit struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
@@ -434,24 +489,26 @@ type ChainTunnelBackup struct {
|
||||
}
|
||||
|
||||
type ForwardBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
UserID int64 `json:"userId"`
|
||||
UserName string `json:"userName"`
|
||||
Name string `json:"name"`
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
RemoteAddr string `json:"remoteAddr"`
|
||||
Strategy string `json:"strategy"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
Inx int `json:"inx"`
|
||||
SpeedID *int64 `json:"speedId,omitempty"`
|
||||
IPMaxConn int `json:"ipMaxConn,omitempty"`
|
||||
IPSpeedID *int64 `json:"ipSpeedId,omitempty"`
|
||||
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
||||
ProxyProtocol int `json:"proxyProtocol"`
|
||||
ID int64 `json:"id"`
|
||||
UserID int64 `json:"userId"`
|
||||
UserName string `json:"userName"`
|
||||
Name string `json:"name"`
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
RemoteAddr string `json:"remoteAddr"`
|
||||
Strategy string `json:"strategy"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
Inx int `json:"inx"`
|
||||
SpeedID *int64 `json:"speedId,omitempty"`
|
||||
IPMaxConn int `json:"ipMaxConn,omitempty"`
|
||||
IPSpeedID *int64 `json:"ipSpeedId,omitempty"`
|
||||
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
||||
ProxyProtocol int `json:"proxyProtocol"`
|
||||
ProxyProtocolReceive int `json:"proxyProtocolReceive,omitempty"`
|
||||
ProxyProtocolSend int `json:"proxyProtocolSend,omitempty"`
|
||||
}
|
||||
|
||||
type ForwardPortBackup struct {
|
||||
@@ -540,19 +597,21 @@ type ImportResult struct {
|
||||
|
||||
// ForwardRecord is a minimal forward view used by control plane and flow policy.
|
||||
type ForwardRecord struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
ProxyProtocol int
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
ProxyProtocol int
|
||||
ProxyProtocolReceive int
|
||||
ProxyProtocolSend int
|
||||
}
|
||||
|
||||
// TunnelRecord is a minimal tunnel view used by control plane.
|
||||
@@ -601,6 +660,7 @@ type NodeRecord struct {
|
||||
UDPListenAddr string
|
||||
InterfaceName string
|
||||
IsRemote int
|
||||
ForwardMode string
|
||||
RemoteURL string
|
||||
RemoteToken string
|
||||
RemoteConfig string
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
package repo
|
||||
|
||||
import "strings"
|
||||
|
||||
type ConfigAccessPolicy string
|
||||
|
||||
const (
|
||||
ConfigAccessPublic ConfigAccessPolicy = "public"
|
||||
ConfigAccessSensitive ConfigAccessPolicy = "sensitive"
|
||||
)
|
||||
|
||||
var publicConfigKeys = map[string]struct{}{
|
||||
"app_name": {},
|
||||
"app_logo": {},
|
||||
"app_favicon": {},
|
||||
"app_bg_image": {},
|
||||
"cloudflare_site_key": {},
|
||||
}
|
||||
|
||||
var sensitiveConfigKeys = map[string]struct{}{
|
||||
"jwt_secret": {},
|
||||
"license_key": {},
|
||||
"cloudflare_secret_key": {},
|
||||
}
|
||||
|
||||
func PolicyForConfig(name string) ConfigAccessPolicy {
|
||||
if IsPublicConfigKey(name) {
|
||||
return ConfigAccessPublic
|
||||
}
|
||||
if IsSensitiveConfigKey(name) {
|
||||
return ConfigAccessSensitive
|
||||
}
|
||||
return ConfigAccessSensitive
|
||||
}
|
||||
|
||||
func IsPublicConfigKey(name string) bool {
|
||||
_, ok := publicConfigKeys[normalizeConfigKey(name)]
|
||||
return ok
|
||||
}
|
||||
|
||||
func IsSensitiveConfigKey(name string) bool {
|
||||
_, ok := sensitiveConfigKeys[normalizeConfigKey(name)]
|
||||
return ok
|
||||
}
|
||||
|
||||
func FilterSensitiveConfigs(in map[string]string) map[string]string {
|
||||
if len(in) == 0 {
|
||||
return map[string]string{}
|
||||
}
|
||||
out := make(map[string]string, len(in))
|
||||
for name, value := range in {
|
||||
if IsSensitiveConfigKey(name) {
|
||||
continue
|
||||
}
|
||||
out[name] = value
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func normalizeConfigKey(name string) string {
|
||||
return strings.ToLower(strings.TrimSpace(name))
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package repo
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestConfigPolicy(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
key string
|
||||
want ConfigAccessPolicy
|
||||
}{
|
||||
{name: "app_name is public", key: "app_name", want: ConfigAccessPublic},
|
||||
{name: "app_logo is public", key: "app_logo", want: ConfigAccessPublic},
|
||||
{name: "app_favicon is public", key: "app_favicon", want: ConfigAccessPublic},
|
||||
{name: "app_bg_image is public", key: "app_bg_image", want: ConfigAccessPublic},
|
||||
{name: "cloudflare_site_key is public", key: "cloudflare_site_key", want: ConfigAccessPublic},
|
||||
{name: "jwt_secret is sensitive", key: "jwt_secret", want: ConfigAccessSensitive},
|
||||
{name: "license_key is sensitive", key: "license_key", want: ConfigAccessSensitive},
|
||||
{name: "cloudflare_secret_key is sensitive", key: "cloudflare_secret_key", want: ConfigAccessSensitive},
|
||||
{name: "trimmed public key is public", key: " APP_NAME ", want: ConfigAccessPublic},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := PolicyForConfig(tt.key); got != tt.want {
|
||||
t.Fatalf("PolicyForConfig(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigPolicyHelpers(t *testing.T) {
|
||||
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "cloudflare_site_key"}
|
||||
for _, key := range publicKeys {
|
||||
if !IsPublicConfigKey(key) {
|
||||
t.Fatalf("expected %q to be public", key)
|
||||
}
|
||||
}
|
||||
|
||||
sensitiveKeys := []string{"jwt_secret", "license_key", "cloudflare_secret_key"}
|
||||
for _, key := range sensitiveKeys {
|
||||
if !IsSensitiveConfigKey(key) {
|
||||
t.Fatalf("expected %q to be sensitive", key)
|
||||
}
|
||||
}
|
||||
|
||||
input := map[string]string{
|
||||
"app_name": "FLVX",
|
||||
"license_key": "secret-license",
|
||||
"cloudflare_secret_key": "secret-cloudflare",
|
||||
"jwt_secret": "secret-jwt",
|
||||
"cloudflare_site_key": "site-key",
|
||||
}
|
||||
filtered := FilterSensitiveConfigs(input)
|
||||
if len(filtered) != 2 {
|
||||
t.Fatalf("expected 2 public configs, got %d", len(filtered))
|
||||
}
|
||||
if filtered["app_name"] != "FLVX" || filtered["cloudflare_site_key"] != "site-key" {
|
||||
t.Fatalf("unexpected filtered configs: %+v", filtered)
|
||||
}
|
||||
if _, ok := filtered["jwt_secret"]; ok {
|
||||
t.Fatal("expected jwt_secret to be filtered out")
|
||||
}
|
||||
if _, ok := filtered["license_key"]; ok {
|
||||
t.Fatal("expected license_key to be filtered out")
|
||||
}
|
||||
if _, ok := filtered["cloudflare_secret_key"]; ok {
|
||||
t.Fatal("expected cloudflare_secret_key to be filtered out")
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,9 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
@@ -19,6 +21,7 @@ import (
|
||||
"gorm.io/gorm/clause"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
@@ -101,6 +104,19 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
|
||||
return nil
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
return applyFlowUploadDeltasTx(tx, deltas)
|
||||
})
|
||||
}
|
||||
|
||||
func applyFlowUploadDeltasTx(tx *gorm.DB, deltas []FlowUploadCounterDelta) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
if len(deltas) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
forwardTotals := make(map[int64][2]int64, len(deltas))
|
||||
userTotals := make(map[int64][2]int64, len(deltas))
|
||||
userTunnelTotals := make(map[int64][2]int64, len(deltas))
|
||||
@@ -125,36 +141,34 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
|
||||
}
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) {
|
||||
total := forwardTotals[forwardID]
|
||||
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) {
|
||||
total := forwardTotals[forwardID]
|
||||
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, userID := range sortedFlowUploadTargetIDs(userTotals) {
|
||||
total := userTotals[userID]
|
||||
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, userID := range sortedFlowUploadTargetIDs(userTotals) {
|
||||
total := userTotals[userID]
|
||||
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) {
|
||||
total := userTunnelTotals[userTunnelID]
|
||||
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) {
|
||||
total := userTunnelTotals[userTunnelID]
|
||||
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{
|
||||
"in_flow": gorm.Expr("in_flow + ?", total[0]),
|
||||
"out_flow": gorm.Expr("out_flow + ?", total[1]),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ─── Open / Close ────────────────────────────────────────────────────
|
||||
@@ -186,7 +200,6 @@ func Open(path string) (*Repository, error) {
|
||||
_ = sqlDB.Close()
|
||||
return nil, fmt.Errorf("prepare sqlite legacy schema: %w", err)
|
||||
}
|
||||
|
||||
if err := autoMigrateAll(db); err != nil {
|
||||
_ = sqlDB.Close()
|
||||
return nil, fmt.Errorf("auto migrate: %w", err)
|
||||
@@ -266,12 +279,21 @@ func (r *Repository) Close() error {
|
||||
}
|
||||
|
||||
func autoMigrateAll(db *gorm.DB) error {
|
||||
if db.Dialector.Name() == "sqlite" {
|
||||
if err := prepareSQLiteNftablesColumns(db); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
models := []interface{}{
|
||||
&model.User{},
|
||||
&model.UserQuota{},
|
||||
&model.Forward{},
|
||||
&model.ForwardPort{},
|
||||
&model.Node{},
|
||||
&model.NodeSSHConfig{},
|
||||
&model.NftRuleBinding{},
|
||||
&model.NftCounterState{},
|
||||
&model.SpeedLimit{},
|
||||
&model.StatisticsFlow{},
|
||||
&model.Tunnel{},
|
||||
@@ -304,6 +326,7 @@ func autoMigrateAll(db *gorm.DB) error {
|
||||
m := db.Migrator()
|
||||
hasNode := m.HasTable(&model.Node{})
|
||||
hasTunnel := m.HasTable(&model.Tunnel{})
|
||||
hasForward := m.HasTable(&model.Forward{})
|
||||
|
||||
for _, item := range models {
|
||||
if hasNode {
|
||||
@@ -316,6 +339,11 @@ func autoMigrateAll(db *gorm.DB) error {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if hasForward {
|
||||
if _, ok := item.(*model.Forward); ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if err := db.AutoMigrate(item); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -396,7 +424,7 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
}
|
||||
|
||||
if m.HasTable(&model.Forward{}) {
|
||||
for _, field := range []string{"ProxyProtocol"} {
|
||||
for _, field := range []string{"MaxConn", "IPMaxConn", "IPSpeedID", "ProxyProtocol", "ProxyProtocolReceive", "ProxyProtocolSend"} {
|
||||
if m.HasColumn(&model.Forward{}, field) {
|
||||
continue
|
||||
}
|
||||
@@ -409,14 +437,39 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func seedData(db *gorm.DB) {
|
||||
adminUser := model.User{
|
||||
ID: 1, User: "admin_user", Pwd: "3c85cdebade1c51cf64ca9f3c09d182d",
|
||||
RoleID: 0, ExpTime: 2727251700000, Flow: 99999, InFlow: 0, OutFlow: 0,
|
||||
FlowResetTime: 1, Num: 99999, CreatedTime: 1748914865000,
|
||||
UpdatedTime: sql.NullInt64{Int64: 1754011744252, Valid: true}, Status: 1,
|
||||
func prepareSQLiteNftablesColumns(db *gorm.DB) error {
|
||||
if db == nil || db.Dialector.Name() != "sqlite" {
|
||||
return nil
|
||||
}
|
||||
if !db.Migrator().HasTable(&model.Node{}) {
|
||||
return nil
|
||||
}
|
||||
if !db.Migrator().HasColumn(&model.Node{}, "forward_mode") {
|
||||
if err := db.Exec("ALTER TABLE node ADD COLUMN forward_mode varchar(20) NOT NULL DEFAULT 'agent'").Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func seedData(db *gorm.DB) {
|
||||
var adminCount int64
|
||||
if err := db.Model(&model.User{}).Where("id = ?", 1).Count(&adminCount).Error; err == nil && adminCount == 0 {
|
||||
adminPwd, err := security.HashPassword("admin_user")
|
||||
if err != nil {
|
||||
log.Printf("seed admin password hash failed: %v", err)
|
||||
} else {
|
||||
adminUser := model.User{
|
||||
ID: 1, User: "admin_user", Pwd: adminPwd,
|
||||
RoleID: 0, ExpTime: 2727251700000, Flow: 99999, InFlow: 0, OutFlow: 0,
|
||||
FlowResetTime: 1, Num: 99999, CreatedTime: 1748914865000,
|
||||
UpdatedTime: sql.NullInt64{Int64: 1754011744252, Valid: true},
|
||||
Status: 1,
|
||||
PasswordChangedAt: 1748914865000,
|
||||
}
|
||||
db.Create(&adminUser)
|
||||
}
|
||||
}
|
||||
db.Where("id = ?", 1).FirstOrCreate(&adminUser)
|
||||
|
||||
appNameConfig := model.ViteConfig{ID: 1, Name: "app_name", Value: "flux", Time: 1755147963000}
|
||||
db.Where("id = ?", 1).FirstOrCreate(&appNameConfig)
|
||||
@@ -480,9 +533,10 @@ func (r *Repository) UpdateUserNameAndPassword(userID int64, username, passwordM
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.User{}).Where("id = ?", userID).Updates(map[string]interface{}{
|
||||
"user": username,
|
||||
"pwd": passwordMD5,
|
||||
"updated_time": now,
|
||||
"user": username,
|
||||
"pwd": passwordMD5,
|
||||
"password_changed_at": now,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -773,6 +827,28 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
|
||||
if err := r.db.Order("inx ASC, id ASC").Find(&nodes).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeIDs := make([]int64, 0, len(nodes))
|
||||
for _, n := range nodes {
|
||||
if defaultNodeForwardMode(n.ForwardMode) == "nftables" {
|
||||
nodeIDs = append(nodeIDs, n.ID)
|
||||
}
|
||||
}
|
||||
sshConfigByNodeID := make(map[int64]map[string]interface{}, len(nodeIDs))
|
||||
if len(nodeIDs) > 0 {
|
||||
var configs []model.NodeSSHConfig
|
||||
if err := r.db.Where("node_id IN ?", nodeIDs).Find(&configs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, cfg := range configs {
|
||||
sshConfigByNodeID[cfg.NodeID] = map[string]interface{}{
|
||||
"host": cfg.Host,
|
||||
"port": cfg.Port,
|
||||
"username": cfg.Username,
|
||||
"authType": cfg.AuthType,
|
||||
"sudoMode": cfg.SudoMode,
|
||||
}
|
||||
}
|
||||
}
|
||||
items := make([]map[string]interface{}, 0, len(nodes))
|
||||
for _, n := range nodes {
|
||||
items = append(items, map[string]interface{}{
|
||||
@@ -790,11 +866,13 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
|
||||
"version": nullableString(n.Version),
|
||||
"http": n.HTTP, "tls": n.TLS, "socks": n.Socks,
|
||||
"status": n.Status, "isRemote": n.IsRemote,
|
||||
"forwardMode": defaultNodeForwardMode(n.ForwardMode),
|
||||
"remoteUrl": nullableString(n.RemoteURL),
|
||||
"remoteToken": nullableString(n.RemoteToken),
|
||||
"remoteConfig": nullableString(n.RemoteConfig),
|
||||
"expiryReminderDismissed": n.ExpiryReminderDismissed,
|
||||
"interfaceName": nullableString(n.InterfaceName),
|
||||
"sshConfig": sshConfigByNodeID[n.ID],
|
||||
})
|
||||
}
|
||||
return items, nil
|
||||
@@ -866,31 +944,33 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
}
|
||||
|
||||
type fwdRow struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
TunnelName string
|
||||
TrafficRatio float64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
CreatedTime int64
|
||||
Status int
|
||||
Inx int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
IPSpeedLimitName string
|
||||
ProxyProtocol int
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
TunnelName string
|
||||
TrafficRatio float64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
CreatedTime int64
|
||||
Status int
|
||||
Inx int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
IPSpeedLimitName string
|
||||
ProxyProtocol int
|
||||
ProxyProtocolReceive int
|
||||
ProxyProtocolSend int
|
||||
}
|
||||
|
||||
var rows []fwdRow
|
||||
err := r.db.Model(&model.Forward{}).
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.ip_max_conn, forward.ip_speed_id, COALESCE(ip_speed_limit.name, '') AS ip_speed_limit_name, forward.proxy_protocol").
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.ip_max_conn, forward.ip_speed_id, COALESCE(ip_speed_limit.name, '') AS ip_speed_limit_name, forward.proxy_protocol, forward.proxy_protocol_receive, forward.proxy_protocol_send").
|
||||
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
|
||||
Joins("LEFT JOIN speed_limit AS ip_speed_limit ON ip_speed_limit.id = forward.ip_speed_id").
|
||||
Order("forward.inx ASC, forward.id ASC").
|
||||
@@ -901,6 +981,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
|
||||
items := make([]map[string]interface{}, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(row.ProxyProtocol, row.ProxyProtocolReceive, row.ProxyProtocolSend)
|
||||
inIP, inPort, err := resolveForwardIngress(r.db, row.ID, row.TunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -913,9 +994,11 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
|
||||
"inFlow": row.InFlow, "outFlow": row.OutFlow,
|
||||
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
|
||||
"maxConn": row.MaxConn,
|
||||
"ipMaxConn": row.IPMaxConn,
|
||||
"proxyProtocol": row.ProxyProtocol,
|
||||
"maxConn": row.MaxConn,
|
||||
"ipMaxConn": row.IPMaxConn,
|
||||
"proxyProtocol": row.ProxyProtocol,
|
||||
"proxyProtocolReceive": proxyProtocolReceive,
|
||||
"proxyProtocolSend": proxyProtocolSend,
|
||||
}
|
||||
if row.SpeedID.Valid {
|
||||
item["speedId"] = row.SpeedID.Int64
|
||||
@@ -1884,7 +1967,7 @@ func (r *Repository) ExportAll() (*model.BackupData, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("export configs failed: %w", err)
|
||||
}
|
||||
backup.Configs = configs
|
||||
backup.Configs = FilterSensitiveConfigs(configs)
|
||||
|
||||
return backup, nil
|
||||
}
|
||||
@@ -1964,7 +2047,7 @@ func (r *Repository) ExportPartial(types []string) (*model.BackupData, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("export configs failed: %w", err)
|
||||
}
|
||||
backup.Configs = v
|
||||
backup.Configs = FilterSensitiveConfigs(v)
|
||||
}
|
||||
return backup, nil
|
||||
}
|
||||
@@ -2117,8 +2200,10 @@ func (r *Repository) exportForwards() ([]model.ForwardBackup, error) {
|
||||
TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy,
|
||||
InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime,
|
||||
UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: f.ProxyProtocolReceive,
|
||||
ProxyProtocolSend: f.ProxyProtocolSend,
|
||||
}
|
||||
if f.SpeedID.Valid {
|
||||
v := f.SpeedID.Int64
|
||||
@@ -2346,26 +2431,31 @@ func (r *Repository) Import(backup *model.BackupData, types []string) (*model.Im
|
||||
func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error) {
|
||||
count := 0
|
||||
for _, u := range users {
|
||||
item := model.User{
|
||||
ID: u.ID,
|
||||
User: u.User,
|
||||
Pwd: u.Pwd,
|
||||
RoleID: u.RoleID,
|
||||
ExpTime: u.ExpTime,
|
||||
Flow: u.Flow,
|
||||
InFlow: u.InFlow,
|
||||
OutFlow: u.OutFlow,
|
||||
FlowResetTime: u.FlowResetTime,
|
||||
Num: u.Num,
|
||||
CreatedTime: u.CreatedTime,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: u.Status,
|
||||
pwdHash, status, err := normalizeImportedUserPassword(u.Pwd, u.Status)
|
||||
if err != nil {
|
||||
return count, err
|
||||
}
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
item := model.User{
|
||||
ID: u.ID,
|
||||
User: u.User,
|
||||
Pwd: pwdHash,
|
||||
RoleID: u.RoleID,
|
||||
ExpTime: u.ExpTime,
|
||||
Flow: u.Flow,
|
||||
InFlow: u.InFlow,
|
||||
OutFlow: u.OutFlow,
|
||||
FlowResetTime: u.FlowResetTime,
|
||||
Num: u.Num,
|
||||
CreatedTime: u.CreatedTime,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: status,
|
||||
PasswordChangedAt: now,
|
||||
}
|
||||
err = tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"user", "pwd", "role_id", "exp_time", "flow", "in_flow", "out_flow",
|
||||
"flow_reset_time", "num", "updated_time", "status",
|
||||
"flow_reset_time", "num", "updated_time", "status", "password_changed_at",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
if err != nil {
|
||||
@@ -2409,6 +2499,33 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func normalizeImportedUserPassword(password string, status int) (string, int, error) {
|
||||
password = strings.TrimSpace(password)
|
||||
if strings.HasPrefix(password, "$2") {
|
||||
return password, status, nil
|
||||
}
|
||||
if security.IsLegacyPasswordHash(password) || password == "" {
|
||||
replacement, err := randomPasswordHash()
|
||||
if err != nil {
|
||||
return "", status, err
|
||||
}
|
||||
return replacement, 0, nil
|
||||
}
|
||||
hash, err := security.HashPassword(password)
|
||||
if err != nil {
|
||||
return "", status, err
|
||||
}
|
||||
return hash, status, nil
|
||||
}
|
||||
|
||||
func randomPasswordHash() (string, error) {
|
||||
buf := make([]byte, 32)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return security.HashPassword(hex.EncodeToString(buf))
|
||||
}
|
||||
|
||||
func importNodes(tx *gorm.DB, nodes []model.NodeBackup, now int64) (int, error) {
|
||||
count := 0
|
||||
for _, n := range nodes {
|
||||
@@ -2520,29 +2637,31 @@ func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int
|
||||
count := 0
|
||||
for _, f := range forwards {
|
||||
item := model.Forward{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
InFlow: f.InFlow,
|
||||
OutFlow: f.OutFlow,
|
||||
CreatedTime: f.CreatedTime,
|
||||
UpdatedTime: now,
|
||||
Status: f.Status,
|
||||
Inx: f.Inx,
|
||||
SpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.SpeedID), Valid: f.SpeedID != nil && *f.SpeedID > 0},
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.IPSpeedID), Valid: f.IPSpeedID != nil && *f.IPSpeedID > 0},
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
InFlow: f.InFlow,
|
||||
OutFlow: f.OutFlow,
|
||||
CreatedTime: f.CreatedTime,
|
||||
UpdatedTime: now,
|
||||
Status: f.Status,
|
||||
Inx: f.Inx,
|
||||
SpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.SpeedID), Valid: f.SpeedID != nil && *f.SpeedID > 0},
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.IPSpeedID), Valid: f.IPSpeedID != nil && *f.IPSpeedID > 0},
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: f.ProxyProtocolReceive,
|
||||
ProxyProtocolSend: f.ProxyProtocolSend,
|
||||
}
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy",
|
||||
"in_flow", "out_flow", "updated_time", "status", "inx", "speed_id", "ip_max_conn", "ip_speed_id", "proxy_protocol",
|
||||
"in_flow", "out_flow", "updated_time", "status", "inx", "speed_id", "ip_max_conn", "ip_speed_id", "proxy_protocol", "proxy_protocol_receive", "proxy_protocol_send",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
if err != nil {
|
||||
@@ -2720,6 +2839,7 @@ func importPermissions(tx *gorm.DB, permissions []model.PermissionBackup, _ int6
|
||||
}
|
||||
|
||||
func importConfigs(tx *gorm.DB, configs map[string]string, now int64) (int, error) {
|
||||
configs = FilterSensitiveConfigs(configs)
|
||||
count := 0
|
||||
for name, value := range configs {
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func (r *Repository) GetUserAuthState(userID int64) (*auth.UserAuthState, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var user struct {
|
||||
ID int64 `gorm:"column:id"`
|
||||
RoleID int `gorm:"column:role_id"`
|
||||
Status int `gorm:"column:status"`
|
||||
PasswordChangedAt int64 `gorm:"column:password_changed_at"`
|
||||
}
|
||||
if err := r.db.Model(&model.User{}).Select("id", "role_id", "status", "password_changed_at").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
return nil, normalizeNotFoundErr(err)
|
||||
}
|
||||
return &auth.UserAuthState{
|
||||
ID: user.ID,
|
||||
RoleID: user.RoleID,
|
||||
Status: user.Status,
|
||||
PasswordChangedAt: user.PasswordChangedAt,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestGetUserAuthStateReturnsPasswordChangedAt(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "auth.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
userID, err := r.CreateUser("admin_user", "pwd", 0, 2727251700000, 99999, 1, 99999, 1, 0, now)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateUser() error = %v", err)
|
||||
}
|
||||
|
||||
state, err := r.GetUserAuthState(userID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUserAuthState() error = %v", err)
|
||||
}
|
||||
if state == nil || state.PasswordChangedAt != now || state.Status != 1 || state.RoleID != 0 {
|
||||
t.Fatalf("unexpected auth state: %+v", state)
|
||||
}
|
||||
}
|
||||
@@ -4,8 +4,33 @@ import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestSeedDataDefaultAdminUsesBcrypt(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "seed.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
admin, err := r.GetUserByUsername("admin_user")
|
||||
if err != nil {
|
||||
t.Fatalf("get admin user: %v", err)
|
||||
}
|
||||
if admin == nil {
|
||||
t.Fatal("expected seeded admin user")
|
||||
}
|
||||
if security.IsLegacyPasswordHash(admin.Pwd) {
|
||||
t.Fatalf("seeded admin password is legacy MD5: %q", admin.Pwd)
|
||||
}
|
||||
if ok, legacy := security.VerifyPassword(admin.Pwd, "admin_user"); !ok || legacy {
|
||||
t.Fatalf("VerifyPassword() = (%v,%v), want (true,false)", ok, legacy)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackupRoundTripsTunnelProbeTarget(t *testing.T) {
|
||||
source, err := Open(filepath.Join(t.TempDir(), "source.db"))
|
||||
if err != nil {
|
||||
@@ -57,3 +82,157 @@ func TestBackupRoundTripsTunnelProbeTarget(t *testing.T) {
|
||||
t.Fatalf("unexpected imported probe target: %+v", items[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "export.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
seedConfig(t, r, "app_name", "FLVX")
|
||||
seedConfig(t, r, "app_logo", "logo")
|
||||
seedConfig(t, r, "app_favicon", "favicon")
|
||||
seedConfig(t, r, "app_bg_image", "bg")
|
||||
seedConfig(t, r, "cloudflare_site_key", "site-key")
|
||||
seedConfig(t, r, "jwt_secret", "jwt-secret")
|
||||
seedConfig(t, r, "license_key", "license-secret")
|
||||
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-secret")
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
export func() (*model.BackupData, error)
|
||||
}{
|
||||
{name: "ExportAll", export: r.ExportAll},
|
||||
{name: "ExportPartial", export: func() (*model.BackupData, error) { return r.ExportPartial([]string{"configs"}) }},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
backup, err := tc.export()
|
||||
if err != nil {
|
||||
t.Fatalf("export backup: %v", err)
|
||||
}
|
||||
if backup.Configs["app_name"] != "FLVX" {
|
||||
t.Fatalf("expected public config in export, got %+v", backup.Configs)
|
||||
}
|
||||
if backup.Configs["cloudflare_site_key"] != "site-key" {
|
||||
t.Fatalf("expected public config in export, got %+v", backup.Configs)
|
||||
}
|
||||
for _, key := range []string{"jwt_secret", "license_key", "cloudflare_secret_key"} {
|
||||
if _, ok := backup.Configs[key]; ok {
|
||||
t.Fatalf("expected %s to be omitted from export, got %+v", key, backup.Configs)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportIgnoresSensitiveConfigs(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "import.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
seedConfig(t, r, "app_name", "before")
|
||||
seedConfig(t, r, "jwt_secret", "jwt-before")
|
||||
seedConfig(t, r, "license_key", "license-before")
|
||||
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
|
||||
backup := &model.BackupData{Configs: map[string]string{
|
||||
"app_name": "after",
|
||||
"jwt_secret": "jwt-after",
|
||||
"license_key": "license-after",
|
||||
"cloudflare_secret_key": "cloudflare-after",
|
||||
}}
|
||||
|
||||
result, err := r.Import(backup, []string{"configs"})
|
||||
if err != nil {
|
||||
t.Fatalf("import backup: %v", err)
|
||||
}
|
||||
if result.ConfigsImported != 1 {
|
||||
t.Fatalf("expected one imported config, got %d", result.ConfigsImported)
|
||||
}
|
||||
|
||||
assertConfigValue(t, r, "app_name", "after")
|
||||
assertConfigValue(t, r, "jwt_secret", "jwt-before")
|
||||
assertConfigValue(t, r, "license_key", "license-before")
|
||||
assertConfigValue(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
}
|
||||
|
||||
func TestImportUsersDoesNotStoreLegacyMD5Passwords(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "legacy-user-import.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
backup := &model.BackupData{
|
||||
Version: "1.0",
|
||||
Users: []model.UserBackup{{
|
||||
ID: 55,
|
||||
User: "legacy-import-user",
|
||||
Pwd: "3c85cdebade1c51cf64ca9f3c09d182d",
|
||||
RoleID: 1,
|
||||
ExpTime: 2727251700000,
|
||||
Flow: 99999,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: 1,
|
||||
Num: 99999,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
}},
|
||||
}
|
||||
|
||||
result, err := r.Import(backup, []string{"users"})
|
||||
if err != nil {
|
||||
t.Fatalf("Import() error = %v", err)
|
||||
}
|
||||
if result.UsersImported != 1 {
|
||||
t.Fatalf("UsersImported = %d, want 1", result.UsersImported)
|
||||
}
|
||||
|
||||
user, err := r.GetUserByUsername("legacy-import-user")
|
||||
if err != nil {
|
||||
t.Fatalf("get imported user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
t.Fatal("expected imported user")
|
||||
}
|
||||
if security.IsLegacyPasswordHash(user.Pwd) {
|
||||
t.Fatalf("imported password remained legacy MD5: %q", user.Pwd)
|
||||
}
|
||||
if ok, _ := security.VerifyPassword(user.Pwd, "admin_user"); ok {
|
||||
t.Fatal("legacy imported password should not remain usable")
|
||||
}
|
||||
if user.Status != 0 {
|
||||
t.Fatalf("legacy imported user status = %d, want disabled status 0", user.Status)
|
||||
}
|
||||
if user.PasswordChangedAt <= 0 {
|
||||
t.Fatalf("PasswordChangedAt = %d, want import revocation timestamp", user.PasswordChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func seedConfig(t *testing.T, r *Repository, name, value string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, name, value, time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertConfigValue(t *testing.T, r *Repository, name, want string) {
|
||||
t.Helper()
|
||||
cfg, err := r.GetConfigByName(name)
|
||||
if err != nil {
|
||||
t.Fatalf("get config %s: %v", name, err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != want {
|
||||
t.Fatalf("expected config %s=%q, got %+v", name, want, cfg)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -44,20 +44,23 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -233,7 +236,8 @@ func nodeRecordFromModel(n *model.Node) *model.NodeRecord {
|
||||
Status: n.Status,
|
||||
PortRange: n.Port,
|
||||
TCPListenAddr: n.TCPListenAddr, UDPListenAddr: n.UDPListenAddr,
|
||||
IsRemote: n.IsRemote,
|
||||
IsRemote: n.IsRemote,
|
||||
ForwardMode: defaultNodeForwardMode(n.ForwardMode),
|
||||
}
|
||||
if n.ServerIPV4.Valid {
|
||||
rec.ServerIPv4 = strings.TrimSpace(n.ServerIPV4.String)
|
||||
|
||||
@@ -11,6 +11,8 @@ import (
|
||||
|
||||
type FlowUploadForwardMeta struct {
|
||||
ForwardID int64
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
TunnelID int64
|
||||
TrafficRatio float64
|
||||
TunnelFlow int64
|
||||
@@ -59,6 +61,8 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
|
||||
type row struct {
|
||||
ForwardID int64 `gorm:"column:forward_id"`
|
||||
UserID int64 `gorm:"column:user_id"`
|
||||
UserTunnelID int64 `gorm:"column:user_tunnel_id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id"`
|
||||
TrafficRatio float64 `gorm:"column:traffic_ratio"`
|
||||
TunnelFlow int64 `gorm:"column:tunnel_flow"`
|
||||
@@ -68,8 +72,9 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
for _, chunk := range chunkFlowUploadForwardIDs(ids) {
|
||||
var rows []row
|
||||
err := r.db.Table("forward AS f").
|
||||
Select("f.id AS forward_id, f.tunnel_id AS tunnel_id, t.traffic_ratio AS traffic_ratio, t.flow AS tunnel_flow").
|
||||
Select("f.id AS forward_id, f.user_id AS user_id, COALESCE(ut.id, 0) AS user_tunnel_id, f.tunnel_id AS tunnel_id, t.traffic_ratio AS traffic_ratio, t.flow AS tunnel_flow").
|
||||
Joins("LEFT JOIN tunnel t ON t.id = f.tunnel_id").
|
||||
Joins("LEFT JOIN user_tunnel ut ON ut.user_id = f.user_id AND ut.tunnel_id = f.tunnel_id").
|
||||
Where("f.id IN ?", chunk).
|
||||
Scan(&rows).Error
|
||||
if err != nil {
|
||||
@@ -84,6 +89,8 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
}
|
||||
out[row.ForwardID] = FlowUploadForwardMeta{
|
||||
ForwardID: row.ForwardID,
|
||||
UserID: row.UserID,
|
||||
UserTunnelID: row.UserTunnelID,
|
||||
TunnelID: row.TunnelID,
|
||||
TrafficRatio: row.TrafficRatio,
|
||||
TunnelFlow: row.TunnelFlow,
|
||||
@@ -113,20 +120,23 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -148,20 +158,23 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -183,20 +196,23 @@ func (r *Repository) ListForwardsByUserAndTunnel(userID, tunnelID int64) ([]mode
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -219,20 +235,23 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
fr := model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
}
|
||||
if strings.TrimSpace(fr.Strategy) == "" {
|
||||
fr.Strategy = "fifo"
|
||||
|
||||
@@ -64,7 +64,7 @@ func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("get metas: %v", err)
|
||||
}
|
||||
if metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
|
||||
if metas[20].UserID != 2 || metas[20].UserTunnelID != 10 || metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
|
||||
t.Fatalf("unexpected meta for forward 20: %#v", metas[20])
|
||||
}
|
||||
if _, ok := metas[99]; ok {
|
||||
@@ -86,6 +86,99 @@ func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNftTrafficAccountingAppliesFlowQuotaAndStates(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
seedFlowBatchRows(t, r, nowMs)
|
||||
|
||||
quotaViews, err := r.ApplyNftTrafficAccounting(
|
||||
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
|
||||
map[int64]int64{2: 1140},
|
||||
[]NftCounterStateInput{{
|
||||
NodeID: 11,
|
||||
ForwardID: 20,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1400,
|
||||
Packets: 14,
|
||||
CollectedTime: nowMs,
|
||||
}},
|
||||
now,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("ApplyNftTrafficAccounting: %v", err)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 480 {
|
||||
t.Fatalf("expected forward in_flow=480, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 660 {
|
||||
t.Fatalf("expected user out_flow=660, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 480 {
|
||||
t.Fatalf("expected user_tunnel in_flow=480, got %d", got)
|
||||
}
|
||||
if quotaViews[2] == nil || quotaViews[2].DailyUsedBytes != 1140 || quotaViews[2].MonthlyUsedBytes != 1140 {
|
||||
t.Fatalf("unexpected quota view: %#v", quotaViews[2])
|
||||
}
|
||||
states, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(states) != 1 || states[0].ForwardID != 20 || states[0].Bytes != 1400 {
|
||||
t.Fatalf("unexpected nft counter state: %+v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNftTrafficAccountingRollsBackFlowAndQuotaWhenStateWriteFails(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting-rollback.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
seedFlowBatchRows(t, r, nowMs)
|
||||
if err := r.DB().Exec(`DROP TABLE nft_counter_state`).Error; err != nil {
|
||||
t.Fatalf("drop nft_counter_state: %v", err)
|
||||
}
|
||||
|
||||
_, err = r.ApplyNftTrafficAccounting(
|
||||
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
|
||||
map[int64]int64{2: 1140},
|
||||
[]NftCounterStateInput{{
|
||||
NodeID: 11,
|
||||
ForwardID: 20,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1400,
|
||||
Packets: 14,
|
||||
CollectedTime: nowMs,
|
||||
}},
|
||||
now,
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatalf("expected ApplyNftTrafficAccounting to fail")
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 0 {
|
||||
t.Fatalf("expected forward flow rollback, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 0 {
|
||||
t.Fatalf("expected user flow rollback, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT COALESCE((SELECT daily_used_bytes FROM user_quota WHERE user_id = 2), 0)`); got != 0 {
|
||||
t.Fatalf("expected quota rollback, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db"))
|
||||
if err != nil {
|
||||
@@ -106,7 +199,7 @@ func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T
|
||||
if !ok {
|
||||
t.Fatalf("expected metadata for forward with missing tunnel row")
|
||||
}
|
||||
if meta.ForwardID != 25 || meta.TunnelID != 99 || meta.TrafficRatio != 1 || meta.TunnelFlow != 1 {
|
||||
if meta.ForwardID != 25 || meta.UserID != 2 || meta.UserTunnelID != 0 || meta.TunnelID != 99 || meta.TrafficRatio != 1 || meta.TunnelFlow != 1 {
|
||||
t.Fatalf("unexpected fallback meta: %#v", meta)
|
||||
}
|
||||
}
|
||||
@@ -167,3 +260,19 @@ func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...inter
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func seedFlowBatchRows(t *testing.T, r *Repository, now int64) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 't1', 2.0, 1, 'tls', 3, ?, ?, 1, NULL, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(20, 2, 'u2', 'f20', 1, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -160,7 +160,7 @@ func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwardID, err := r.CreateForwardTx(1, "admin", "per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, []int64{3}, 24000, "", nil, 0, 5, int64(21), 0)
|
||||
forwardID, err := r.CreateForwardTx(1, "admin", "per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, []int64{3}, 24000, "", nil, 0, 5, int64(21), 0, 0, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateForwardTx: %v", err)
|
||||
}
|
||||
@@ -175,7 +175,7 @@ func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
|
||||
t.Fatalf("expected created ipSpeedId 21, got %+v", record.IPSpeedID)
|
||||
}
|
||||
|
||||
if err := r.UpdateForward(forwardID, "per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 9, int64(22), 0); err != nil {
|
||||
if err := r.UpdateForward(forwardID, "per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 9, int64(22), 0, 0, 0); err != nil {
|
||||
t.Fatalf("UpdateForward: %v", err)
|
||||
}
|
||||
record, err = r.GetForwardRecord(forwardID)
|
||||
@@ -216,6 +216,83 @@ func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardRepositoryPersistsProxyProtocolReceiveAndSend(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwardID, err := r.CreateForwardTx(1, "admin", "proxy-protocol-forward", 2, "1.1.1.1:443", "fifo", now, 1, []int64{3}, 24000, "", nil, 0, 0, nil, 0, 1, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateForwardTx: %v", err)
|
||||
}
|
||||
|
||||
record, err := r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetForwardRecord after create: %v", err)
|
||||
}
|
||||
if record.ProxyProtocolReceive != 1 || record.ProxyProtocolSend != 2 {
|
||||
t.Fatalf("expected proxyProtocol receive/send 1/2 after create, got %d/%d", record.ProxyProtocolReceive, record.ProxyProtocolSend)
|
||||
}
|
||||
|
||||
if err := r.UpdateForward(forwardID, "proxy-protocol-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 0, nil, 0, 2, 1); err != nil {
|
||||
t.Fatalf("UpdateForward: %v", err)
|
||||
}
|
||||
record, err = r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetForwardRecord after update: %v", err)
|
||||
}
|
||||
if record.ProxyProtocolReceive != 2 || record.ProxyProtocolSend != 1 {
|
||||
t.Fatalf("expected proxyProtocol receive/send 2/1 after update, got %d/%d", record.ProxyProtocolReceive, record.ProxyProtocolSend)
|
||||
}
|
||||
|
||||
records, err := r.ListForwardsByTunnel(2)
|
||||
if err != nil {
|
||||
t.Fatalf("ListForwardsByTunnel: %v", err)
|
||||
}
|
||||
if len(records) != 1 {
|
||||
t.Fatalf("expected 1 listed record, got %d", len(records))
|
||||
}
|
||||
if records[0].ProxyProtocolReceive != 2 || records[0].ProxyProtocolSend != 1 {
|
||||
t.Fatalf("expected listed proxyProtocol receive/send 2/1, got %d/%d", records[0].ProxyProtocolReceive, records[0].ProxyProtocolSend)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardRepositoryMapsLegacyProxyProtocolToSend(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Create(&model.Forward{
|
||||
UserID: 2,
|
||||
UserName: "user",
|
||||
Name: "legacy-proxy-protocol-forward",
|
||||
TunnelID: 9,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
ProxyProtocol: 2,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create legacy forward: %v", err)
|
||||
}
|
||||
forwardID := mustRepoLastInsertID(t, r)
|
||||
|
||||
record, err := r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetForwardRecord: %v", err)
|
||||
}
|
||||
if record.ProxyProtocolReceive != 0 || record.ProxyProtocolSend != 2 {
|
||||
t.Fatalf("expected legacy proxyProtocol to map to receive/send 0/2, got %d/%d", record.ProxyProtocolReceive, record.ProxyProtocolSend)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackForwardFieldsRestoresPerIPLimits(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
@@ -224,15 +301,15 @@ func TestRollbackForwardFieldsRestoresPerIPLimits(t *testing.T) {
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwardID, err := r.CreateForwardTx(1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, nil, 0, "", nil, 7, 5, int64(21), 2)
|
||||
forwardID, err := r.CreateForwardTx(1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, nil, 0, "", nil, 7, 5, int64(21), 2, 0, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateForwardTx: %v", err)
|
||||
}
|
||||
if err := r.UpdateForward(forwardID, "rollback-per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 0, nil, 0); err != nil {
|
||||
if err := r.UpdateForward(forwardID, "rollback-per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 0, nil, 0, 0, 0); err != nil {
|
||||
t.Fatalf("UpdateForward: %v", err)
|
||||
}
|
||||
|
||||
r.RollbackForwardFields(forwardID, 1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", 1, nil, 7, 5, int64(21), 2, now+2)
|
||||
r.RollbackForwardFields(forwardID, 1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", 1, nil, 7, 5, int64(21), 2, 0, 2, now+2)
|
||||
|
||||
record, err := r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
|
||||
@@ -117,6 +117,68 @@ func TestOpenBackfillsSQLiteLegacyTunnelProbeTargetColumns(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenBackfillsSQLiteLegacyForwardColumns(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "legacy-forward.db")
|
||||
db, err := gorm.Open(gsqlite.Open(dbPath), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open legacy sqlite: %v", err)
|
||||
}
|
||||
|
||||
if err := db.Exec(`
|
||||
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
|
||||
)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create legacy forward table: %v", err)
|
||||
}
|
||||
if err := 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, speed_id)
|
||||
VALUES(1, 2, 'legacy-user', 'legacy-forward', 3, '127.0.0.1:9000', 'fifo', 0, 0, 1, 1, 1, 0, NULL)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert legacy forward: %v", err)
|
||||
}
|
||||
if sqlDB, _ := db.DB(); sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open migrated sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
m := r.DB().Migrator()
|
||||
for _, field := range []string{"MaxConn", "IPMaxConn", "IPSpeedID", "ProxyProtocol"} {
|
||||
if !m.HasColumn(&model.Forward{}, field) {
|
||||
t.Fatalf("expected forward.%s column to exist", field)
|
||||
}
|
||||
}
|
||||
|
||||
var maxConn, ipMaxConn, proxyProtocol int
|
||||
var ipSpeedID sql.NullInt64
|
||||
if err := r.DB().Raw(`SELECT max_conn, ip_max_conn, ip_speed_id, proxy_protocol FROM forward WHERE id = 1`).Row().Scan(&maxConn, &ipMaxConn, &ipSpeedID, &proxyProtocol); err != nil {
|
||||
t.Fatalf("query forward defaults: %v", err)
|
||||
}
|
||||
if maxConn != 0 || ipMaxConn != 0 || ipSpeedID.Valid || proxyProtocol != 0 {
|
||||
t.Fatalf("expected default forward columns 0/0/NULL/0, got max_conn=%d ip_max_conn=%d ip_speed_id=%+v proxy_protocol=%d", maxConn, ipMaxConn, ipSpeedID, proxyProtocol)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaRunsPostgresIDRepairEvenAtCurrentVersion(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
|
||||
@@ -132,3 +132,31 @@ func TestUpsertTunnelMetricBucketsIsSafeUnderConcurrency(t *testing.T) {
|
||||
t.Fatalf("expected bytesOut %d, got %d", wantOut, rows[0].BytesOut)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetLatestTunnelQualitiesIncludesChainDetails(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
if err := r.InsertTunnelQuality(&model.TunnelQuality{
|
||||
TunnelID: 7,
|
||||
Timestamp: time.Now().UnixMilli(),
|
||||
Success: 1,
|
||||
ChainDetails: `{"primaryPath":[],"candidateHops":[{"fromNodeId":10,"toNodeId":31}]}`,
|
||||
}); err != nil {
|
||||
t.Fatalf("insert tunnel quality: %v", err)
|
||||
}
|
||||
|
||||
items, err := r.GetLatestTunnelQualities()
|
||||
if err != nil {
|
||||
t.Fatalf("get latest tunnel qualities: %v", err)
|
||||
}
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("expected one latest tunnel quality, got %+v", items)
|
||||
}
|
||||
if items[0].ChainDetails == "" {
|
||||
t.Fatalf("expected chain details in latest quality row, got %+v", items[0])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,19 +42,20 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
user := model.User{
|
||||
User: username,
|
||||
Pwd: pwdHash,
|
||||
RoleID: roleID,
|
||||
ExpTime: expTime,
|
||||
Flow: flow,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
Num: num,
|
||||
MaxConn: maxConn,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: status,
|
||||
User: username,
|
||||
Pwd: pwdHash,
|
||||
RoleID: roleID,
|
||||
ExpTime: expTime,
|
||||
Flow: flow,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
Num: num,
|
||||
MaxConn: maxConn,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: status,
|
||||
PasswordChangedAt: now,
|
||||
}
|
||||
if err := r.db.Create(&user).Error; err != nil {
|
||||
return 0, err
|
||||
@@ -81,15 +82,16 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
|
||||
return r.db.Model(&model.User{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"user": username,
|
||||
"pwd": pwdHash,
|
||||
"flow": flow,
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
"status": status,
|
||||
"max_conn": maxConn,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
"user": username,
|
||||
"pwd": pwdHash,
|
||||
"flow": flow,
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
"status": status,
|
||||
"max_conn": maxConn,
|
||||
"password_changed_at": now,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -111,6 +113,19 @@ func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow i
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserPassword(userID int64, pwdHash string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.User{}).
|
||||
Where("id = ?", userID).
|
||||
Updates(map[string]interface{}{
|
||||
"pwd": pwdHash,
|
||||
"password_changed_at": now,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
@@ -182,6 +197,19 @@ func (r *Repository) ResetUserFlowByUserTunnel(userTunnelID int64) {
|
||||
Updates(map[string]interface{}{"in_flow": 0, "out_flow": 0}).Error
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func (r *Repository) GetUsernameByID(userID int64) string {
|
||||
if r == nil || r.db == nil {
|
||||
return ""
|
||||
@@ -205,7 +233,7 @@ func (r *Repository) GetUserDefaultsForTunnel(userID int64) (flow int64, num int
|
||||
return user.Flow, user.Num, user.ExpTime, user.FlowResetTime, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}) error {
|
||||
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}, forwardMode string) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -232,6 +260,7 @@ func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serve
|
||||
UDPListenAddr: udpAddr,
|
||||
Inx: inx,
|
||||
IsRemote: isRemote,
|
||||
ForwardMode: defaultNodeForwardMode(forwardMode),
|
||||
RemoteURL: nullStringFromInterface(remoteURL),
|
||||
RemoteToken: nullStringFromInterface(remoteToken),
|
||||
RemoteConfig: nullStringFromInterface(remoteConfig),
|
||||
@@ -251,31 +280,44 @@ func (r *Repository) GetNodeStatusFields(nodeID int64) (status, httpFlag, tlsFla
|
||||
return node.Status, node.HTTP, node.TLS, node.Socks, nil
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
|
||||
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, expiryTime, renewalCycle interface{}, forwardMode string, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
updates := map[string]interface{}{
|
||||
"name": name,
|
||||
"remark": nullStringFromInterface(remark),
|
||||
"expiry_time": nullInt64FromInterface(expiryTime),
|
||||
"renewal_cycle": nullStringFromInterface(renewalCycle),
|
||||
"server_ip": serverIP,
|
||||
"server_ip_v4": nullStringFromInterface(serverIPV4),
|
||||
"server_ip_v6": nullStringFromInterface(serverIPV6),
|
||||
"extra_ips": nullStringFromInterface(extraIPs),
|
||||
"port": stringFromInterface(port),
|
||||
"interface_name": nullStringFromInterface(interfaceName),
|
||||
"http": httpFlag,
|
||||
"tls": tlsFlag,
|
||||
"socks": socksFlag,
|
||||
"tcp_listen_addr": tcpAddr,
|
||||
"udp_listen_addr": udpAddr,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
"expiry_reminder_dismissed": 0,
|
||||
}
|
||||
if strings.TrimSpace(forwardMode) != "" {
|
||||
updates["forward_mode"] = defaultNodeForwardMode(forwardMode)
|
||||
}
|
||||
return r.db.Model(&model.Node{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"remark": nullStringFromInterface(remark),
|
||||
"expiry_time": nullInt64FromInterface(expiryTime),
|
||||
"renewal_cycle": nullStringFromInterface(renewalCycle),
|
||||
"server_ip": serverIP,
|
||||
"server_ip_v4": nullStringFromInterface(serverIPV4),
|
||||
"server_ip_v6": nullStringFromInterface(serverIPV6),
|
||||
"extra_ips": nullStringFromInterface(extraIPs),
|
||||
"port": stringFromInterface(port),
|
||||
"interface_name": nullStringFromInterface(interfaceName),
|
||||
"http": httpFlag,
|
||||
"tls": tlsFlag,
|
||||
"socks": socksFlag,
|
||||
"tcp_listen_addr": tcpAddr,
|
||||
"udp_listen_addr": udpAddr,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
"expiry_reminder_dismissed": 0,
|
||||
}).Error
|
||||
Updates(updates).Error
|
||||
}
|
||||
|
||||
func defaultNodeForwardMode(mode string) string {
|
||||
switch strings.TrimSpace(strings.ToLower(mode)) {
|
||||
case "nftables":
|
||||
return "nftables"
|
||||
default:
|
||||
return "agent"
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeSecret(nodeID int64) (string, error) {
|
||||
@@ -697,23 +739,26 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
|
||||
return p
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) error {
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, proxyProtocolReceive int, proxyProtocolSend int) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
_, proxyProtocolSend = normalizeForwardProxyProtocol(proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
|
||||
return r.db.Model(&model.Forward{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"updated_time": now,
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"proxy_protocol_receive": proxyProtocolReceive,
|
||||
"proxy_protocol_send": proxyProtocolSend,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -740,6 +785,9 @@ func (r *Repository) DeleteForwardCascade(forwardID int64) error {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.NftCounterState{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.ForwardPort{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -787,26 +835,29 @@ func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int,
|
||||
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, now int64) {
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, proxyProtocolReceive int, proxyProtocolSend int, now int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
}
|
||||
_, proxyProtocolSend = normalizeForwardProxyProtocol(proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
|
||||
_ = r.db.Model(&model.Forward{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"user_id": userID,
|
||||
"user_name": userName,
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"status": status,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"updated_time": now,
|
||||
"user_id": userID,
|
||||
"user_name": userName,
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"status": status,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"proxy_protocol_receive": proxyProtocolReceive,
|
||||
"proxy_protocol_send": proxyProtocolSend,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -1266,30 +1317,33 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
return ut.ID, true, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) (int64, error) {
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, proxyProtocolReceive int, proxyProtocolSend int) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
_, proxyProtocolSend = normalizeForwardProxyProtocol(proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
|
||||
var forwardID int64
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
fwd := model.Forward{
|
||||
UserID: userID,
|
||||
UserName: userName,
|
||||
Name: name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: remoteAddr,
|
||||
Strategy: strategy,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
MaxConn: maxConn,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
IPMaxConn: ipMaxConn,
|
||||
IPSpeedID: nullInt64FromInterface(ipSpeedID),
|
||||
ProxyProtocol: proxyProtocol,
|
||||
UserID: userID,
|
||||
UserName: userName,
|
||||
Name: name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: remoteAddr,
|
||||
Strategy: strategy,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
MaxConn: maxConn,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
IPMaxConn: ipMaxConn,
|
||||
IPSpeedID: nullInt64FromInterface(ipSpeedID),
|
||||
ProxyProtocol: proxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
}
|
||||
if err := tx.Create(&fwd).Error; err != nil {
|
||||
return err
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const (
|
||||
nftCounterProtocolTCP = "tcp"
|
||||
nftCounterProtocolUDP = "udp"
|
||||
|
||||
nftCounterDirectionToTarget = "to-target"
|
||||
nftCounterDirectionFromTarget = "from-target"
|
||||
)
|
||||
|
||||
type NftCounterStateInput struct {
|
||||
NodeID int64
|
||||
ForwardID int64
|
||||
Protocol string
|
||||
Direction string
|
||||
RuleHash string
|
||||
Bytes uint64
|
||||
Packets uint64
|
||||
CollectedTime int64
|
||||
}
|
||||
|
||||
type NftablesCollectionNode struct {
|
||||
NodeID int64
|
||||
Config model.NodeSSHConfig
|
||||
}
|
||||
|
||||
func (r *Repository) ListNftablesNodesForCollection() ([]NftablesCollectionNode, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
type collectionRow struct {
|
||||
NodeID int64 `gorm:"column:node_id"`
|
||||
ConfigID int64 `gorm:"column:config_id"`
|
||||
Host string `gorm:"column:host"`
|
||||
Port int `gorm:"column:port"`
|
||||
Username string `gorm:"column:username"`
|
||||
AuthType string `gorm:"column:auth_type"`
|
||||
Password string `gorm:"column:password"`
|
||||
PrivateKey string `gorm:"column:private_key"`
|
||||
Passphrase string `gorm:"column:passphrase"`
|
||||
SudoMode string `gorm:"column:sudo_mode"`
|
||||
CreatedTime int64 `gorm:"column:created_time"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time"`
|
||||
}
|
||||
|
||||
var rows []collectionRow
|
||||
if err := r.db.Table("node").
|
||||
Select("node.id AS node_id, node_ssh_config.id AS config_id, node_ssh_config.host, node_ssh_config.port, node_ssh_config.username, node_ssh_config.auth_type, node_ssh_config.password, node_ssh_config.private_key, node_ssh_config.passphrase, node_ssh_config.sudo_mode, node_ssh_config.created_time, node_ssh_config.updated_time").
|
||||
Joins("JOIN node_ssh_config ON node_ssh_config.node_id = node.id").
|
||||
Where("node.status = ? AND LOWER(TRIM(node.forward_mode)) = ?", 1, "nftables").
|
||||
Order("node.id ASC").
|
||||
Scan(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nodes := make([]NftablesCollectionNode, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
nodes = append(nodes, NftablesCollectionNode{
|
||||
NodeID: row.NodeID,
|
||||
Config: model.NodeSSHConfig{
|
||||
ID: row.ConfigID,
|
||||
NodeID: row.NodeID,
|
||||
Host: row.Host,
|
||||
Port: row.Port,
|
||||
Username: row.Username,
|
||||
AuthType: row.AuthType,
|
||||
Password: nullStringFromInterface(row.Password),
|
||||
PrivateKey: nullStringFromInterface(row.PrivateKey),
|
||||
Passphrase: nullStringFromInterface(row.Passphrase),
|
||||
SudoMode: row.SudoMode,
|
||||
CreatedTime: row.CreatedTime,
|
||||
UpdatedTime: row.UpdatedTime,
|
||||
},
|
||||
})
|
||||
}
|
||||
return nodes, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetNftCounterStatesByNode(nodeID int64) ([]model.NftCounterState, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var rows []model.NftCounterState
|
||||
err := r.db.Where("node_id = ?", nodeID).
|
||||
Order("forward_id ASC, protocol ASC, direction ASC").
|
||||
Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertNftCounterStates(inputs []NftCounterStateInput, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if len(inputs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
return upsertNftCounterStatesTx(tx, inputs, now)
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) ApplyNftTrafficAccounting(deltas []FlowUploadCounterDelta, quotaUsage map[int64]int64, states []NftCounterStateInput, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
quotaViews := map[int64]*model.UserQuotaView{}
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := applyFlowUploadDeltasTx(tx, deltas); err != nil {
|
||||
return err
|
||||
}
|
||||
var err error
|
||||
quotaViews, err = r.addUserQuotaUsageBatchTx(tx, quotaUsage, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return upsertNftCounterStatesTx(tx, states, now.UnixMilli())
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return quotaViews, nil
|
||||
}
|
||||
|
||||
func upsertNftCounterStatesTx(tx *gorm.DB, inputs []NftCounterStateInput, now int64) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
for _, input := range inputs {
|
||||
row, ok := nftCounterStateFromInput(input, now)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{
|
||||
{Name: "node_id"},
|
||||
{Name: "forward_id"},
|
||||
{Name: "protocol"},
|
||||
{Name: "direction"},
|
||||
},
|
||||
DoUpdates: clause.Assignments(map[string]interface{}{
|
||||
"rule_hash": row.RuleHash,
|
||||
"bytes": row.Bytes,
|
||||
"packets": row.Packets,
|
||||
"collected_time": row.CollectedTime,
|
||||
"updated_time": row.UpdatedTime,
|
||||
}),
|
||||
}).Create(&row).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNftCounterStatesByForward(forwardID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Where("forward_id = ?", forwardID).Delete(&model.NftCounterState{}).Error
|
||||
}
|
||||
|
||||
func nftCounterStateFromInput(input NftCounterStateInput, now int64) (model.NftCounterState, bool) {
|
||||
protocol := strings.ToLower(strings.TrimSpace(input.Protocol))
|
||||
direction := strings.ToLower(strings.TrimSpace(input.Direction))
|
||||
if input.NodeID <= 0 || input.ForwardID <= 0 || !isValidNftCounterProtocol(protocol) || !isValidNftCounterDirection(direction) {
|
||||
return model.NftCounterState{}, false
|
||||
}
|
||||
if input.Bytes > uint64(math.MaxInt64) || input.Packets > uint64(math.MaxInt64) {
|
||||
return model.NftCounterState{}, false
|
||||
}
|
||||
return model.NftCounterState{
|
||||
NodeID: input.NodeID,
|
||||
ForwardID: input.ForwardID,
|
||||
Protocol: protocol,
|
||||
Direction: direction,
|
||||
RuleHash: strings.TrimSpace(input.RuleHash),
|
||||
Bytes: int64(input.Bytes),
|
||||
Packets: int64(input.Packets),
|
||||
CollectedTime: input.CollectedTime,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}, true
|
||||
}
|
||||
|
||||
func isValidNftCounterProtocol(protocol string) bool {
|
||||
return protocol == nftCounterProtocolTCP || protocol == nftCounterProtocolUDP
|
||||
}
|
||||
|
||||
func isValidNftCounterDirection(direction string) bool {
|
||||
return direction == nftCounterDirectionToTarget || direction == nftCounterDirectionFromTarget
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"math"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestNftCounterStateUpsertUpdatesExistingKey(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
first := []NftCounterStateInput{
|
||||
{
|
||||
NodeID: 11,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 100,
|
||||
Packets: 10,
|
||||
CollectedTime: 1000,
|
||||
},
|
||||
{
|
||||
NodeID: 0,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
Bytes: 999,
|
||||
},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(first, 2000); err != nil {
|
||||
t.Fatalf("first UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
second := []NftCounterStateInput{
|
||||
{
|
||||
NodeID: 11,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-b",
|
||||
Bytes: 250,
|
||||
Packets: 25,
|
||||
CollectedTime: 3000,
|
||||
},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(second, 4000); err != nil {
|
||||
t.Fatalf("second UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected one counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
got := rows[0]
|
||||
if got.ForwardID != 42 || got.Protocol != "tcp" || got.Direction != "to-target" {
|
||||
t.Fatalf("unexpected counter state key: %+v", got)
|
||||
}
|
||||
if got.RuleHash != "hash-b" || got.Bytes != 250 || got.Packets != 25 || got.CollectedTime != 3000 {
|
||||
t.Fatalf("counter state was not updated: %+v", got)
|
||||
}
|
||||
if got.CreatedTime != 2000 || got.UpdatedTime != 4000 {
|
||||
t.Fatalf("unexpected timestamps after upsert: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateDeleteByForwardRemovesOnlyMatchingRows(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "a", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "b", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||
{NodeID: 12, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "c", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
if err := r.DeleteNftCounterStatesByForward(42); err != nil {
|
||||
t.Fatalf("DeleteNftCounterStatesByForward: %v", err)
|
||||
}
|
||||
|
||||
node11, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode(11): %v", err)
|
||||
}
|
||||
if len(node11) != 1 || node11[0].ForwardID != 43 {
|
||||
t.Fatalf("expected only forward 43 for node 11, got %+v", node11)
|
||||
}
|
||||
node12, err := r.GetNftCounterStatesByNode(12)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode(12): %v", err)
|
||||
}
|
||||
if len(node12) != 0 {
|
||||
t.Fatalf("expected forward 42 state removed from node 12, got %+v", node12)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteForwardCascadeRemovesNftCounterStateOnlyForDeletedForward(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwards := []model.Forward{
|
||||
{ID: 42, UserID: 1, UserName: "admin", Name: "forward-a", TunnelID: 10, RemoteAddr: "203.0.113.1:80", Strategy: "fifo", CreatedTime: now, UpdatedTime: now, Status: 1},
|
||||
{ID: 43, UserID: 1, UserName: "admin", Name: "forward-b", TunnelID: 10, RemoteAddr: "203.0.113.2:80", Strategy: "fifo", CreatedTime: now, UpdatedTime: now, Status: 1},
|
||||
}
|
||||
if err := r.DB().Create(&forwards).Error; err != nil {
|
||||
t.Fatalf("seed forwards: %v", err)
|
||||
}
|
||||
if err := r.UpsertNftCounterStates([]NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "a", Bytes: 100, Packets: 10, CollectedTime: now},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "b", Bytes: 200, Packets: 20, CollectedTime: now},
|
||||
}, now); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DeleteForwardCascade(42); err != nil {
|
||||
t.Fatalf("DeleteForwardCascade: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 || rows[0].ForwardID != 43 {
|
||||
t.Fatalf("expected only forward 43 counter state to remain, got %+v", rows)
|
||||
}
|
||||
var deletedForwardCount int64
|
||||
if err := r.DB().Model(&model.Forward{}).Where("id = ?", int64(42)).Count(&deletedForwardCount).Error; err != nil {
|
||||
t.Fatalf("count deleted forward: %v", err)
|
||||
}
|
||||
if deletedForwardCount != 0 {
|
||||
t.Fatalf("expected forward 42 deleted, count=%d", deletedForwardCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateUpsertSkipsInvalidProtocolAndDirection(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "icmp", Direction: "to-target", RuleHash: "bad-protocol", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "tcp", Direction: "sideways", RuleHash: "bad-direction", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 44, Protocol: " UDP ", Direction: " FROM-TARGET ", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
if rows[0].ForwardID != 44 || rows[0].Protocol != "udp" || rows[0].Direction != "from-target" {
|
||||
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateUpsertSkipsCountersAboveInt64(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "too-large", Bytes: uint64(math.MaxInt64) + 1, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
if rows[0].ForwardID != 43 || rows[0].Bytes != 300 || rows[0].Packets != 30 {
|
||||
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestListNftablesNodesForCollectionReturnsActiveNftablesWithSSHOrdered(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-collection.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
seedCollectionNode(t, r, 1, "agent", 1, now)
|
||||
seedCollectionNode(t, r, 2, " nftables ", 1, now)
|
||||
seedCollectionNode(t, r, 3, "NFTABLES", 0, now)
|
||||
seedCollectionNode(t, r, 4, "nftables", 1, now)
|
||||
seedCollectionNode(t, r, 5, "nftables", 1, now)
|
||||
|
||||
if err := r.UpsertNodeSSHConfig(4, NftSSHConfigInput{
|
||||
Host: "203.0.113.4",
|
||||
Port: 2222,
|
||||
Username: "root",
|
||||
AuthType: "password",
|
||||
Password: "secret-4",
|
||||
SudoMode: "none",
|
||||
}, now); err != nil {
|
||||
t.Fatalf("upsert ssh config 4: %v", err)
|
||||
}
|
||||
if err := r.UpsertNodeSSHConfig(2, NftSSHConfigInput{
|
||||
Host: "203.0.113.2",
|
||||
Port: 22,
|
||||
Username: "admin",
|
||||
AuthType: "private_key",
|
||||
SudoMode: "sudo",
|
||||
}, now); err != nil {
|
||||
t.Fatalf("upsert ssh config 2: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNftablesNodesForCollection()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftablesNodesForCollection: %v", err)
|
||||
}
|
||||
if len(nodes) != 2 {
|
||||
t.Fatalf("expected 2 collection nodes, got %d: %+v", len(nodes), nodes)
|
||||
}
|
||||
if nodes[0].NodeID != 2 || nodes[1].NodeID != 4 {
|
||||
t.Fatalf("expected nodes ordered by id [2 4], got [%d %d]", nodes[0].NodeID, nodes[1].NodeID)
|
||||
}
|
||||
if nodes[0].Config.NodeID != 2 || nodes[0].Config.Host != "203.0.113.2" || nodes[0].Config.Username != "admin" {
|
||||
t.Fatalf("unexpected first config: %+v", nodes[0].Config)
|
||||
}
|
||||
if nodes[1].Config.NodeID != 4 || nodes[1].Config.Port != 2222 || nodes[1].Config.Password.String != "secret-4" {
|
||||
t.Fatalf("unexpected second config: %+v", nodes[1].Config)
|
||||
}
|
||||
}
|
||||
|
||||
func seedCollectionNode(t *testing.T, r *Repository, id int64, forwardMode string, status int, now int64) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, port, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, forward_mode)
|
||||
VALUES(?, ?, 'secret', ?, '1000-2000', ?, ?, ?, '[::]', '[::]', 0, ?)
|
||||
`, id, "node", "198.51.100.1", now, now, status, forwardMode).Error; err != nil {
|
||||
t.Fatalf("insert node %d: %v", id, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,249 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
type NftSSHConfigInput struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
AuthType string
|
||||
Password string
|
||||
PrivateKey string
|
||||
Passphrase string
|
||||
SudoMode string
|
||||
}
|
||||
|
||||
type NftRuleBindingInput struct {
|
||||
ForwardID int64
|
||||
NodeID int64
|
||||
InPort int
|
||||
Protocols string
|
||||
TargetAddr string
|
||||
BindIP string
|
||||
RuleHash string
|
||||
Status string
|
||||
LastError string
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertNodeSSHConfig(nodeID int64, cfg NftSSHConfigInput, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if nodeID <= 0 {
|
||||
return errors.New("node id is required")
|
||||
}
|
||||
port := cfg.Port
|
||||
if port <= 0 {
|
||||
port = 22
|
||||
}
|
||||
authType := strings.TrimSpace(strings.ToLower(cfg.AuthType))
|
||||
if authType == "" {
|
||||
authType = "private_key"
|
||||
}
|
||||
sudoMode := strings.TrimSpace(strings.ToLower(cfg.SudoMode))
|
||||
if sudoMode == "" {
|
||||
sudoMode = "none"
|
||||
}
|
||||
row := model.NodeSSHConfig{
|
||||
NodeID: nodeID,
|
||||
Host: strings.TrimSpace(cfg.Host),
|
||||
Port: port,
|
||||
Username: strings.TrimSpace(cfg.Username),
|
||||
AuthType: authType,
|
||||
Password: nullStringFromInterface(cfg.Password),
|
||||
PrivateKey: nullStringFromInterface(cfg.PrivateKey),
|
||||
Passphrase: nullStringFromInterface(cfg.Passphrase),
|
||||
SudoMode: sudoMode,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
return r.db.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "node_id"}},
|
||||
DoUpdates: clause.Assignments(map[string]interface{}{
|
||||
"host": row.Host,
|
||||
"port": row.Port,
|
||||
"username": row.Username,
|
||||
"auth_type": row.AuthType,
|
||||
"password": row.Password,
|
||||
"private_key": row.PrivateKey,
|
||||
"passphrase": row.Passphrase,
|
||||
"sudo_mode": row.SudoMode,
|
||||
"updated_time": row.UpdatedTime,
|
||||
}),
|
||||
}).Create(&row).Error
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeSSHConfig(nodeID int64) (*model.NodeSSHConfig, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
return r.GetNodeSSHConfigTx(r.db, nodeID)
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeSSHConfigTx(tx *gorm.DB, nodeID int64) (*model.NodeSSHConfig, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
var cfg model.NodeSSHConfig
|
||||
if err := tx.Where("node_id = ?", nodeID).First(&cfg).Error; err != nil {
|
||||
return nil, normalizeNotFoundErr(err)
|
||||
}
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNodeSSHConfig(nodeID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Where("node_id = ?", nodeID).Delete(&model.NodeSSHConfig{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertNftRuleBinding(input NftRuleBindingInput, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
row := model.NftRuleBinding{
|
||||
ForwardID: input.ForwardID,
|
||||
NodeID: input.NodeID,
|
||||
InPort: input.InPort,
|
||||
Protocols: defaultString(strings.TrimSpace(input.Protocols), "tcp,udp"),
|
||||
TargetAddr: strings.TrimSpace(input.TargetAddr),
|
||||
BindIP: strings.TrimSpace(input.BindIP),
|
||||
RuleHash: strings.TrimSpace(input.RuleHash),
|
||||
Status: defaultString(strings.TrimSpace(input.Status), "pending"),
|
||||
LastError: strings.TrimSpace(input.LastError),
|
||||
AppliedTime: now,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
return r.db.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "forward_id"}, {Name: "node_id"}},
|
||||
DoUpdates: clause.Assignments(map[string]interface{}{
|
||||
"in_port": row.InPort,
|
||||
"protocols": row.Protocols,
|
||||
"target_addr": row.TargetAddr,
|
||||
"bind_ip": row.BindIP,
|
||||
"rule_hash": row.RuleHash,
|
||||
"status": row.Status,
|
||||
"last_error": row.LastError,
|
||||
"applied_time": row.AppliedTime,
|
||||
"updated_time": row.UpdatedTime,
|
||||
}),
|
||||
}).Create(&row).Error
|
||||
}
|
||||
|
||||
func (r *Repository) MarkNftRuleBindingError(forwardID, nodeID int64, message string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.NftRuleBinding{}).
|
||||
Where("forward_id = ? AND node_id = ?", forwardID, nodeID).
|
||||
Updates(map[string]interface{}{
|
||||
"status": "error",
|
||||
"last_error": strings.TrimSpace(message),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ListNftRuleBindingsByNode(nodeID int64) ([]model.NftRuleBinding, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var rows []model.NftRuleBinding
|
||||
err := r.db.Where("node_id = ?", nodeID).Order("forward_id ASC").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNftRuleBindingsByForward(forwardID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Where("forward_id = ?", forwardID).Delete(&model.NftRuleBinding{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeForwardMode(nodeID int64) (string, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return "", errors.New("repository not initialized")
|
||||
}
|
||||
return r.GetNodeForwardModeTx(r.db, nodeID)
|
||||
}
|
||||
|
||||
func (r *Repository) GetNodeForwardModeTx(tx *gorm.DB, nodeID int64) (string, error) {
|
||||
if tx == nil {
|
||||
return "", errors.New("database unavailable")
|
||||
}
|
||||
var row struct {
|
||||
ForwardMode sql.NullString `gorm:"column:forward_mode"`
|
||||
}
|
||||
err := tx.Model(&model.Node{}).Select("forward_mode").Where("id = ?", nodeID).First(&row).Error
|
||||
if err != nil {
|
||||
return "", normalizeNotFoundErr(err)
|
||||
}
|
||||
return defaultNodeForwardMode(row.ForwardMode.String), nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListActiveForwardsByNode(nodeID int64) ([]model.ForwardRecord, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var forwards []model.Forward
|
||||
err := r.db.Model(&model.Forward{}).
|
||||
Joins("JOIN forward_port ON forward_port.forward_id = forward.id").
|
||||
Where("forward_port.node_id = ? AND forward.status = 1", nodeID).
|
||||
Order("forward.id ASC").
|
||||
Distinct("forward.*").
|
||||
Find(&forwards).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
if strings.TrimSpace(rows[i].Strategy) == "" {
|
||||
rows[i].Strategy = "fifo"
|
||||
}
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func defaultString(value, fallback string) string {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fallback
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
|
||||
if send == 0 && legacy > 0 {
|
||||
send = legacy
|
||||
}
|
||||
return receive, send
|
||||
}
|
||||
@@ -0,0 +1,323 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestNftablesNodeModeSSHConfigAndBindingPersistence(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreateNode(
|
||||
"nft-node",
|
||||
"secret",
|
||||
"203.0.113.10",
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
now,
|
||||
1,
|
||||
"[::]",
|
||||
"[::]",
|
||||
1,
|
||||
0,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"nftables",
|
||||
); err != nil {
|
||||
t.Fatalf("CreateNode: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes: %v", err)
|
||||
}
|
||||
if len(nodes) != 1 {
|
||||
t.Fatalf("expected 1 node, got %d", len(nodes))
|
||||
}
|
||||
nodeID := nodes[0]["id"].(int64)
|
||||
if got := nodes[0]["forwardMode"]; got != "nftables" {
|
||||
t.Fatalf("expected forwardMode nftables, got %#v", got)
|
||||
}
|
||||
|
||||
cfg := NftSSHConfigInput{
|
||||
Host: "203.0.113.10",
|
||||
Port: 22,
|
||||
Username: "root",
|
||||
AuthType: "private_key",
|
||||
PrivateKey: "encrypted-private-key",
|
||||
SudoMode: "none",
|
||||
}
|
||||
if err := r.UpsertNodeSSHConfig(nodeID, cfg, now); err != nil {
|
||||
t.Fatalf("UpsertNodeSSHConfig: %v", err)
|
||||
}
|
||||
loaded, err := r.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNodeSSHConfig: %v", err)
|
||||
}
|
||||
if loaded.Host != cfg.Host || loaded.Port != cfg.Port || loaded.Username != cfg.Username || loaded.AuthType != cfg.AuthType {
|
||||
t.Fatalf("unexpected ssh config: %+v", loaded)
|
||||
}
|
||||
|
||||
binding := NftRuleBindingInput{
|
||||
ForwardID: 42,
|
||||
NodeID: nodeID,
|
||||
InPort: 24000,
|
||||
Protocols: "tcp,udp",
|
||||
TargetAddr: "198.51.100.20:443",
|
||||
BindIP: "",
|
||||
RuleHash: "hash-a",
|
||||
Status: "applied",
|
||||
LastError: "",
|
||||
}
|
||||
if err := r.UpsertNftRuleBinding(binding, now); err != nil {
|
||||
t.Fatalf("UpsertNftRuleBinding: %v", err)
|
||||
}
|
||||
bindings, err := r.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftRuleBindingsByNode: %v", err)
|
||||
}
|
||||
if len(bindings) != 1 {
|
||||
t.Fatalf("expected 1 binding, got %d", len(bindings))
|
||||
}
|
||||
if bindings[0].ForwardID != 42 || bindings[0].RuleHash != "hash-a" || bindings[0].Status != "applied" {
|
||||
t.Fatalf("unexpected binding: %+v", bindings[0])
|
||||
}
|
||||
|
||||
if err := r.MarkNftRuleBindingError(42, nodeID, "nft failed", now+1); err != nil {
|
||||
t.Fatalf("MarkNftRuleBindingError: %v", err)
|
||||
}
|
||||
bindings, err = r.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftRuleBindingsByNode after error: %v", err)
|
||||
}
|
||||
if bindings[0].Status != "error" || !strings.Contains(bindings[0].LastError, "nft failed") {
|
||||
t.Fatalf("expected error binding, got %+v", bindings[0])
|
||||
}
|
||||
|
||||
if err := r.DeleteNftRuleBindingsByForward(42); err != nil {
|
||||
t.Fatalf("DeleteNftRuleBindingsByForward: %v", err)
|
||||
}
|
||||
bindings, err = r.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftRuleBindingsByNode after delete: %v", err)
|
||||
}
|
||||
if len(bindings) != 0 {
|
||||
t.Fatalf("expected no bindings after delete, got %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftablesNodeSSHConfigSurvivesRepositoryReopen(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg NftSSHConfigInput
|
||||
}{
|
||||
{
|
||||
name: "password",
|
||||
cfg: NftSSHConfigInput{
|
||||
Host: "203.0.113.10",
|
||||
Port: 2222,
|
||||
Username: "root",
|
||||
AuthType: "password",
|
||||
Password: "ssh-password",
|
||||
Passphrase: "key-passphrase",
|
||||
SudoMode: "sudo",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "private-key",
|
||||
cfg: NftSSHConfigInput{
|
||||
Host: "203.0.113.11",
|
||||
Port: 2223,
|
||||
Username: "admin",
|
||||
AuthType: "private_key",
|
||||
PrivateKey: "PRIVATE-KEY-SHOULD-PERSIST",
|
||||
Passphrase: "key-passphrase",
|
||||
SudoMode: "sudo_su",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "nftables-ssh.sqlite")
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreateNode(
|
||||
"nft-node",
|
||||
"secret",
|
||||
tt.cfg.Host,
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
now,
|
||||
1,
|
||||
"[::]",
|
||||
"[::]",
|
||||
1,
|
||||
0,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"nftables",
|
||||
); err != nil {
|
||||
t.Fatalf("CreateNode: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes: %v", err)
|
||||
}
|
||||
nodeID := nodes[0]["id"].(int64)
|
||||
if err := r.UpsertNodeSSHConfig(nodeID, tt.cfg, now); err != nil {
|
||||
t.Fatalf("UpsertNodeSSHConfig: %v", err)
|
||||
}
|
||||
if err := r.Close(); err != nil {
|
||||
t.Fatalf("close repo: %v", err)
|
||||
}
|
||||
|
||||
reopened, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("reopen repo: %v", err)
|
||||
}
|
||||
defer reopened.Close()
|
||||
|
||||
loaded, err := reopened.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNodeSSHConfig after reopen: %v", err)
|
||||
}
|
||||
if loaded.Host != tt.cfg.Host || loaded.Port != tt.cfg.Port || loaded.Username != tt.cfg.Username || loaded.AuthType != tt.cfg.AuthType || loaded.SudoMode != tt.cfg.SudoMode {
|
||||
t.Fatalf("unexpected ssh config after reopen: %+v", loaded)
|
||||
}
|
||||
if tt.cfg.Password != "" && (!loaded.Password.Valid || loaded.Password.String != tt.cfg.Password) {
|
||||
t.Fatalf("expected password to persist after reopen, got %+v", loaded.Password)
|
||||
}
|
||||
if tt.cfg.PrivateKey != "" && (!loaded.PrivateKey.Valid || loaded.PrivateKey.String != tt.cfg.PrivateKey) {
|
||||
t.Fatalf("expected private key to persist after reopen, got %+v", loaded.PrivateKey)
|
||||
}
|
||||
if tt.cfg.Passphrase != "" && (!loaded.Passphrase.Valid || loaded.Passphrase.String != tt.cfg.Passphrase) {
|
||||
t.Fatalf("expected passphrase to persist after reopen, got %+v", loaded.Passphrase)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateNodeWithoutForwardModePreservesExistingMode(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreateNode(
|
||||
"nft-node",
|
||||
"secret",
|
||||
"203.0.113.11",
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
now,
|
||||
1,
|
||||
"[::]",
|
||||
"[::]",
|
||||
1,
|
||||
0,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"nftables",
|
||||
); err != nil {
|
||||
t.Fatalf("CreateNode: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes: %v", err)
|
||||
}
|
||||
if len(nodes) != 1 {
|
||||
t.Fatalf("expected 1 node, got %d", len(nodes))
|
||||
}
|
||||
nodeID := nodes[0]["id"].(int64)
|
||||
|
||||
if err := r.UpdateNode(
|
||||
nodeID,
|
||||
"nft-node-updated",
|
||||
"203.0.113.11",
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"",
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
"[::]",
|
||||
"[::]",
|
||||
now+1,
|
||||
); err != nil {
|
||||
t.Fatalf("UpdateNode: %v", err)
|
||||
}
|
||||
|
||||
gotNode, err := r.GetNodeRecord(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNodeRecord: %v", err)
|
||||
}
|
||||
if gotNode == nil {
|
||||
t.Fatal("expected node record, got nil")
|
||||
}
|
||||
if gotNode.ForwardMode != "nftables" {
|
||||
t.Fatalf("expected mapped forward mode nftables, got %q", gotNode.ForwardMode)
|
||||
}
|
||||
|
||||
nodes, err = r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes after update: %v", err)
|
||||
}
|
||||
if got := nodes[0]["forwardMode"]; got != "nftables" {
|
||||
t.Fatalf("expected persisted forwardMode nftables after update, got %#v", got)
|
||||
}
|
||||
}
|
||||
@@ -44,7 +44,8 @@ func (r *Repository) GetLatestTunnelQualities() ([]model.TunnelQuality, error) {
|
||||
// Use window function (works on modern SQLite 3.25+ and PostgreSQL).
|
||||
q := `
|
||||
SELECT id, tunnel_id, entry_to_exit_latency, exit_to_bing_latency,
|
||||
entry_to_exit_loss, exit_to_bing_loss, success, error_message, timestamp
|
||||
entry_to_exit_loss, exit_to_bing_loss, success, error_message, timestamp,
|
||||
chain_details
|
||||
FROM (
|
||||
SELECT *, ROW_NUMBER() OVER (PARTITION BY tunnel_id ORDER BY timestamp DESC, id DESC) AS rn
|
||||
FROM tunnel_quality
|
||||
|
||||
@@ -264,39 +264,11 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
var result map[int64]*model.UserQuotaView
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
userIDs := make([]int64, 0, len(usages))
|
||||
for userID := range usages {
|
||||
if userID > 0 {
|
||||
userIDs = append(userIDs, userID)
|
||||
}
|
||||
}
|
||||
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
|
||||
|
||||
for _, userID := range userIDs {
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
if usages[userID] > 0 {
|
||||
q.DailyUsedBytes += usages[userID]
|
||||
q.MonthlyUsedBytes += usages[userID]
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
|
||||
}
|
||||
return nil
|
||||
var err error
|
||||
result, err = r.addUserQuotaUsageBatchTx(tx, usages, now)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -304,6 +276,48 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *Repository) addUserQuotaUsageBatchTx(tx *gorm.DB, usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
if len(usages) == 0 {
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
userIDs := make([]int64, 0, len(usages))
|
||||
for userID := range usages {
|
||||
if userID > 0 {
|
||||
userIDs = append(userIDs, userID)
|
||||
}
|
||||
}
|
||||
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
|
||||
|
||||
for _, userID := range userIDs {
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
if usages[userID] > 0 {
|
||||
q.DailyUsedBytes += usages[userID]
|
||||
q.MonthlyUsedBytes += usages[userID]
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
|
||||
@@ -42,6 +42,18 @@ type nodeSession struct {
|
||||
crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建
|
||||
}
|
||||
|
||||
type adminSession struct {
|
||||
userID int64
|
||||
claims auth.Claims
|
||||
conn *connWrap
|
||||
}
|
||||
|
||||
type monitorSession struct {
|
||||
userID int64
|
||||
claims auth.Claims
|
||||
conn *connWrap
|
||||
}
|
||||
|
||||
type commandResponse struct {
|
||||
Type string `json:"type"`
|
||||
Success bool `json:"success"`
|
||||
@@ -74,12 +86,14 @@ type Server struct {
|
||||
upgrader websocket.Upgrader
|
||||
onNodeOnline func(nodeID int64)
|
||||
onNodeMetric func(nodeID int64, info SystemInfo)
|
||||
getUserAuthState func(userID int64) (*auth.UserAuthState, error)
|
||||
|
||||
mu sync.RWMutex
|
||||
admins map[*connWrap]struct{}
|
||||
nodes map[int64]*nodeSession
|
||||
byConn map[*websocket.Conn]*nodeSession
|
||||
pending map[string]pendingRequest
|
||||
admins map[*adminSession]struct{}
|
||||
monitors map[*monitorSession]struct{}
|
||||
nodes map[int64]*nodeSession
|
||||
byConn map[*websocket.Conn]*nodeSession
|
||||
pending map[string]pendingRequest
|
||||
}
|
||||
|
||||
type SystemInfo struct {
|
||||
@@ -123,13 +137,23 @@ func NewServer(repo *repo.Repository, jwtSecret string) *Server {
|
||||
upgrader: websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool { return true },
|
||||
},
|
||||
admins: make(map[*connWrap]struct{}),
|
||||
nodes: make(map[int64]*nodeSession),
|
||||
byConn: make(map[*websocket.Conn]*nodeSession),
|
||||
pending: make(map[string]pendingRequest),
|
||||
admins: make(map[*adminSession]struct{}),
|
||||
monitors: make(map[*monitorSession]struct{}),
|
||||
nodes: make(map[int64]*nodeSession),
|
||||
byConn: make(map[*websocket.Conn]*nodeSession),
|
||||
pending: make(map[string]pendingRequest),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) SetUserAuthStateLookup(fn func(userID int64) (*auth.UserAuthState, error)) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.getUserAuthState = fn
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
query := r.URL.Query()
|
||||
typeVal := query.Get("type")
|
||||
@@ -146,18 +170,36 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if typeVal == "0" {
|
||||
if _, ok := auth.ValidateToken(secret, s.jwtSecret); !ok {
|
||||
claims, ok := auth.ValidateToken(secret, s.jwtSecret)
|
||||
if !ok {
|
||||
http.Error(w, "forbidden", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
s.handleAdmin(w, r)
|
||||
userID, err := strconv.ParseInt(claims.Sub, 10, 64)
|
||||
if err != nil {
|
||||
http.Error(w, "forbidden", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
if claims.RoleID == 0 {
|
||||
if !s.validateAdminSession(userID, claims) {
|
||||
http.Error(w, "forbidden", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
s.handleAdmin(w, r, userID, claims)
|
||||
return
|
||||
}
|
||||
if !s.validateMonitorSession(userID, claims) {
|
||||
http.Error(w, "forbidden", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
s.handleMonitor(w, r, userID, claims)
|
||||
return
|
||||
}
|
||||
|
||||
http.Error(w, "bad request", http.StatusBadRequest)
|
||||
}
|
||||
|
||||
func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) {
|
||||
func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request, userID int64, claims auth.Claims) {
|
||||
conn, err := s.upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
@@ -168,16 +210,54 @@ func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
go startKeepalive(cw, done)
|
||||
session := &adminSession{userID: userID, claims: claims, conn: cw}
|
||||
go startKeepalive(cw, done, func() bool {
|
||||
return s.validateAdminSession(session.userID, session.claims)
|
||||
})
|
||||
|
||||
s.mu.Lock()
|
||||
s.admins[cw] = struct{}{}
|
||||
s.admins[session] = struct{}{}
|
||||
s.mu.Unlock()
|
||||
|
||||
defer func() {
|
||||
close(done)
|
||||
s.mu.Lock()
|
||||
delete(s.admins, cw)
|
||||
delete(s.admins, session)
|
||||
s.mu.Unlock()
|
||||
_ = conn.Close()
|
||||
}()
|
||||
|
||||
for {
|
||||
if _, _, err := conn.ReadMessage(); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleMonitor(w http.ResponseWriter, r *http.Request, userID int64, claims auth.Claims) {
|
||||
conn, err := s.upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
cw := &connWrap{conn: conn}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
session := &monitorSession{userID: userID, claims: claims, conn: cw}
|
||||
go startKeepalive(cw, done, func() bool {
|
||||
return s.validateMonitorSession(session.userID, session.claims)
|
||||
})
|
||||
|
||||
s.mu.Lock()
|
||||
s.monitors[session] = struct{}{}
|
||||
s.mu.Unlock()
|
||||
|
||||
defer func() {
|
||||
close(done)
|
||||
s.mu.Lock()
|
||||
delete(s.monitors, session)
|
||||
s.mu.Unlock()
|
||||
_ = conn.Close()
|
||||
}()
|
||||
@@ -200,7 +280,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
go startKeepalive(cw, done)
|
||||
go startKeepalive(cw, done, nil)
|
||||
|
||||
version := r.URL.Query().Get("version")
|
||||
httpVal := parseIntDefault(r.URL.Query().Get("http"), 0)
|
||||
@@ -237,7 +317,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
needOfflineBroadcast := false
|
||||
s.mu.Lock()
|
||||
current, ok := s.nodes[nodeID]
|
||||
if ok && current.conn.conn == conn {
|
||||
if ok && current.conn != nil && current.conn.conn == conn {
|
||||
delete(s.nodes, nodeID)
|
||||
needOfflineBroadcast = true
|
||||
}
|
||||
@@ -417,11 +497,7 @@ func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, tim
|
||||
}
|
||||
}
|
||||
|
||||
ns.conn.mu.Lock()
|
||||
_ = ns.conn.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err = ns.conn.conn.WriteMessage(websocket.TextMessage, messageData)
|
||||
_ = ns.conn.conn.SetWriteDeadline(time.Time{})
|
||||
ns.conn.mu.Unlock()
|
||||
err = writeWSMessage(ns.conn, websocket.TextMessage, messageData)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return CommandResult{}, err
|
||||
@@ -537,35 +613,47 @@ func (s *Server) broadcastStatus(nodeID int64, status int) {
|
||||
"data": status,
|
||||
}
|
||||
raw, _ := json.Marshal(payload)
|
||||
s.broadcastToAdmins(string(raw))
|
||||
s.broadcastToRealtime(string(raw))
|
||||
}
|
||||
|
||||
func (s *Server) broadcastInfo(nodeID int64, data string) {
|
||||
payload := broadcastMessage{ID: nodeID, Type: "info", Data: data}
|
||||
raw, _ := json.Marshal(payload)
|
||||
s.broadcastToAdmins(string(raw))
|
||||
s.broadcastToRealtime(string(raw))
|
||||
}
|
||||
|
||||
func (s *Server) broadcastTyped(nodeID int64, msgType string, data string) {
|
||||
payload := broadcastMessage{ID: nodeID, Type: msgType, Data: data}
|
||||
raw, _ := json.Marshal(payload)
|
||||
s.broadcastToAdmins(string(raw))
|
||||
s.broadcastToRealtime(string(raw))
|
||||
}
|
||||
|
||||
func (s *Server) broadcastToAdmins(message string) {
|
||||
func (s *Server) broadcastToRealtime(message string) {
|
||||
s.mu.RLock()
|
||||
admins := make([]*connWrap, 0, len(s.admins))
|
||||
admins := make([]*adminSession, 0, len(s.admins))
|
||||
for c := range s.admins {
|
||||
admins = append(admins, c)
|
||||
}
|
||||
monitors := make([]*monitorSession, 0, len(s.monitors))
|
||||
for c := range s.monitors {
|
||||
monitors = append(monitors, c)
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
|
||||
for _, c := range admins {
|
||||
c.mu.Lock()
|
||||
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
|
||||
_ = c.conn.SetWriteDeadline(time.Time{})
|
||||
c.mu.Unlock()
|
||||
if c == nil || c.conn == nil || c.conn.conn == nil {
|
||||
continue
|
||||
}
|
||||
err := writeWSMessage(c.conn, websocket.TextMessage, []byte(message))
|
||||
if err != nil {
|
||||
log.Printf("websocket broadcast failed: %v", err)
|
||||
}
|
||||
}
|
||||
for _, c := range monitors {
|
||||
if c == nil || c.conn == nil || c.conn.conn == nil {
|
||||
continue
|
||||
}
|
||||
err := writeWSMessage(c.conn, websocket.TextMessage, []byte(message))
|
||||
if err != nil {
|
||||
log.Printf("websocket broadcast failed: %v", err)
|
||||
}
|
||||
@@ -602,7 +690,60 @@ func parseIntDefault(v string, fallback int) int {
|
||||
return x
|
||||
}
|
||||
|
||||
func startKeepalive(cw *connWrap, done <-chan struct{}) {
|
||||
func (s *Server) validateAdminSession(userID int64, claims auth.Claims) bool {
|
||||
if s == nil {
|
||||
return false
|
||||
}
|
||||
if claims.Exp <= time.Now().Unix() {
|
||||
return false
|
||||
}
|
||||
if s.getUserAuthState == nil {
|
||||
return true
|
||||
}
|
||||
state, err := s.getUserAuthState(userID)
|
||||
if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *Server) validateMonitorSession(userID int64, claims auth.Claims) bool {
|
||||
if s == nil {
|
||||
return false
|
||||
}
|
||||
if claims.Exp <= time.Now().Unix() {
|
||||
return false
|
||||
}
|
||||
if s.getUserAuthState != nil {
|
||||
state, err := s.getUserAuthState(userID)
|
||||
if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if s.repo == nil {
|
||||
return false
|
||||
}
|
||||
allowed, err := s.repo.HasMonitorPermission(userID)
|
||||
if err != nil || !allowed {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func writeWSMessage(cw *connWrap, messageType int, payload []byte) error {
|
||||
if cw == nil || cw.conn == nil {
|
||||
return errors.New("websocket connection not initialized")
|
||||
}
|
||||
cw.mu.Lock()
|
||||
defer cw.mu.Unlock()
|
||||
|
||||
_ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := cw.conn.WriteMessage(messageType, payload)
|
||||
_ = cw.conn.SetWriteDeadline(time.Time{})
|
||||
return err
|
||||
}
|
||||
|
||||
func startKeepalive(cw *connWrap, done <-chan struct{}, validate func() bool) {
|
||||
if cw == nil || cw.conn == nil {
|
||||
return
|
||||
}
|
||||
@@ -614,11 +755,11 @@ func startKeepalive(cw *connWrap, done <-chan struct{}) {
|
||||
case <-done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
cw.mu.Lock()
|
||||
_ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := cw.conn.WriteMessage(websocket.PingMessage, nil)
|
||||
_ = cw.conn.SetWriteDeadline(time.Time{})
|
||||
cw.mu.Unlock()
|
||||
if validate != nil && !validate() {
|
||||
_ = cw.conn.Close()
|
||||
return
|
||||
}
|
||||
err := writeWSMessage(cw, websocket.PingMessage, nil)
|
||||
if err != nil {
|
||||
_ = cw.conn.Close()
|
||||
return
|
||||
|
||||
@@ -0,0 +1,215 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/store/repo"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
func TestServeHTTPRejectsDisabledAdminToken(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)
|
||||
}
|
||||
|
||||
server := NewServer(nil, secret)
|
||||
server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 0, PasswordChangedAt: 0}, nil
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/system-info?type=0&secret="+url.QueryEscape(token), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
server.ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusForbidden {
|
||||
t.Fatalf("expected forbidden for disabled admin token, got %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeHTTPRejectsNonAdminToken(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
server := NewServer(nil, secret)
|
||||
server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: 0}, nil
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/system-info?type=0&secret="+url.QueryEscape(token), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
server.ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusForbidden {
|
||||
t.Fatalf("expected forbidden for non-admin token, got %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateAdminSessionRejectsAuthStateChanges(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)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
state *auth.UserAuthState
|
||||
}{
|
||||
{name: "disabled", state: &auth.UserAuthState{ID: 1, RoleID: 0, Status: 0, PasswordChangedAt: 0}},
|
||||
{name: "role changed", state: &auth.UserAuthState{ID: 1, RoleID: 1, Status: 1, PasswordChangedAt: 0}},
|
||||
{name: "password changed", state: &auth.UserAuthState{ID: 1, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs + 1}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server := NewServer(nil, secret)
|
||||
server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) {
|
||||
return tt.state, nil
|
||||
})
|
||||
|
||||
if ok := server.validateAdminSession(1, claims); ok {
|
||||
t.Fatalf("expected session validation to fail for %s state", tt.name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateAdminSessionRejectsExpiredToken(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)
|
||||
}
|
||||
claims.Exp = 1
|
||||
|
||||
server := NewServer(nil, secret)
|
||||
server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: 0}, nil
|
||||
})
|
||||
|
||||
if ok := server.validateAdminSession(1, claims); ok {
|
||||
t.Fatal("expected expired token to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeHTTPAllowsMonitorTokenWithPermission(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
r, err := repo.Open(t.TempDir() + "/monitor.db")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
if err := r.InsertMonitorPermission(2, 123); err != nil {
|
||||
t.Fatalf("insert permission: %v", err)
|
||||
}
|
||||
|
||||
server := NewServer(r, secret)
|
||||
server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: 0}, nil
|
||||
})
|
||||
|
||||
ts := httptest.NewServer(server)
|
||||
defer ts.Close()
|
||||
|
||||
conn, resp, err := websocket.DefaultDialer.Dial(
|
||||
"ws"+strings.TrimPrefix(ts.URL, "http")+"/system-info?type=0&secret="+url.QueryEscape(token),
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
if resp != nil {
|
||||
t.Fatalf("dial websocket error = %v, status=%d", err, resp.StatusCode)
|
||||
}
|
||||
t.Fatalf("dial websocket error = %v", err)
|
||||
}
|
||||
_ = conn.Close()
|
||||
}
|
||||
|
||||
func TestConnWrapSerializesConcurrentWrites(t *testing.T) {
|
||||
serverConn, clientConn := websocketTestPipe(t)
|
||||
defer serverConn.Close()
|
||||
defer clientConn.Close()
|
||||
|
||||
cw := &connWrap{conn: serverConn}
|
||||
readerDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(readerDone)
|
||||
for i := 0; i < 64; i++ {
|
||||
if _, _, err := clientConn.ReadMessage(); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 64; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
if err := writeWSMessage(cw, websocket.TextMessage, []byte("x")); err != nil {
|
||||
t.Errorf("writeWSMessage() error = %v", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
_ = clientConn.Close()
|
||||
<-readerDone
|
||||
}
|
||||
|
||||
func websocketTestPipe(t *testing.T) (*websocket.Conn, *websocket.Conn) {
|
||||
t.Helper()
|
||||
|
||||
upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
|
||||
serverConnCh := make(chan *websocket.Conn, 1)
|
||||
serverErrCh := make(chan error, 1)
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
serverErrCh <- err
|
||||
return
|
||||
}
|
||||
serverConnCh <- conn
|
||||
}))
|
||||
t.Cleanup(ts.Close)
|
||||
|
||||
clientConn, _, err := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(ts.URL, "http"), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("dial websocket: %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-serverErrCh:
|
||||
t.Fatalf("upgrade websocket: %v", err)
|
||||
case serverConn := <-serverConnCh:
|
||||
return serverConn, clientConn
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -9,6 +10,7 @@ import (
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/security"
|
||||
)
|
||||
|
||||
func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
@@ -18,7 +20,9 @@ func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
})
|
||||
|
||||
wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret})(next)
|
||||
wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret, GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: 0}, nil
|
||||
}})(next)
|
||||
|
||||
t.Run("login path is excluded", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", nil)
|
||||
@@ -59,6 +63,9 @@ func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret, GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: 0}, nil
|
||||
}})(next)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
@@ -67,6 +74,83 @@ func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestLoginTokenValidatesThroughRouter(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
seedLegacyUser(t, r, 9110, "router-login-user", "router-login-pass")
|
||||
|
||||
body := bytes.NewBufferString(`{"username":"router-login-user","password":"router-login-pass","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode login response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected login code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
data, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected login data map, got %T", out.Data)
|
||||
}
|
||||
token, _ := data["token"].(string)
|
||||
if token == "" {
|
||||
t.Fatal("expected login token")
|
||||
}
|
||||
|
||||
checkReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/tunnel", nil)
|
||||
checkReq.Header.Set("Authorization", token)
|
||||
checkResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(checkResp, checkReq)
|
||||
assertCode(t, checkResp, 0)
|
||||
}
|
||||
|
||||
func TestLegacyPasswordMigratesOnLogin(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUser(t, r, 9101, "legacy-login-user", "legacy-login-pass")
|
||||
|
||||
body := bytes.NewBufferString(`{"username":"legacy-login-user","password":"legacy-login-pass","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "legacy-login-user", "legacy-login-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "legacy-login-user"); changedAt <= legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to advance on login migration, got %d <= %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisabledLegacyPasswordIsRejectedWithoutMigration(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUserWithStatus(t, r, 9105, "disabled-legacy-user", "disabled-legacy-pass", 0)
|
||||
|
||||
body := bytes.NewBufferString(`{"username":"disabled-legacy-user","password":"disabled-legacy-pass","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCodeMsg(t, resp, -1, "账号被停用")
|
||||
|
||||
user, err := r.GetUserByUsername("disabled-legacy-user")
|
||||
if err != nil {
|
||||
t.Fatalf("get user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
t.Fatal("expected disabled user to exist")
|
||||
}
|
||||
if ok, migrated := security.VerifyPassword(user.Pwd, "disabled-legacy-pass"); !ok || !migrated {
|
||||
t.Fatalf("expected disabled user to remain legacy MD5, got (%v,%v) with hash %q", ok, migrated, user.Pwd)
|
||||
}
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "disabled-legacy-user"); changedAt != legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to remain unchanged, got %d want %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
httpserver "go-backend/internal/http"
|
||||
"go-backend/internal/http/handler"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/repo"
|
||||
|
||||
"gorm.io/gorm"
|
||||
@@ -128,6 +129,52 @@ func TestCaptchaVerifyLoginContract(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestPublicConfigGetAndAuthConfigContract(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
|
||||
for name, value := range map[string]string{
|
||||
"app_name": "FLVX Public",
|
||||
"app_logo": "logo",
|
||||
"app_favicon": "favicon",
|
||||
"app_bg_image": "bg",
|
||||
"cloudflare_site_key": "site-key",
|
||||
"cloudflare_secret_key": "secret-key",
|
||||
"jwt_secret": "jwt-secret",
|
||||
} {
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, name, value, time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
publicReq := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
publicReq.Header.Set("Content-Type", "application/json")
|
||||
publicResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(publicResp, publicReq)
|
||||
assertCode(t, publicResp, 0)
|
||||
|
||||
secretReq := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"jwt_secret"}`))
|
||||
secretReq.Header.Set("Content-Type", "application/json")
|
||||
secretResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(secretResp, secretReq)
|
||||
assertCodeMsg(t, secretResp, 403, "禁止访问敏感配置")
|
||||
|
||||
configReq := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"cloudflare_site_key"}`))
|
||||
configReq.Header.Set("Content-Type", "application/json")
|
||||
configResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(configResp, configReq)
|
||||
assertCode(t, configResp, 0)
|
||||
|
||||
configSecretReq := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"jwt_secret"}`))
|
||||
configSecretReq.Header.Set("Content-Type", "application/json")
|
||||
configSecretResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(configSecretResp, configSecretReq)
|
||||
assertCodeMsg(t, configSecretResp, 403, "禁止访问敏感配置")
|
||||
}
|
||||
|
||||
func TestOpenAPISubStoreContracts(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
|
||||
@@ -209,6 +256,122 @@ func TestOpenAPISubStoreContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestLegacyPasswordMigratesOnSubStore(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUser(t, r, 9102, "legacy-substore-user", "legacy-substore-pass")
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=legacy-substore-user&pwd=legacy-substore-pass", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("read body: %v", err)
|
||||
}
|
||||
expected := "upload=0; download=0; total=107373108658176; expire=2727251700"
|
||||
if string(body) != expected {
|
||||
t.Fatalf("expected body %q, got %q", expected, string(body))
|
||||
}
|
||||
assertUserPasswordIsBcrypt(t, r, "legacy-substore-user", "legacy-substore-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "legacy-substore-user"); changedAt <= legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to advance on sub-store migration, got %d <= %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisabledLegacyPasswordIsRejectedOnSubStoreWithoutMigration(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUserWithStatus(t, r, 9106, "disabled-substore-user", "disabled-substore-pass", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=disabled-substore-user&pwd=disabled-substore-pass", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCodeMsg(t, resp, -1, "账号被停用")
|
||||
|
||||
user, err := r.GetUserByUsername("disabled-substore-user")
|
||||
if err != nil {
|
||||
t.Fatalf("get user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
t.Fatal("expected disabled user to exist")
|
||||
}
|
||||
if ok, migrated := security.VerifyPassword(user.Pwd, "disabled-substore-pass"); !ok || !migrated {
|
||||
t.Fatalf("expected disabled user to remain legacy MD5, got (%v,%v) with hash %q", ok, migrated, user.Pwd)
|
||||
}
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "disabled-substore-user"); changedAt != legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to remain unchanged, got %d want %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserCreateStoresStrongHash(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
startedAt := time.Now().UnixMilli()
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, "contract-jwt-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
body := bytes.NewBufferString(`{"user":"created-user","pwd":"created-pass"}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/create", body)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "created-user", "created-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "created-user"); changedAt < startedAt {
|
||||
t.Fatalf("expected password_changed_at to be set on create, got %d < %d", changedAt, startedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserUpdateStoresStrongHash(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
seedLegacyUser(t, r, 9103, "legacy-update-user", "legacy-update-pass")
|
||||
startedAt := time.Now().UnixMilli()
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, "contract-jwt-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
body := bytes.NewBufferString(`{"id":9103,"user":"updated-user","pwd":"updated-pass","flow":99999,"num":99999,"expTime":2727251700000,"flowResetTime":1,"status":1,"maxConn":0}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/update", body)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "updated-user", "updated-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "updated-user"); changedAt < startedAt {
|
||||
t.Fatalf("expected password_changed_at to be updated on user update, got %d < %d", changedAt, startedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdatePasswordStoresStrongHash(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
seedLegacyUser(t, r, 9104, "legacy-self-user", "legacy-self-pass")
|
||||
startedAt := time.Now().UnixMilli()
|
||||
token, err := auth.GenerateToken(9104, "legacy-self-user", 1, "contract-jwt-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
body := bytes.NewBufferString(`{"newUsername":"self-updated-user","currentPassword":"legacy-self-pass","newPassword":"self-updated-pass","confirmPassword":"self-updated-pass"}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/updatePassword", body)
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "self-updated-user", "self-updated-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "self-updated-user"); changedAt < startedAt {
|
||||
t.Fatalf("expected password_changed_at to be updated on password change, got %d < %d", changedAt, startedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpeedLimitTunnelsRouteRemoved(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
@@ -232,6 +395,7 @@ func TestSpeedLimitTunnelsRouteRemoved(t *testing.T) {
|
||||
func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
seedContractUser(t, r, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
@@ -559,6 +723,71 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestBackupConfigFilteringContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
configs := map[string]string{
|
||||
"app_name": "contract-before",
|
||||
"cloudflare_site_key": "site-key-before",
|
||||
"jwt_secret": "jwt-before",
|
||||
"license_key": "license-before",
|
||||
"cloudflare_secret_key": "cloudflare-before",
|
||||
}
|
||||
for name, value := range configs {
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, name, value, time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
payload := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
for _, key := range []string{"jwt_secret", "license_key", "cloudflare_secret_key"} {
|
||||
if _, ok := payload.Configs[key]; ok {
|
||||
t.Fatalf("expected %s to be omitted from exported configs: %+v", key, payload.Configs)
|
||||
}
|
||||
}
|
||||
if payload.Configs["app_name"] != "contract-before" {
|
||||
t.Fatalf("expected public config to be exported, got %+v", payload.Configs)
|
||||
}
|
||||
|
||||
payload.Configs["app_name"] = "contract-after"
|
||||
payload.Configs["jwt_secret"] = "jwt-after"
|
||||
payload.Configs["license_key"] = "license-after"
|
||||
payload.Configs["cloudflare_secret_key"] = "cloudflare-after"
|
||||
raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: payload})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal import payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/import", bytes.NewReader(raw))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
var out response.R
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode import response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected import code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
assertConfigValue(t, r, "app_name", "contract-after")
|
||||
assertConfigValue(t, r, "jwt_secret", "jwt-before")
|
||||
assertConfigValue(t, r, "license_key", "license-before")
|
||||
assertConfigValue(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
}
|
||||
|
||||
type backupExportPayload struct {
|
||||
Version string `json:"version"`
|
||||
ExportedAt int64 `json:"exportedAt"`
|
||||
@@ -619,6 +848,66 @@ func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *repo.Re
|
||||
return httpserver.NewRouter(h, jwtSecret), r
|
||||
}
|
||||
|
||||
func assertConfigValue(t *testing.T, r *repo.Repository, name, want string) {
|
||||
t.Helper()
|
||||
cfg, err := r.GetConfigByName(name)
|
||||
if err != nil {
|
||||
t.Fatalf("get config %s: %v", name, err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != want {
|
||||
t.Fatalf("expected config %s=%q, got %+v", name, want, cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func seedLegacyUser(t *testing.T, r *repo.Repository, id int64, username, password string) int64 {
|
||||
return seedLegacyUserWithStatus(t, r, id, username, password, 1)
|
||||
}
|
||||
|
||||
func seedContractUser(t *testing.T, r *repo.Repository, id int64, username string, roleID, status int) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
passwordChangedAt := now - 10_000
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, max_conn, created_time, updated_time, status, password_changed_at)
|
||||
VALUES(?, ?, ?, ?, 2727251700000, 99999, 0, 0, 1, 99999, 0, ?, ?, ?, ?)
|
||||
`, id, username, security.MD5("contract-pass"), roleID, now, now, status, passwordChangedAt).Error; err != nil {
|
||||
t.Fatalf("seed contract user %s: %v", username, err)
|
||||
}
|
||||
return passwordChangedAt
|
||||
}
|
||||
|
||||
func seedLegacyUserWithStatus(t *testing.T, r *repo.Repository, id int64, username, password string, status int) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
legacyChangedAt := now - 10_000
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, max_conn, created_time, updated_time, status, password_changed_at)
|
||||
VALUES(?, ?, ?, 1, 2727251700000, 99999, 0, 0, 1, 99999, 0, ?, ?, ?, ?)
|
||||
`, id, username, security.MD5(password), now, now, status, legacyChangedAt).Error; err != nil {
|
||||
t.Fatalf("seed legacy user %s: %v", username, err)
|
||||
}
|
||||
return legacyChangedAt
|
||||
}
|
||||
|
||||
func mustQueryPasswordChangedAtByUsername(t *testing.T, r *repo.Repository, username string) int64 {
|
||||
t.Helper()
|
||||
return mustQueryInt64(t, r, `SELECT password_changed_at FROM user WHERE user = ?`, username)
|
||||
}
|
||||
|
||||
func assertUserPasswordIsBcrypt(t *testing.T, r *repo.Repository, username, password string) {
|
||||
t.Helper()
|
||||
user, err := r.GetUserByUsername(username)
|
||||
if err != nil {
|
||||
t.Fatalf("get user %s: %v", username, err)
|
||||
}
|
||||
if user == nil {
|
||||
t.Fatalf("expected user %s to exist", username)
|
||||
}
|
||||
if ok, migrated := security.VerifyPassword(user.Pwd, password); !ok || migrated {
|
||||
t.Fatalf("expected bcrypt password for %s, got %q", username, user.Pwd)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "legacy-2.0.7-beta.db")
|
||||
legacyDB, err := sql.Open("sqlite", dbPath)
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
func TestNodeMetricsEndpoints(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
seedContractUser(t, repo, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
@@ -1302,6 +1303,7 @@ func TestMonitoringAuthRequired(t *testing.T) {
|
||||
func TestMonitorAccessEndpoint(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
seedContractUser(t, repo, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
@@ -1357,6 +1359,7 @@ func TestMonitorAccessEndpoint(t *testing.T) {
|
||||
func TestMonitoringPermissionRequired(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
seedContractUser(t, repo, 2, "normal_user", 1, 1)
|
||||
|
||||
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
|
||||
@@ -12,7 +12,8 @@ import (
|
||||
|
||||
func TestStorageSummaryRequiresAdminAndReturnsSize(t *testing.T) {
|
||||
secret := "storage-contract-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
router, r := setupContractRouter(t, secret)
|
||||
seedContractUser(t, r, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
|
||||
+8
-12
@@ -22,6 +22,7 @@ require (
|
||||
github.com/coreos/go-iptables v0.7.0 // indirect
|
||||
github.com/danieljoos/wincred v1.2.0 // indirect
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
||||
github.com/dunglas/httpsfv v1.1.0 // indirect
|
||||
github.com/fsnotify/fsnotify v1.7.0 // indirect
|
||||
github.com/gabriel-vasile/mimetype v1.4.3 // indirect
|
||||
github.com/gin-contrib/cors v1.7.2 // indirect
|
||||
@@ -37,13 +38,11 @@ require (
|
||||
github.com/go-playground/universal-translator v0.18.1 // indirect
|
||||
github.com/go-playground/validator/v10 v10.20.0 // indirect
|
||||
github.com/go-redis/redis/v8 v8.11.5 // indirect
|
||||
github.com/go-task/slim-sprig/v3 v3.0.0 // indirect
|
||||
github.com/gobwas/glob v0.2.3 // indirect
|
||||
github.com/goccy/go-json v0.10.2 // indirect
|
||||
github.com/godbus/dbus/v5 v5.1.0 // indirect
|
||||
github.com/golang/snappy v0.0.4 // indirect
|
||||
github.com/google/gopacket v1.1.20-0.20220810144506-32ee38206866 // indirect
|
||||
github.com/google/pprof v0.0.0-20241210010833-40e02aabc2ad // indirect
|
||||
github.com/google/uuid v1.6.0 // indirect
|
||||
github.com/gorilla/websocket v1.5.3 // indirect
|
||||
github.com/gravitational/trace v1.1.16-0.20220114165159-14a9a7dd6aaf // indirect
|
||||
@@ -61,13 +60,11 @@ require (
|
||||
github.com/mitchellh/mapstructure v1.5.0 // indirect
|
||||
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
|
||||
github.com/modern-go/reflect2 v1.0.2 // indirect
|
||||
github.com/onsi/ginkgo/v2 v2.22.0 // indirect
|
||||
github.com/patrickmn/go-cache v2.1.0+incompatible // indirect
|
||||
github.com/pelletier/go-toml/v2 v2.2.2 // indirect
|
||||
github.com/pion/dtls/v2 v2.2.6 // indirect
|
||||
github.com/pion/logging v0.2.2 // indirect
|
||||
github.com/pion/transport/v2 v2.0.2 // indirect
|
||||
github.com/pion/udp/v2 v2.0.1 // indirect
|
||||
github.com/pion/dtls/v3 v3.0.11 // indirect
|
||||
github.com/pion/logging v0.2.4 // indirect
|
||||
github.com/pion/transport/v4 v4.0.1 // indirect
|
||||
github.com/pires/go-proxyproto v0.7.0 // indirect
|
||||
github.com/pkg/errors v0.9.1 // indirect
|
||||
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
|
||||
@@ -75,9 +72,9 @@ require (
|
||||
github.com/prometheus/client_model v0.6.0 // indirect
|
||||
github.com/prometheus/common v0.48.0 // indirect
|
||||
github.com/prometheus/procfs v0.12.0 // indirect
|
||||
github.com/quic-go/qpack v0.5.1 // indirect
|
||||
github.com/quic-go/quic-go v0.49.1 // indirect
|
||||
github.com/quic-go/webtransport-go v0.8.1-0.20241018022711-4ac2c9250e66 // indirect
|
||||
github.com/quic-go/qpack v0.6.0 // indirect
|
||||
github.com/quic-go/quic-go v0.59.0 // indirect
|
||||
github.com/quic-go/webtransport-go v0.10.0 // indirect
|
||||
github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3 // indirect
|
||||
github.com/rs/xid v1.3.0 // indirect
|
||||
github.com/sagikazarmark/locafero v0.4.0 // indirect
|
||||
@@ -86,7 +83,7 @@ require (
|
||||
github.com/shadowsocks/shadowsocks-go v0.0.0-20200409064450-3e585ff90601 // indirect
|
||||
github.com/shirou/gopsutil/v3 v3.24.5 // indirect
|
||||
github.com/shoenig/go-m1cpu v0.1.6 // indirect
|
||||
github.com/sirupsen/logrus v1.8.1 // indirect
|
||||
github.com/sirupsen/logrus v1.8.3 // indirect
|
||||
github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8 // indirect
|
||||
github.com/sourcegraph/conc v0.3.0 // indirect
|
||||
github.com/spf13/afero v1.11.0 // indirect
|
||||
@@ -110,7 +107,6 @@ require (
|
||||
github.com/yl2chen/cidranger v1.0.2 // indirect
|
||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||
github.com/zalando/go-keyring v0.2.4 // indirect
|
||||
go.uber.org/mock v0.5.0 // indirect
|
||||
go.uber.org/multierr v1.11.0 // indirect
|
||||
golang.org/x/arch v0.8.0 // indirect
|
||||
golang.org/x/crypto v0.50.0 // indirect
|
||||
|
||||
+21
-51
@@ -33,12 +33,12 @@ github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1
|
||||
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
|
||||
github.com/dunglas/httpsfv v1.1.0 h1:Jw76nAyKWKZKFrpMMcL76y35tOpYHqQPzHQiwDvpe54=
|
||||
github.com/dunglas/httpsfv v1.1.0/go.mod h1:zID2mqw9mFsnt7YC3vYQ9/cjq30q41W+1AnDwH8TiMg=
|
||||
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
|
||||
github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
|
||||
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
|
||||
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
|
||||
github.com/francoispqt/gojay v1.2.13 h1:d2m3sFjloqoIUQU3TsHBgj6qg/BVGlTBeHDUmyJnXKk=
|
||||
github.com/francoispqt/gojay v1.2.13/go.mod h1:ehT5mTG4ua4581f1++1WLG0vPdaA9HaiDsoyrBGkyDY=
|
||||
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
||||
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
||||
github.com/fsnotify/fsnotify v1.7.0 h1:8JEhPFa5W2WU7YfeZzPNqzMP6Lwt7L2715Ggo0nosvA=
|
||||
@@ -79,8 +79,6 @@ github.com/go-playground/validator/v10 v10.20.0 h1:K9ISHbSaI0lyB2eWMPJo+kOS/FBEx
|
||||
github.com/go-playground/validator/v10 v10.20.0/go.mod h1:dbuPbCMFw/DrkbEynArYaCwl3amGuJotoKCe95atGMM=
|
||||
github.com/go-redis/redis/v8 v8.11.5 h1:AcZZR7igkdvfVmQTPnu9WE37LRrO/YrBH5zWyjDC0oI=
|
||||
github.com/go-redis/redis/v8 v8.11.5/go.mod h1:gREzHqY1hg6oD9ngVRbLStwAWKhA0FEgq8Jd4h5lpwo=
|
||||
github.com/go-task/slim-sprig/v3 v3.0.0 h1:sUs3vkvUymDpBKi3qH1YSqBQk9+9D/8M2mN1vB6EwHI=
|
||||
github.com/go-task/slim-sprig/v3 v3.0.0/go.mod h1:W848ghGpv3Qj3dhTPRyJypKRiqCdHZiAzKg9hl15HA8=
|
||||
github.com/gobwas/glob v0.2.3 h1:A4xDbljILXROh+kObIiy5kIaPYD8e96x1tgBhUI5J+Y=
|
||||
github.com/gobwas/glob v0.2.3/go.mod h1:d3Ez4x06l9bZtSvzIay5+Yzi0fmZzPgnTbPcKjJAkT8=
|
||||
github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
|
||||
@@ -114,8 +112,6 @@ github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX
|
||||
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
|
||||
github.com/google/gopacket v1.1.20-0.20220810144506-32ee38206866 h1:NaJi58bCZZh0jjPw78EqDZekPEfhlzYE01C5R+zh1tE=
|
||||
github.com/google/gopacket v1.1.20-0.20220810144506-32ee38206866/go.mod h1:riddUzxTSBpJXk3qBHtYr4qOhFhT6k/1c0E3qkQjQpA=
|
||||
github.com/google/pprof v0.0.0-20241210010833-40e02aabc2ad h1:a6HEuzUHeKH6hwfN/ZoQgRgVIWFJljSWa/zetS2WTvg=
|
||||
github.com/google/pprof v0.0.0-20241210010833-40e02aabc2ad/go.mod h1:vavhavw2zAxS5dIdcRluK6cSGGPlZynqzFM8NdvU144=
|
||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||
@@ -163,22 +159,18 @@ github.com/nxadm/tail v1.4.8 h1:nPr65rt6Y5JFSKQO7qToXr7pePgD6Gwiw05lkbyAQTE=
|
||||
github.com/nxadm/tail v1.4.8/go.mod h1:+ncqLTQzXmGhMZNUePPaPqPvBxHAIsmXswZKocGu+AU=
|
||||
github.com/onsi/ginkgo v1.16.5 h1:8xi0RTUf59SOSfEtZMvwTvXYMzG4gV23XVHOZiXNtnE=
|
||||
github.com/onsi/ginkgo v1.16.5/go.mod h1:+E8gABHa3K6zRBolWtd+ROzc/U5bkGt0FwiG042wbpU=
|
||||
github.com/onsi/ginkgo/v2 v2.22.0 h1:Yed107/8DjTr0lKCNt7Dn8yQ6ybuDRQoMGrNFKzMfHg=
|
||||
github.com/onsi/ginkgo/v2 v2.22.0/go.mod h1:7Du3c42kxCUegi0IImZ1wUQzMBVecgIHjR1C+NkhLQo=
|
||||
github.com/onsi/gomega v1.34.2 h1:pNCwDkzrsv7MS9kpaQvVb1aVLahQXyJ/Tv5oAZMI3i8=
|
||||
github.com/onsi/gomega v1.34.2/go.mod h1:v1xfxRgk0KIsG+QOdm7p8UosrOzPYRo60fd3B/1Dukc=
|
||||
github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc=
|
||||
github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ=
|
||||
github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6Wq+LM=
|
||||
github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs=
|
||||
github.com/pion/dtls/v2 v2.2.6 h1:yXMxKr0Skd+Ub6A8UqXTRLSywskx93ooMRHsQUtd+Z4=
|
||||
github.com/pion/dtls/v2 v2.2.6/go.mod h1:t8fWJCIquY5rlQZwA2yWxUS1+OCrAdXrhVKXB5oD/wY=
|
||||
github.com/pion/logging v0.2.2 h1:M9+AIj/+pxNsDfAT64+MAVgJO0rsyLnoJKCqf//DoeY=
|
||||
github.com/pion/logging v0.2.2/go.mod h1:k0/tDVsRCX2Mb2ZEmTqNa7CWsQPc+YYCB7Q+5pahoms=
|
||||
github.com/pion/transport/v2 v2.0.2 h1:St+8o+1PEzPT51O9bv+tH/KYYLMNR5Vwm5Z3Qkjsywg=
|
||||
github.com/pion/transport/v2 v2.0.2/go.mod h1:vrz6bUbFr/cjdwbnxq8OdDDzHf7JJfGsIRkxfpZoTA0=
|
||||
github.com/pion/udp/v2 v2.0.1 h1:xP0z6WNux1zWEjhC7onRA3EwwSliXqu1ElUZAQhUP54=
|
||||
github.com/pion/udp/v2 v2.0.1/go.mod h1:B7uvTMP00lzWdyMr/1PVZXtV3wpPIxBRd4Wl6AksXn8=
|
||||
github.com/pion/dtls/v3 v3.0.11 h1:zqn8YhoAU7d9whsWLhNiQlbB8QdpJj8XQVSc5ImUons=
|
||||
github.com/pion/dtls/v3 v3.0.11/go.mod h1:YEmmBYIoBsY3jmG56dsziTv/Lca9y4Om83370CXfqJ8=
|
||||
github.com/pion/logging v0.2.4 h1:tTew+7cmQ+Mc1pTBLKH2puKsOvhm32dROumOZ655zB8=
|
||||
github.com/pion/logging v0.2.4/go.mod h1:DffhXTKYdNZU+KtJ5pyQDjvOAh/GsNSyv1lbkFbe3so=
|
||||
github.com/pion/transport/v4 v4.0.1 h1:sdROELU6BZ63Ab7FrOLn13M6YdJLY20wldXW2Cu2k8o=
|
||||
github.com/pion/transport/v4 v4.0.1/go.mod h1:nEuEA4AD5lPdcIegQDpVLgNoDGreqM/YqmEx3ovP4jM=
|
||||
github.com/pires/go-proxyproto v0.7.0 h1:IukmRewDQFWC7kfnb66CSomk2q/seBuilHBYFwyq0Hs=
|
||||
github.com/pires/go-proxyproto v0.7.0/go.mod h1:Vz/1JPY/OACxWGQNIRY2BeyDmpoaWmEP40O9LbuiFR4=
|
||||
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
|
||||
@@ -197,12 +189,12 @@ github.com/prometheus/common v0.48.0 h1:QO8U2CdOzSn1BBsmXJXduaaW+dY/5QLjfB8svtSz
|
||||
github.com/prometheus/common v0.48.0/go.mod h1:0/KsvlIEfPQCQ5I2iNSAWKPZziNCvRs5EC6ILDTlAPc=
|
||||
github.com/prometheus/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo=
|
||||
github.com/prometheus/procfs v0.12.0/go.mod h1:pcuDEFsWDnvcgNzo4EEweacyhjeA9Zk3cnaOZAZEfOo=
|
||||
github.com/quic-go/qpack v0.5.1 h1:giqksBPnT/HDtZ6VhtFKgoLOWmlyo9Ei6u9PqzIMbhI=
|
||||
github.com/quic-go/qpack v0.5.1/go.mod h1:+PC4XFrEskIVkcLzpEkbLqq1uCoxPhQuvK5rH1ZgaEg=
|
||||
github.com/quic-go/quic-go v0.49.1 h1:e5JXpUyF0f2uFjckQzD8jTghZrOUK1xxDqqZhlwixo0=
|
||||
github.com/quic-go/quic-go v0.49.1/go.mod h1:s2wDnmCdooUQBmQfpUSTCYBl1/D4FcqbULMMkASvR6s=
|
||||
github.com/quic-go/webtransport-go v0.8.1-0.20241018022711-4ac2c9250e66 h1:4WFk6u3sOT6pLa1kQ50ZVdm8BQFgJNA117cepZxtLIg=
|
||||
github.com/quic-go/webtransport-go v0.8.1-0.20241018022711-4ac2c9250e66/go.mod h1:Vp72IJajgeOL6ddqrAhmp7IM9zbTcgkQxD/YdxrVwMw=
|
||||
github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8=
|
||||
github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII=
|
||||
github.com/quic-go/quic-go v0.59.0 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw=
|
||||
github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU=
|
||||
github.com/quic-go/webtransport-go v0.10.0 h1:LqXXPOXuETY5Xe8ITdGisBzTYmUOy5eSj+9n4hLTjHI=
|
||||
github.com/quic-go/webtransport-go v0.10.0/go.mod h1:LeGIXr5BQKE3UsynwVBeQrU1TPrbh73MGoC6jd+V7ow=
|
||||
github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3 h1:f/FNXud6gA3MNr8meMVVGxhp+QBTqY91tM8HjEuMjGg=
|
||||
github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3/go.mod h1:HgjTstvQsPGkxUsCd2KWxErBblirPizecHcpD3ffK+s=
|
||||
github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ=
|
||||
@@ -224,8 +216,8 @@ github.com/shoenig/go-m1cpu v0.1.6/go.mod h1:1JJMcUBvfNwpq05QDQVAnx3gUHr9IYF7GNg
|
||||
github.com/shoenig/test v0.6.4 h1:kVTaSd7WLz5WZ2IaoM0RSzRsUD+m8wRR+5qvntpn4LU=
|
||||
github.com/shoenig/test v0.6.4/go.mod h1:byHiCGXqrVaflBLAMq/srcZIHynQPQgeyvkvXnjqq0k=
|
||||
github.com/sirupsen/logrus v1.7.0/go.mod h1:yWOB1SBYBC5VeMP7gHvWumXLIWorT60ONWic61uBYv0=
|
||||
github.com/sirupsen/logrus v1.8.1 h1:dJKuHgqk1NNQlqoA6BTlM1Wf9DOH3NBjQyu0h9+AZZE=
|
||||
github.com/sirupsen/logrus v1.8.1/go.mod h1:yWOB1SBYBC5VeMP7gHvWumXLIWorT60ONWic61uBYv0=
|
||||
github.com/sirupsen/logrus v1.8.3 h1:DBBfY8eMYazKEJHb3JKpSPfpgd2mBCoNFlQx6C5fftU=
|
||||
github.com/sirupsen/logrus v1.8.3/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ=
|
||||
github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8 h1:TG/diQgUe0pntT/2D9tmUCz4VNwm9MfrtPr0SU2qSX8=
|
||||
github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8/go.mod h1:P5HUIBuIWKbyjl083/loAegFkfbFNx5i2qEP4CNbm7E=
|
||||
github.com/sourcegraph/conc v0.3.0 h1:OQTbbt6P72L20UqAkXXuLOj79LfEanQ+YQFNpLA9ySo=
|
||||
@@ -252,8 +244,9 @@ github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/
|
||||
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
|
||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
||||
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
|
||||
github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
|
||||
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
|
||||
github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU=
|
||||
github.com/templexxx/cpu v0.1.0 h1:wVM+WIJP2nYaxVxqgHPD4wGA2aJ9rvrQRV8CvFzNb40=
|
||||
@@ -290,7 +283,6 @@ github.com/xtaci/tcpraw v1.2.25 h1:VDlqo0op17JeXBM6e2G9ocCNLOJcw9mZbobMbJjo0vk=
|
||||
github.com/xtaci/tcpraw v1.2.25/go.mod h1:dKyZ2V75s0cZ7cbgJYdxPvms7af0joIeOyx1GgJQbLk=
|
||||
github.com/yl2chen/cidranger v1.0.2 h1:lbOWZVCG1tCRX4u24kuM1Tb4nHqWkDxwLdoS+SevawU=
|
||||
github.com/yl2chen/cidranger v1.0.2/go.mod h1:9U1yz7WPYDwf0vpNWFaeRh0bjwz5RVgRy/9UEQfHl0g=
|
||||
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
||||
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
||||
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||
github.com/zalando/go-keyring v0.2.4 h1:wi2xxTqdiwMKbM6TWwi+uJCG/Tum2UV0jqaQhCa9/68=
|
||||
@@ -307,8 +299,8 @@ go.opentelemetry.io/otel/sdk/metric v1.39.0 h1:cXMVVFVgsIf2YL6QkRF4Urbr/aMInf+2W
|
||||
go.opentelemetry.io/otel/sdk/metric v1.39.0/go.mod h1:xq9HEVH7qeX69/JnwEfp6fVq5wosJsY1mt4lLfYdVew=
|
||||
go.opentelemetry.io/otel/trace v1.39.0 h1:2d2vfpEDmCJ5zVYz7ijaJdOF59xLomrvj7bjt6/qCJI=
|
||||
go.opentelemetry.io/otel/trace v1.39.0/go.mod h1:88w4/PnZSazkGzz/w84VHpQafiU4EtqqlVdxWy+rNOA=
|
||||
go.uber.org/mock v0.5.0 h1:KAMbZvZPyBPWgD14IrIQ38QCyjwpvVVV6K/bHl1IwQU=
|
||||
go.uber.org/mock v0.5.0/go.mod h1:ge71pBPLYDk7QIi1LupWxdAykm7KIEFchiOqd6z7qMM=
|
||||
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
||||
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
||||
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
||||
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||
golang.org/x/arch v0.0.0-20210923205945-b76863e36670/go.mod h1:5om86z9Hs0C8fWVUuoMHwpExlXzs5Tkyp9hOrfG7pp8=
|
||||
@@ -320,8 +312,6 @@ golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPh
|
||||
golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.0.0-20201016220609-9e8e0b390897/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.0.0-20210220033148-5ea612d1eb83/go.mod h1:jdWPYTVW3xRLrWPugEBEK3UY2ZEsg3UU495nc5E+M+I=
|
||||
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
|
||||
golang.org/x/crypto v0.5.0/go.mod h1:NK/OQwhpMQP3MwtdjgLlYHnH9ebylxKWv3e0fK+mkQU=
|
||||
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
||||
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
||||
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
|
||||
@@ -332,7 +322,6 @@ golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvx
|
||||
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
|
||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
|
||||
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
||||
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
|
||||
golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
|
||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
@@ -343,17 +332,12 @@ golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn
|
||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||
golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||
golang.org/x/net v0.0.0-20201031054903-ff519b6c9102/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||
golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c=
|
||||
golang.org/x/net v0.5.0/go.mod h1:DivGGAXEgPSlEBzxGzZI+ZLohi+xUj054jfeKui00ws=
|
||||
golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
|
||||
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
|
||||
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
|
||||
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
|
||||
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
@@ -365,13 +349,9 @@ golang.org/x/sys v0.0.0-20191026070338-33540a1f6037/go.mod h1:h1NjWce9XRLGQEsW7w
|
||||
golang.org/x/sys v0.0.0-20200217220822-9197077df867/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20200728102440-3e129f6d46b1/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20201204225414-ed752295db88/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.4.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
@@ -379,17 +359,10 @@ golang.org/x/sys v0.11.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
||||
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/term v0.0.0-20201117132131-f5c789dd3221/go.mod h1:Nr5EML6q2oocZ2LXRh80K7BxOlk5/8JxuGnuhpl+muw=
|
||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
|
||||
golang.org/x/term v0.4.0/go.mod h1:9P2UbLfCdcvo3p/nzKvsmas4TnlujnuoV9hGgYzW1lQ=
|
||||
golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k=
|
||||
golang.org/x/term v0.42.0 h1:UiKe+zDFmJobeJ5ggPwOshJIVt6/Ft0rcfrXZDLWAWY=
|
||||
golang.org/x/term v0.42.0/go.mod h1:Dq/D+snpsbazcBG5+F9Q1n2rXV8Ma+71xEjTRufARgY=
|
||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||
golang.org/x/text v0.6.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
|
||||
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
|
||||
golang.org/x/text v0.36.0 h1:JfKh3XmcRPqZPKevfXVpI1wXPTqbkE5f7JA92a55Yxg=
|
||||
golang.org/x/text v0.36.0/go.mod h1:NIdBknypM8iqVmPiuco0Dh6P5Jcdk8lJL0CUebqK164=
|
||||
golang.org/x/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE=
|
||||
@@ -399,12 +372,9 @@ golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGm
|
||||
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
|
||||
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
|
||||
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
|
||||
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
||||
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
||||
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
|
||||
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
|
||||
golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
|
||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
|
||||
|
||||
@@ -2,6 +2,8 @@ package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
|
||||
"github.com/go-gost/core/chain"
|
||||
"github.com/go-gost/core/hop"
|
||||
@@ -38,11 +40,12 @@ type chainNamer interface {
|
||||
}
|
||||
|
||||
type Chain struct {
|
||||
name string
|
||||
hops []hop.Hop
|
||||
marker selector.Marker
|
||||
metadata metadata.Metadata
|
||||
logger logger.Logger
|
||||
name string
|
||||
hops []hop.Hop
|
||||
ownedHops []hop.Hop
|
||||
marker selector.Marker
|
||||
metadata metadata.Metadata
|
||||
logger logger.Logger
|
||||
}
|
||||
|
||||
func NewChain(name string, opts ...ChainOption) *Chain {
|
||||
@@ -61,8 +64,15 @@ func NewChain(name string, opts ...ChainOption) *Chain {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Chain) AddHop(hop hop.Hop) {
|
||||
func (c *Chain) AddHop(hop hop.Hop, owned ...bool) {
|
||||
c.hops = append(c.hops, hop)
|
||||
isOwned := true
|
||||
if len(owned) > 0 {
|
||||
isOwned = owned[0]
|
||||
}
|
||||
if isOwned {
|
||||
c.ownedHops = append(c.ownedHops, hop)
|
||||
}
|
||||
}
|
||||
|
||||
// Metadata implements metadata.Metadatable interface.
|
||||
@@ -112,6 +122,36 @@ func (c *Chain) Route(ctx context.Context, network, address string, opts ...chai
|
||||
return rt
|
||||
}
|
||||
|
||||
// Retire gracefully drains resources owned by a chain that has been replaced.
|
||||
func (c *Chain) Retire() {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
for _, h := range c.ownedHops {
|
||||
if retirer, ok := h.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
continue
|
||||
}
|
||||
if closer, ok := h.(io.Closer); ok {
|
||||
_ = closer.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases all resources owned by the chain.
|
||||
func (c *Chain) Close() error {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
var errs []error
|
||||
for _, h := range c.ownedHops {
|
||||
if closer, ok := h.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
type chainGroup struct {
|
||||
chains []chain.Chainer
|
||||
selector selector.Selector[chain.Chainer]
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
corechain "github.com/go-gost/core/chain"
|
||||
corehop "github.com/go-gost/core/hop"
|
||||
)
|
||||
|
||||
type lifecycleTestHop struct {
|
||||
selected int
|
||||
retired int
|
||||
closed int
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Select(context.Context, ...corehop.SelectOption) *corechain.Node {
|
||||
h.selected++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Retire() {
|
||||
h.retired++
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Close() error {
|
||||
h.closed++
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestChainRoutesThroughSharedHopWithoutOwningLifecycle(t *testing.T) {
|
||||
hop := &lifecycleTestHop{}
|
||||
chain := NewChain("shared-hop")
|
||||
chain.AddHop(hop, false)
|
||||
|
||||
if route := chain.Route(context.Background(), "tcp", "example.com:443"); route == nil {
|
||||
t.Fatal("route is nil")
|
||||
}
|
||||
if hop.selected != 1 {
|
||||
t.Fatalf("shared hop selected %d times, want 1", hop.selected)
|
||||
}
|
||||
|
||||
chain.Retire()
|
||||
if err := chain.Close(); err != nil {
|
||||
t.Fatalf("close chain: %v", err)
|
||||
}
|
||||
if hop.retired != 0 || hop.closed != 0 {
|
||||
t.Fatalf("shared hop lifecycle changed: retired=%d closed=%d", hop.retired, hop.closed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainRetiresAndClosesOwnedHop(t *testing.T) {
|
||||
hop := &lifecycleTestHop{}
|
||||
chain := NewChain("owned-hop")
|
||||
chain.AddHop(hop)
|
||||
|
||||
chain.Retire()
|
||||
if err := chain.Close(); err != nil {
|
||||
t.Fatalf("close chain: %v", err)
|
||||
}
|
||||
if hop.retired != 1 || hop.closed != 1 {
|
||||
t.Fatalf("owned hop lifecycle: retired=%d closed=%d, want 1/1", hop.retired, hop.closed)
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,8 @@ package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
|
||||
"github.com/go-gost/core/chain"
|
||||
@@ -102,5 +104,34 @@ func (tr *Transport) Options() *chain.TransportOptions {
|
||||
func (tr *Transport) Copy() chain.Transporter {
|
||||
tr2 := &Transport{}
|
||||
*tr2 = *tr
|
||||
return tr
|
||||
return tr2
|
||||
}
|
||||
|
||||
// Retire prevents long-lived dialer sessions owned by an obsolete chain from
|
||||
// accepting new streams while allowing existing streams to drain.
|
||||
func (tr *Transport) Retire() {
|
||||
if tr == nil {
|
||||
return
|
||||
}
|
||||
if retirer, ok := tr.dialer.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
if retirer, ok := tr.connector.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases transport-owned dialer and connector resources.
|
||||
func (tr *Transport) Close() error {
|
||||
if tr == nil {
|
||||
return nil
|
||||
}
|
||||
var errs []error
|
||||
if closer, ok := tr.dialer.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
if closer, ok := tr.connector.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
corechain "github.com/go-gost/core/chain"
|
||||
)
|
||||
|
||||
type copyTestRoute struct{}
|
||||
|
||||
func (copyTestRoute) Dial(context.Context, string, string, ...corechain.DialOption) (net.Conn, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (copyTestRoute) Bind(context.Context, string, string, ...corechain.BindOption) (net.Listener, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (copyTestRoute) Nodes() []*corechain.Node {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestTransportCopyReturnsIndependentTransport(t *testing.T) {
|
||||
originalRoute := copyTestRoute{}
|
||||
replacementRoute := ©TestRoute{}
|
||||
original := NewTransport(nil, nil, corechain.RouteTransportOption(originalRoute))
|
||||
|
||||
copied, ok := original.Copy().(*Transport)
|
||||
if !ok {
|
||||
t.Fatalf("copy type = %T, want *Transport", original.Copy())
|
||||
}
|
||||
if copied == original {
|
||||
t.Fatal("Copy returned the original transport")
|
||||
}
|
||||
|
||||
copied.Options().Route = replacementRoute
|
||||
if original.Options().Route != originalRoute {
|
||||
t.Fatal("mutating copied transport changed original route")
|
||||
}
|
||||
if copied.Options().Route != replacementRoute {
|
||||
t.Fatal("copied transport did not retain its independent route")
|
||||
}
|
||||
}
|
||||
@@ -35,16 +35,18 @@ func ParseChain(cfg *config.ChainConfig, log logger.Logger) (chain.Chainer, erro
|
||||
for _, ch := range cfg.Hops {
|
||||
var hop hop.Hop
|
||||
var err error
|
||||
owned := false
|
||||
|
||||
if ch.Nodes != nil || ch.Plugin != nil {
|
||||
if hop, err = hop_parser.ParseHop(ch, log); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
owned = true
|
||||
} else {
|
||||
hop = registry.HopRegistry().Get(ch.Name)
|
||||
}
|
||||
if hop != nil {
|
||||
c.AddHop(hop)
|
||||
c.AddHop(hop, owned)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -10,7 +10,8 @@ import (
|
||||
md "github.com/go-gost/core/metadata"
|
||||
xdtls "github.com/go-gost/x/internal/util/dtls"
|
||||
"github.com/go-gost/x/registry"
|
||||
"github.com/pion/dtls/v2"
|
||||
"github.com/pion/dtls/v3"
|
||||
dtlsnet "github.com/pion/dtls/v3/pkg/net"
|
||||
)
|
||||
|
||||
func init() {
|
||||
@@ -66,9 +67,13 @@ func (d *dtlsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialO
|
||||
MTU: d.md.mtu,
|
||||
}
|
||||
|
||||
c, err := dtls.ClientWithContext(ctx, conn, &config)
|
||||
c, err := dtls.Client(dtlsnet.PacketConnFromConn(conn), conn.RemoteAddr(), &config)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := c.HandshakeContext(ctx); err != nil {
|
||||
_ = c.Close()
|
||||
return nil, err
|
||||
}
|
||||
return xdtls.Conn(c, d.md.bufferSize), nil
|
||||
}
|
||||
|
||||
@@ -72,7 +72,7 @@ func (d *http3Dialer) Dial(ctx context.Context, addr string, opts ...dialer.Dial
|
||||
// Timeout: 60 * time.Second,
|
||||
Transport: &http3.Transport{
|
||||
TLSClientConfig: d.options.TLSConfig,
|
||||
Dial: func(ctx context.Context, adr string, tlsCfg *tls.Config, cfg *quic.Config) (quic.EarlyConnection, error) {
|
||||
Dial: func(ctx context.Context, adr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
// d.options.Logger.Infof("dial: %s/%s, %s", addr, network, host)
|
||||
udpAddr, err := net.ResolveUDPAddr("udp", addr)
|
||||
if err != nil {
|
||||
|
||||
@@ -74,7 +74,7 @@ func (d *wtDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOpt
|
||||
header: d.md.header,
|
||||
dialer: &wt.Dialer{
|
||||
TLSClientConfig: d.options.TLSConfig,
|
||||
DialAddr: func(ctx context.Context, adr string, tlsCfg *tls.Config, cfg *quic.Config) (quic.EarlyConnection, error) {
|
||||
DialAddr: func(ctx context.Context, adr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
// d.options.Logger.Infof("dial: %s, %s, %s", addr, adr, host)
|
||||
udpAddr, err := net.ResolveUDPAddr("udp", addr)
|
||||
if err != nil {
|
||||
@@ -97,8 +97,9 @@ func (d *wtDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOpt
|
||||
quic.Version1,
|
||||
},
|
||||
*/
|
||||
MaxIncomingStreams: int64(d.md.maxStreams),
|
||||
EnableDatagrams: true,
|
||||
MaxIncomingStreams: int64(d.md.maxStreams),
|
||||
EnableDatagrams: true,
|
||||
EnableStreamResetPartialDelivery: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
)
|
||||
|
||||
type quicSession struct {
|
||||
session quic.EarlyConnection
|
||||
session *quic.Conn
|
||||
}
|
||||
|
||||
func (session *quicSession) GetConn() (*quicConn, error) {
|
||||
@@ -28,7 +28,7 @@ func (session *quicSession) Close() error {
|
||||
}
|
||||
|
||||
type quicConn struct {
|
||||
quic.Stream
|
||||
*quic.Stream
|
||||
laddr net.Addr
|
||||
raddr net.Addr
|
||||
}
|
||||
|
||||
@@ -19,20 +19,23 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return true
|
||||
}
|
||||
return session.session.IsClosed()
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user